fromsafetune.hardenimportConstrainedSFTHFTrainer,ConstrainedSFTConfigconfig=ConstrainedSFTConfig(output_dir="csft_out",csft_beta=0.5,csft_decay_rate=0.1)trainer=ConstrainedSFTHFTrainer(model=model,args=config,train_dataset=task_ds,reference_model=ref_model,# frozen aligned model, before fine-tuning)trainer.train()
ConstrainedSFTHFTrainer is the transformers.Trainer subclass behind
harden.ConstrainedSFTTrainer. The high-level trainer (also used by the CLI)
accepts csft_beta / csft_decay_rate as keyword arguments, loads the frozen
reference from reference_model_path and passes it to
ConstrainedSFTHFTrainer, so the KL constraint is on. Before, it passed no
reference and trained plain SFT; use_reference=False or
safetune.configure(legacy_constrained_sft=True) gives that behaviour.
Best for: a lightweight KL-regularized SFT that penalizes first-token drift from the aligned model.
Trade-offs: Uses a KL-regularized SFT with a position-decaying first-token penalty rather than the paper's bounded-DPO Eq. 3 + step-function β schedule; trains cleanly but cite the implementation, not the paper name.
@article{constrainedsft2024,title={Safety Alignment Should Be Made More Than Just a Few Tokens Deep},author={Qi, et al.},year={2024},note={ICLR 2025, arXiv:2406.05946},}