mirror of
https://github.com/elder-plinius/OBLITERATUS.git
synced 2026-08-30 06:30:37 +02:00
Merge pull request #176 from jmagly/fix/175-qwen-multi-eos-probes
fix(qwen): preserve checkpoint multi-EOS stops
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user