SupritiVijay's picture
Add track-codes (training scripts + eval)
ae6dcfe verified
Raw
History Blame Contribute Delete
15.9 kB
"""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 = ["</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)
# 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 <safety_check>...</safety_check> and <think>...</think>
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 = [
"<safety_check>", "</safety_check>",
"<safety_check_score>", "</safety_check_score>",
"<think>", "</think>",
]
else: # "all"
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)}")
# 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()