"""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"]) # Weighted loss 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) # NEFTune parser.add_argument("--neftune-alpha", type=float, default=None, help="NEFTune noise alpha (e.g. 5.0)") # Length filter for curriculum parser.add_argument("--max-token-length", type=int, default=None, help="Filter examples longer than this") # Special tokens 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 = ["", "", "", ""] 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 = ["", ""] markers_end = ["", ""] 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) # Build weight mask 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": # Weight all tokens inside ... and ... 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 # Check if we're entering a structure block 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 # Check if we're exiting 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) # Apply weights only to valid tokens 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 # Register special tokens if requested if args.special_tokens: if args.special_tokens == "structure": new_tokens = [ "", "", "", "", "", "", ] else: # "all" new_tokens = [ "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "", ] 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)}") # Parse save epochs save_epochs = [int(x) for x in args.save_epochs.split(",")] # Calculate steps per epoch for save_steps 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, ) # Custom save callback to only save at desired epochs 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): # Ensure tokenizer (with any added special tokens) is saved alongside model 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): # Force save final epoch — on_step_end can't trigger saves on the last step 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 # Override save strategy with our callback training_args.save_strategy = "no" trainer.add_callback(EpochSaveCallback(save_epochs, steps_per_epoch, output_dir, tokenizer)) trainer.train() # Rename checkpoint dirs to epoch-{IDX} (only on main process) 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()