| """SFT training script v2 for HBRT safety project. |
| Supports: weighted loss (structure/transitions), NEFTune, completion curriculum. |
| Designed for 8-GPU node (accelerate launch --num_processes 8). |
| """ |
| import argparse |
| import os |
| import re |
| import pandas as pd |
| import numpy as np |
| from datasets import Dataset |
| from transformers import AutoTokenizer, AutoModelForCausalLM |
| from trl import SFTTrainer, SFTConfig |
| import torch |
| import torch.nn.functional as F |
|
|
|
|
| def parse_args(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--track-id", type=int, required=True) |
| parser.add_argument("--track-name", type=str, required=True) |
| parser.add_argument("--data-path", type=str, default="data/sft-training-data/convergent_data_10k.parquet") |
| parser.add_argument("--model-path", type=str, default="models/starting-checkpoint") |
| parser.add_argument("--output-dir", type=str, default=None) |
| parser.add_argument("--lr", type=float, default=5e-5) |
| parser.add_argument("--epochs", type=int, default=10) |
| parser.add_argument("--save-epochs", type=str, default="3,5,7,10", help="Comma-separated epochs to save at") |
| parser.add_argument("--max-seq-length", type=int, default=8192) |
| parser.add_argument("--per-device-batch-size", type=int, default=2) |
| parser.add_argument("--gradient-accumulation-steps", type=int, default=2) |
| parser.add_argument("--warmup-ratio", type=float, default=0.05) |
| parser.add_argument("--val-split", type=float, default=0.1) |
| parser.add_argument("--subset", type=str, default=None) |
| parser.add_argument("--seed", type=int, default=42) |
| parser.add_argument("--resume-from", type=str, default=None) |
| parser.add_argument("--lr-scheduler", type=str, default="linear", choices=["linear", "cosine"]) |
| |
| parser.add_argument("--loss-weight-mode", type=str, default=None, |
| choices=[None, "structure", "transitions"], |
| help="structure: 5x on safety_check+think tokens. transitions: 5x on 4 transition tokens only.") |
| parser.add_argument("--loss-weight-multiplier", type=float, default=5.0) |
| |
| parser.add_argument("--neftune-alpha", type=float, default=None, help="NEFTune noise alpha (e.g. 5.0)") |
| |
| parser.add_argument("--max-token-length", type=int, default=None, help="Filter examples longer than this") |
| |
| parser.add_argument("--special-tokens", type=str, default=None, |
| choices=[None, "structure", "all"], |
| help="structure: register 6 key XML boundary tokens. all: register all XML tags.") |
| return parser.parse_args() |
|
|
|
|
| def load_data(path, tokenizer, subset=None, val_split=0.1, seed=42, max_token_length=None): |
| df = pd.read_parquet(path) |
|
|
| if subset == "harmful": |
| df = df[df["gold_label"] == "harmful"] |
| elif subset == "benign": |
| df = df[df["gold_label"] == "benign"] |
| elif subset == "high-confidence": |
| df = df[(df["harmfulness_score"] > 0.8) | (df["harmfulness_score"] < 0.15)] |
| elif subset and subset.startswith("source:"): |
| prefix = subset.split(":", 1)[1] |
| df = df[df["model"].str.startswith(prefix)] |
|
|
| messages_list = [] |
| for _, row in df.iterrows(): |
| messages_list.append([ |
| {"role": "user", "content": row["prompt"]}, |
| {"role": "assistant", "content": row["convergent_final_sft_format"]} |
| ]) |
|
|
| if max_token_length: |
| filtered = [] |
| for msgs in messages_list: |
| text = tokenizer.apply_chat_template(msgs, tokenize=False) |
| toks = tokenizer.encode(text) |
| if len(toks) <= max_token_length: |
| filtered.append(msgs) |
| print(f" Length filter: {len(filtered)}/{len(messages_list)} kept (max {max_token_length} tokens)") |
| messages_list = filtered |
|
|
| ds = Dataset.from_dict({"messages": messages_list}) |
| split = ds.train_test_split(test_size=val_split, seed=seed) |
| return split["train"], split["test"] |
|
|
|
|
| class WeightedLossTrainer(SFTTrainer): |
| """SFTTrainer subclass that applies per-token loss weighting.""" |
|
|
| def __init__(self, *args, loss_weight_mode=None, loss_weight_multiplier=5.0, |
| weight_tokenizer=None, **kwargs): |
| super().__init__(*args, **kwargs) |
| self.loss_weight_mode = loss_weight_mode |
| self.loss_weight_multiplier = loss_weight_multiplier |
| self.weight_tokenizer = weight_tokenizer |
|
|
| if loss_weight_mode and weight_tokenizer: |
| self._precompute_marker_ids() |
|
|
| def _precompute_marker_ids(self): |
| tok = self.weight_tokenizer |
| if self.loss_weight_mode == "transitions": |
| markers = ["</safety_check>", "<safety_check_score>", "<think>", "</think>"] |
| self.marker_token_ids = set() |
| for m in markers: |
| ids = tok.encode(m, add_special_tokens=False) |
| self.marker_token_ids.update(ids) |
| elif self.loss_weight_mode == "structure": |
| markers_start = ["<safety_check>", "<think>"] |
| markers_end = ["</safety_check>", "</think>"] |
| self.structure_start_ids = [] |
| self.structure_end_ids = [] |
| for m in markers_start: |
| self.structure_start_ids.append(tok.encode(m, add_special_tokens=False)) |
| for m in markers_end: |
| self.structure_end_ids.append(tok.encode(m, add_special_tokens=False)) |
|
|
| def compute_loss(self, model, inputs, return_outputs=False, **kwargs): |
| if self.loss_weight_mode is None: |
| return super().compute_loss(model, inputs, return_outputs=return_outputs, **kwargs) |
|
|
| labels = inputs.pop("labels") |
| outputs = model(**inputs) |
| logits = outputs.logits |
|
|
| shift_logits = logits[..., :-1, :].contiguous() |
| shift_labels = labels[..., 1:].contiguous() |
|
|
| loss_fct = torch.nn.CrossEntropyLoss(reduction="none") |
| flat_logits = shift_logits.view(-1, shift_logits.size(-1)) |
| flat_labels = shift_labels.view(-1) |
| per_token_loss = loss_fct(flat_logits, flat_labels) |
|
|
| |
| weights = torch.ones_like(per_token_loss) |
| valid_mask = flat_labels != -100 |
|
|
| if self.loss_weight_mode == "transitions": |
| token_ids = flat_labels.clone() |
| for tid in self.marker_token_ids: |
| weights[token_ids == tid] = self.loss_weight_multiplier |
| elif self.loss_weight_mode == "structure": |
| |
| batch_size = shift_labels.size(0) |
| seq_len = shift_labels.size(1) |
| weight_2d = torch.ones(batch_size, seq_len, device=shift_labels.device) |
|
|
| for b in range(batch_size): |
| seq = shift_labels[b] |
| seq_list = seq.tolist() |
| in_block = False |
| for i, tid in enumerate(seq_list): |
| if tid == -100: |
| continue |
| |
| for start_ids in self.structure_start_ids: |
| if i + len(start_ids) <= len(seq_list): |
| if seq_list[i:i+len(start_ids)] == start_ids: |
| in_block = True |
| break |
| if in_block: |
| weight_2d[b, i] = self.loss_weight_multiplier |
| |
| for end_ids in self.structure_end_ids: |
| if i >= len(end_ids) - 1: |
| if seq_list[i-len(end_ids)+1:i+1] == end_ids: |
| in_block = False |
| break |
|
|
| weights = weight_2d.view(-1) |
|
|
| |
| weights = weights * valid_mask.float() |
| loss = (per_token_loss * weights).sum() / weights.sum() |
|
|
| return (loss, outputs) if return_outputs else loss |
|
|
|
|
| def main(): |
| args = parse_args() |
|
|
| output_dir = args.output_dir or f"models/track-{args.track_id}-{args.track_name}" |
| os.makedirs(output_dir, exist_ok=True) |
|
|
| model_path = args.resume_from if args.resume_from else args.model_path |
|
|
| print(f"=== Track {args.track_id}: {args.track_name} ===") |
| print(f" LR: {args.lr}, Epochs: {args.epochs}, Scheduler: {args.lr_scheduler}") |
| print(f" Data: {args.data_path}, Subset: {args.subset}") |
| print(f" Model: {model_path}") |
| print(f" Output: {output_dir}") |
| print(f" Loss weight: {args.loss_weight_mode} (x{args.loss_weight_multiplier})") |
| print(f" NEFTune: {args.neftune_alpha}") |
| print(f" Save at epochs: {args.save_epochs}") |
|
|
| tokenizer = AutoTokenizer.from_pretrained(model_path) |
| if tokenizer.pad_token is None: |
| tokenizer.pad_token = tokenizer.eos_token |
|
|
| |
| if args.special_tokens: |
| if args.special_tokens == "structure": |
| new_tokens = [ |
| "<safety_check>", "</safety_check>", |
| "<safety_check_score>", "</safety_check_score>", |
| "<think>", "</think>", |
| ] |
| else: |
| new_tokens = [ |
| "<safety_check>", "</safety_check>", |
| "<safety_check_score>", "</safety_check_score>", |
| "<think>", "</think>", |
| "<stakeholder>", "</stakeholder>", |
| "<harms>", "</harms>", |
| "<benefits>", "</benefits>", |
| "<harm_score>", "</harm_score>", |
| "<benefit_score>", "</benefit_score>", |
| "<action>", "</action>", |
| "<action_name>", "</action_name>", |
| "<effects>", "</effects>", |
| "<effect>", "</effect>", |
| "<effect_name>", "</effect_name>", |
| "<immediacy>", "</immediacy>", |
| "<extent>", "</extent>", |
| "<likelihood>", "</likelihood>", |
| "<effect_score>", "</effect_score>", |
| "<harms_total>", "</harms_total>", |
| "<benefits_total>", "</benefits_total>", |
| "<raw_score>", "</raw_score>", |
| "<final_score>", "</final_score>", |
| "<label>", "</label>", |
| ] |
| num_added = tokenizer.add_special_tokens({"additional_special_tokens": new_tokens}) |
| print(f" Added {num_added} special tokens: {args.special_tokens}") |
|
|
| model = AutoModelForCausalLM.from_pretrained( |
| model_path, |
| torch_dtype=torch.bfloat16, |
| attn_implementation="sdpa", |
| ) |
|
|
| if args.special_tokens: |
| model.resize_token_embeddings(len(tokenizer)) |
|
|
| train_ds, val_ds = load_data( |
| args.data_path, tokenizer, subset=args.subset, |
| val_split=args.val_split, seed=args.seed, |
| max_token_length=args.max_token_length |
| ) |
| print(f" Train: {len(train_ds)}, Val: {len(val_ds)}") |
|
|
| |
| save_epochs = [int(x) for x in args.save_epochs.split(",")] |
|
|
| |
| num_gpus = int(os.environ.get("WORLD_SIZE", torch.cuda.device_count())) |
| steps_per_epoch = len(train_ds) // (args.per_device_batch_size * args.gradient_accumulation_steps * num_gpus) |
| if steps_per_epoch == 0: |
| steps_per_epoch = 1 |
| save_steps = [e * steps_per_epoch for e in save_epochs] |
| print(f" Steps/epoch: {steps_per_epoch}, Save at steps: {save_steps}") |
|
|
| sft_kwargs = {} |
| if args.neftune_alpha: |
| sft_kwargs["neftune_noise_alpha"] = args.neftune_alpha |
|
|
| training_args = SFTConfig( |
| output_dir=output_dir, |
| num_train_epochs=args.epochs, |
| per_device_train_batch_size=args.per_device_batch_size, |
| per_device_eval_batch_size=args.per_device_batch_size, |
| gradient_accumulation_steps=args.gradient_accumulation_steps, |
| learning_rate=args.lr, |
| lr_scheduler_type=args.lr_scheduler, |
| warmup_ratio=args.warmup_ratio, |
| bf16=True, |
| logging_steps=10, |
| eval_strategy="epoch", |
| save_strategy="steps", |
| save_steps=save_steps[0] if save_steps else steps_per_epoch, |
| save_total_limit=len(save_epochs) + 1, |
| max_length=args.max_seq_length, |
| seed=args.seed, |
| report_to="none", |
| **sft_kwargs, |
| ) |
|
|
| TrainerClass = WeightedLossTrainer if args.loss_weight_mode else SFTTrainer |
| trainer_kwargs = {} |
| if args.loss_weight_mode: |
| trainer_kwargs["loss_weight_mode"] = args.loss_weight_mode |
| trainer_kwargs["loss_weight_multiplier"] = args.loss_weight_multiplier |
| trainer_kwargs["weight_tokenizer"] = tokenizer |
|
|
| trainer = TrainerClass( |
| model=model, |
| processing_class=tokenizer, |
| train_dataset=train_ds, |
| eval_dataset=val_ds, |
| args=training_args, |
| **trainer_kwargs, |
| ) |
|
|
| |
| from transformers import TrainerCallback |
|
|
| class EpochSaveCallback(TrainerCallback): |
| def __init__(self, save_epochs, steps_per_epoch, output_dir, tokenizer_ref): |
| self.save_epochs = save_epochs |
| self.steps_per_epoch = steps_per_epoch |
| self.output_dir = output_dir |
| self.tokenizer_ref = tokenizer_ref |
| self.saved_epochs = set() |
|
|
| def on_step_end(self, args, state, control, **kwargs): |
| current_epoch = int(state.epoch) if state.epoch else 0 |
| if current_epoch in self.save_epochs and current_epoch not in self.saved_epochs: |
| if abs(state.epoch - current_epoch) < 0.01: |
| control.should_save = True |
| self.saved_epochs.add(current_epoch) |
| else: |
| control.should_save = False |
| return control |
|
|
| def on_save(self, args, state, control, **kwargs): |
| |
| ckpt_dir = os.path.join(args.output_dir, f"checkpoint-{state.global_step}") |
| if os.path.isdir(ckpt_dir): |
| self.tokenizer_ref.save_pretrained(ckpt_dir) |
| return control |
|
|
| def on_train_end(self, args, state, control, model=None, tokenizer=None, **kwargs): |
| |
| final_epoch = max(self.save_epochs) |
| if final_epoch not in self.saved_epochs: |
| save_dir = os.path.join(self.output_dir, f"epoch-{final_epoch}") |
| os.makedirs(save_dir, exist_ok=True) |
| if model is not None: |
| model.save_pretrained(save_dir) |
| self.tokenizer_ref.save_pretrained(save_dir) |
| self.saved_epochs.add(final_epoch) |
| print(f" [on_train_end] Saved epoch-{final_epoch} to {save_dir}") |
| return control |
|
|
| |
| training_args.save_strategy = "no" |
| trainer.add_callback(EpochSaveCallback(save_epochs, steps_per_epoch, output_dir, tokenizer)) |
|
|
| trainer.train() |
|
|
| |
| local_rank = int(os.environ.get("LOCAL_RANK", 0)) |
| if local_rank == 0: |
| for entry in sorted(os.listdir(output_dir)): |
| if entry.startswith("checkpoint-"): |
| step = int(entry.split("-")[1]) |
| epoch_idx = round(step / steps_per_epoch) |
| epoch_dir = os.path.join(output_dir, f"epoch-{epoch_idx}") |
| ckpt_dir = os.path.join(output_dir, entry) |
| if not os.path.exists(epoch_dir): |
| os.rename(ckpt_dir, epoch_dir) |
| print(f" Renamed {entry} -> epoch-{epoch_idx}") |
|
|
| print(f"=== Track {args.track_id} complete ===") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|