fix(qwen): preserve checkpoint multi-EOS stops

This commit is contained in:
Joseph Magly
2026-08-26 23:45:01 -04:00
parent a1de7ff473
commit 270db39e27
2 changed files with 48 additions and 4 deletions
+40 -1
View File
@@ -10,6 +10,7 @@ Tests the new capabilities added to the OBLITERATUS abliteration pipeline:
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
@@ -244,7 +245,12 @@ class TestChatTemplate:
def test_deterministic_generation_uses_explicit_eos_and_pad_fallback(self):
pipeline = AbliterationPipeline(model_name="chat-model")
tokenizer = MagicMock(eos_token_id=42, pad_token_id=None)
pipeline.handle = MagicMock(architecture="llama", tokenizer=tokenizer)
model = SimpleNamespace(
generation_config=SimpleNamespace(eos_token_id=None, pad_token_id=None)
)
pipeline.handle = MagicMock(
architecture="llama", tokenizer=tokenizer, model=model
)
assert pipeline._deterministic_generation_kwargs(64) == {
"max_new_tokens": 64,
@@ -253,6 +259,39 @@ class TestChatTemplate:
"pad_token_id": 42,
}
def test_deterministic_generation_preserves_model_multi_eos_contract(self):
pipeline = AbliterationPipeline(model_name="Qwen/Qwen3.8-27B")
tokenizer = MagicMock(eos_token_id=248044, pad_token_id=None)
model = SimpleNamespace(
generation_config=SimpleNamespace(
eos_token_id=[248046, 248044],
pad_token_id=248044,
)
)
pipeline.handle = MagicMock(
architecture="qwen3_5", tokenizer=tokenizer, model=model
)
assert pipeline._deterministic_generation_kwargs(64) == {
"max_new_tokens": 64,
"do_sample": False,
"eos_token_id": [248046, 248044],
"pad_token_id": 248044,
}
def test_deterministic_generation_omits_absent_stop_tokens(self):
pipeline = AbliterationPipeline(model_name="bare-model")
tokenizer = SimpleNamespace(eos_token_id=None, pad_token_id=None)
pipeline.handle = SimpleNamespace(
tokenizer=tokenizer,
model=SimpleNamespace(generation_config=None),
)
assert pipeline._deterministic_generation_kwargs(12) == {
"max_new_tokens": 12,
"do_sample": False,
}
def test_no_wrap_when_disabled(self):
"""Should not wrap prompts when use_chat_template is False."""
pipeline = AbliterationPipeline(