diff --git a/obliteratus/abliterate.py b/obliteratus/abliterate.py index be5eae4..de87d40 100644 --- a/obliteratus/abliterate.py +++ b/obliteratus/abliterate.py @@ -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: diff --git a/tests/test_abliterate_extended.py b/tests/test_abliterate_extended.py index 33c7a5c..03bda27 100644 --- a/tests/test_abliterate_extended.py +++ b/tests/test_abliterate_extended.py @@ -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(