add configured max tokens for refusal test generation. Exposes --refusal-max-tokens via CLI and the abliterator constructor (default to 128, not changed). Better for resoning models like qwen3. :>

This commit is contained in:
Tokard
2026-08-14 20:03:46 -04:00
committed by Joseph Magly
parent 9746c21b63
commit 44473e5919
2 changed files with 10 additions and 1 deletions
+3 -1
View File
@@ -843,6 +843,7 @@ class AbliterationPipeline:
max_seq_length: int | None = None,
# Verify stage sample size
verify_sample_size: int | None = None,
refusal_max_tokens: int | None = None,
on_stage: Callable[[StageResult], None] | None = None,
on_log: Callable[[str], None] | None = None,
):
@@ -972,6 +973,7 @@ class AbliterationPipeline:
# refusal rate measurement. Default 30 gives ~3.3% resolution;
# increase for tighter confidence intervals (reviewer feedback).
self.verify_sample_size = verify_sample_size if verify_sample_size is not None else 30
self.refusal_max_tokens = refusal_max_tokens if refusal_max_tokens is not None else 128
# Large model mode: conservative defaults for 120B+ models.
# Reduces memory footprint by limiting SAE features, directions,
@@ -6338,7 +6340,7 @@ class AbliterationPipeline:
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=128,
max_new_tokens=self.refusal_max_tokens,
do_sample=False,
)
+7
View File
@@ -240,6 +240,10 @@ def main(argv: list[str] | None = None):
help="Number of harmful prompts to test for refusal rate (default: 30). "
"Increase for tighter confidence intervals (e.g. 100 for ~1%% resolution).",
)
p.add_argument(
"--refusal-max-tokens", type=int, default=None,
help="Max new tokens to generate per response in the refusal test (default: 128).",
)
p.add_argument(
"--dataset", type=str, default="builtin",
help="Prompt dataset source for contrastive extraction when using residue mining (default: builtin).",
@@ -1043,6 +1047,7 @@ def _cmd_abliterate(args):
quantization=args.quantization,
large_model_mode=getattr(args, "large_model", False),
verify_sample_size=getattr(args, "verify_sample_size", None),
refusal_max_tokens=getattr(args, "refusal_max_tokens", None),
on_stage=on_stage,
on_log=on_log,
**prompt_kwargs,
@@ -1368,6 +1373,8 @@ def _cmd_remote_abliterate(args):
kwargs["large_model"] = True
if getattr(args, "verify_sample_size", None) is not None:
kwargs["verify_sample_size"] = args.verify_sample_size
if getattr(args, "refusal_max_tokens", None) is not None:
kwargs["refusal_max_tokens"] = args.refusal_max_tokens
result_path = runner.run_obliterate(
model=args.model,