diff --git a/obliteratus/abliterate.py b/obliteratus/abliterate.py index 8cb8249..bad1873 100644 --- a/obliteratus/abliterate.py +++ b/obliteratus/abliterate.py @@ -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, ) diff --git a/obliteratus/cli.py b/obliteratus/cli.py index 4bb9996..c762d61 100644 --- a/obliteratus/cli.py +++ b/obliteratus/cli.py @@ -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,