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
+8 -3
View File
@@ -1904,14 +1904,19 @@ class AbliterationPipeline:
return prompt
def _deterministic_generation_kwargs(self, max_new_tokens: int) -> dict[str, Any]:
"""Return explicit, tokenizer-compatible settings for quality probes."""
"""Return deterministic settings without overriding checkpoint stop tokens."""
tokenizer = self.handle.tokenizer
generation_config = getattr(self.handle.model, "generation_config", None)
kwargs: dict[str, Any] = {
"max_new_tokens": max_new_tokens,
"do_sample": False,
}
eos_token_id = getattr(tokenizer, "eos_token_id", None)
pad_token_id = getattr(tokenizer, "pad_token_id", None)
eos_token_id = getattr(generation_config, "eos_token_id", None)
pad_token_id = getattr(generation_config, "pad_token_id", None)
if eos_token_id is None:
eos_token_id = getattr(tokenizer, "eos_token_id", None)
if pad_token_id is None:
pad_token_id = getattr(tokenizer, "pad_token_id", None)
if eos_token_id is not None:
kwargs["eos_token_id"] = eos_token_id
if pad_token_id is not None:
+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(