""" Supervised fine-tuning (SFT) script for the search agent. Key differences from pretraining: 1. Uses the agent tokenizer (32009 vocab, includes special tokens) 2. Formats traces as chat: ...<|end|>...<|end|>...<|end|> 3. Loss masking: only compute loss on ASSISTANT tokens, not on system/user/result tokens (the model should learn to GENERATE the agent responses, not predict the inputs) 4. Lower learning rate (5e-5) — fine-tuning, not pretraining 5. Resizes model embedding to 32009 (from 32000) The trace format in the training data: system_prompt<|end|> user_query<|end|> assistant_turn_1<|end|> ← loss computed here [result injected by harness] result_content<|end|> assistant_turn_2<|end|> ← loss computed here ... evidence<|finish|><|end|> ← loss computed here Usage: python src/sft.py --steps N [--resume] [--batch_size N] [--lr F] """ import argparse import json import os import sys import time from dataclasses import asdict import numpy as np import torch import torch.nn.functional as F from tqdm import tqdm sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from model import ModelConfig, Retriever500M from tokenizers import Tokenizer # ─── Paths ─────────────────────────────────────────────────────────────────── PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) DATA_DIR = os.path.join(PROJECT_DIR, "data") TOKENIZER_DIR = os.path.join(PROJECT_DIR, "tokenizer") CHECKPOINT_DIR = os.path.join(PROJECT_DIR, "checkpoints") LOGS_DIR = os.path.join(PROJECT_DIR, "logs") TRACES_PATH = os.path.join(DATA_DIR, "sft_traces.jsonl") GOLD_PATH = os.path.join(DATA_DIR, "gold_traces.jsonl") TOKENIZER_PATH = os.path.join(TOKENIZER_DIR, "tokenizer_agent.json") SPECIAL_TOKENS_PATH = os.path.join(TOKENIZER_DIR, "special_tokens.json") # ─── Special token IDs ─────────────────────────────────────────────────────── # Loaded from special_tokens.json, but hardcoded as fallback SYSTEM_ID = 32000 USER_ID = 32001 ASSISTANT_ID = 32002 SEARCH_ID = 32003 RESULT_ID = 32004 EVIDENCE_ID = 32005 REASONING_ID = 32006 FINISH_ID = 32007 END_ID = 32008 def load_special_tokens(): """Load special token IDs from the mapping file.""" global SYSTEM_ID, USER_ID, ASSISTANT_ID, SEARCH_ID, RESULT_ID global EVIDENCE_ID, REASONING_ID, FINISH_ID, END_ID if os.path.exists(SPECIAL_TOKENS_PATH): with open(SPECIAL_TOKENS_PATH, "r") as f: data = json.load(f) ids = data["token_ids"] SYSTEM_ID = ids.get("", 32000) USER_ID = ids.get("", 32001) ASSISTANT_ID = ids.get("", 32002) SEARCH_ID = ids.get("<|search|>", 32003) RESULT_ID = ids.get("<|result|>", 32004) EVIDENCE_ID = ids.get("<|evidence|>", 32005) REASONING_ID = ids.get("<|reasoning|>", 32006) FINISH_ID = ids.get("<|finish|>", 32007) END_ID = ids.get("<|end|>", 32008) # ─── Data formatting ───────────────────────────────────────────────────────── def format_trace_to_tokens(trace: dict, tokenizer: Tokenizer, max_seq_len: int = 768) -> tuple[np.ndarray, np.ndarray]: """Convert a trace to (input_ids, loss_mask) arrays. loss_mask[i] = 1 if we should compute loss on token i, 0 otherwise. Loss is only computed on ASSISTANT turns (and the special tokens within them). """ messages = trace["trace"] all_tokens = [] loss_mask = [] for msg in messages: role = msg["role"] content = msg["content"] if role == "system": # System: content<|end|> — no loss role_id = SYSTEM_ID tokens = [role_id] + tokenizer.encode(content).ids + [END_ID] all_tokens.extend(tokens) loss_mask.extend([0] * len(tokens)) elif role == "user": # User: content<|end|> — no loss role_id = USER_ID tokens = [role_id] + tokenizer.encode(content).ids + [END_ID] all_tokens.extend(tokens) loss_mask.extend([0] * len(tokens)) elif role == "assistant": # Assistant: content<|end|> — LOSS on all tokens role_id = ASSISTANT_ID # The content already contains <|search|>, <|reasoning|>, etc. as text # We need to encode them properly tokens = [role_id] + tokenizer.encode(content).ids + [END_ID] all_tokens.extend(tokens) loss_mask.extend([1] * len(tokens)) elif role == "result": # Result: injected by harness, no loss # Format as <|result|>content<|end|> if content: tokens = [RESULT_ID] + tokenizer.encode(content).ids + [END_ID] else: tokens = [RESULT_ID, END_ID] all_tokens.extend(tokens) loss_mask.extend([0] * len(tokens)) # Truncate to max_seq_len if len(all_tokens) > max_seq_len: all_tokens = all_tokens[:max_seq_len] loss_mask = loss_mask[:max_seq_len] return np.array(all_tokens, dtype=np.int32), np.array(loss_mask, dtype=np.int32) def load_sft_dataset(traces_path: str, tokenizer: Tokenizer, max_seq_len: int = 768) -> list[tuple[np.ndarray, np.ndarray]]: """Load all SFT traces and convert to (input_ids, loss_mask) pairs.""" print(f"Loading SFT traces from {traces_path}...") dataset = [] with open(traces_path, "r", encoding="utf-8") as f: for line in f: trace = json.loads(line) ids, mask = format_trace_to_tokens(trace, tokenizer, max_seq_len) if len(ids) > 10: # skip traces that are too short dataset.append((ids, mask)) print(f" Loaded {len(dataset):,} traces") return dataset def get_sft_batch( dataset: list[tuple[np.ndarray, np.ndarray]], batch_size: int, seq_len: int, device: torch.device, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Sample a batch of SFT traces. Returns (input_ids, targets, loss_mask) where: - input_ids: (B, T) token IDs - targets: (B, T) shifted targets (next token prediction) - loss_mask: (B, T) 1 where loss should be computed, 0 elsewhere """ # Randomly sample traces indices = np.random.randint(0, len(dataset), size=batch_size) input_ids_list = [] loss_mask_list = [] for idx in indices: ids, mask = dataset[idx] # Pad or truncate to seq_len if len(ids) < seq_len: pad_len = seq_len - len(ids) ids = np.concatenate([ids, np.zeros(pad_len, dtype=np.int32)]) mask = np.concatenate([mask, np.zeros(pad_len, dtype=np.int32)]) else: ids = ids[:seq_len] mask = mask[:seq_len] input_ids_list.append(ids) loss_mask_list.append(mask) input_ids = torch.from_numpy(np.stack(input_ids_list)).long().to(device) loss_mask = torch.from_numpy(np.stack(loss_mask_list)).long().to(device) # Targets are shifted input_ids targets = torch.cat([input_ids[:, 1:], torch.zeros_like(input_ids[:, :1])], dim=1) return input_ids, targets, loss_mask # ─── Training ──────────────────────────────────────────────────────────────── def setup_optimizer(model, lr, use_8bit=True): """Same as pretraining optimizer setup.""" decay_params, no_decay_params = [], [] for name, param in model.named_parameters(): if not param.requires_grad: continue if "embedding" in name or "norm" in name: no_decay_params.append(param) else: decay_params.append(param) param_groups = [ {"params": decay_params, "weight_decay": 0.1}, {"params": no_decay_params, "weight_decay": 0.0}, ] if use_8bit: try: import bitsandbytes as bnb optimizer = bnb.optim.AdamW8bit(param_groups, lr=lr, betas=(0.9, 0.95), eps=1e-8) print("Using 8-bit AdamW (bitsandbytes)") return optimizer except Exception as e: print(f"8-bit optimizer unavailable ({e}), falling back to AdamW") optimizer = torch.optim.AdamW(param_groups, lr=lr, betas=(0.9, 0.95), eps=1e-8) print("Using standard AdamW") return optimizer def get_lr(step, warmup, max_steps, max_lr, min_lr): """Cosine LR schedule with linear warmup.""" if step < warmup: return max_lr * (step + 1) / warmup if step > max_steps: return min_lr decay_ratio = (step - warmup) / (max_steps - warmup) coeff = 0.5 * (1.0 + np.cos(np.pi * decay_ratio)) return min_lr + coeff * (max_lr - min_lr) def resize_embeddings(model, new_vocab_size): """Resize the model's token embedding to accommodate new tokens.""" old_size = model.token_embedding.weight.shape[0] if old_size == new_vocab_size: return print(f"Resizing embeddings: {old_size} -> {new_vocab_size}") d_model = model.config.d_model # Create new embedding with the old weights + new random weights old_weight = model.token_embedding.weight.data new_embedding = nn.Embedding(new_vocab_size, d_model) new_embedding.weight.data[:old_size] = old_weight # Initialize new tokens with small random values nn.init.normal_(new_embedding.weight.data[old_size:], mean=0.0, std=0.02) model.token_embedding = new_embedding # If tied, the output weight is the same, so nothing else to do # If not tied, resize lm_head too if not model.config.tie_embeddings and model.lm_head is not None: model.lm_head = nn.Linear(d_model, new_vocab_size, bias=False) with torch.no_grad(): model.lm_head.weight.data[:old_size] = old_weight nn.init.normal_(model.lm_head.weight.data[old_size:], mean=0.0, std=0.02) from torch import nn def train(args): load_special_tokens() device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device}") if device.type == "cuda": print(f"GPU: {torch.cuda.get_device_name(0)}") os.makedirs(CHECKPOINT_DIR, exist_ok=True) os.makedirs(LOGS_DIR, exist_ok=True) # ─── Tokenizer ─────────────────────────────────────────────────────────── print("Loading agent tokenizer...") tokenizer = Tokenizer.from_file(TOKENIZER_PATH) vocab_size = tokenizer.get_vocab_size() print(f"Vocab size: {vocab_size}") # ─── Data ──────────────────────────────────────────────────────────────── dataset = load_sft_dataset(TRACES_PATH, tokenizer, args.seq_len) # Also load gold traces gold_dataset = load_sft_dataset(GOLD_PATH, tokenizer, args.seq_len) dataset.extend(gold_dataset) print(f" Total (with gold): {len(dataset):,}") # ─── Model ─────────────────────────────────────────────────────────────── config = ModelConfig( vocab_size=vocab_size, # 32009 d_model=1_280, n_layers=23, n_heads=20, d_ff=3_456, max_seq_len=args.seq_len, dropout=0.0, tie_embeddings=True, ) model = Retriever500M(config).to(device) # ─── Load pretrained checkpoint ────────────────────────────────────────── ckpt_path = args.resume_path or os.path.join(CHECKPOINT_DIR, "latest.pt") if os.path.exists(ckpt_path): print(f"Loading pretrained weights from {ckpt_path}...") ckpt = torch.load(ckpt_path, map_location=device, weights_only=False) old_config = ModelConfig(**ckpt["config"]) # Load state dict, handling vocab size mismatch state_dict = ckpt["model_state_dict"] old_vocab = old_config.vocab_size if old_vocab != vocab_size: print(f" Vocab size mismatch: {old_vocab} -> {vocab_size}") print(f" Resizing embeddings in state dict...") # Resize token_embedding in the state dict old_weight = state_dict["token_embedding.weight"] d_model = old_weight.shape[1] new_weight = torch.zeros(vocab_size, d_model) new_weight[:old_vocab] = old_weight nn.init.normal_(new_weight[old_vocab:], mean=0.0, std=0.02) state_dict["token_embedding.weight"] = new_weight model.load_state_dict(state_dict) print(f" Loaded (step {ckpt.get('step', '?')}, loss {ckpt.get('loss', '?')})") else: print(f"WARNING: No checkpoint at {ckpt_path}, starting from scratch!") total_params = model.count_parameters() print(f"Model parameters: {total_params:,} ({total_params / 1e6:.1f}M)") # ─── Optimizer ─────────────────────────────────────────────────────────── optimizer = setup_optimizer(model, args.lr, use_8bit=args.use_8bit_adam) # ─── Training loop ─────────────────────────────────────────────────────── effective_batch = args.batch_size * args.grad_accum print(f"\nSFT configuration:") print(f" Batch size: {args.batch_size}") print(f" Grad accum: {args.grad_accum}") print(f" Effective batch: {effective_batch}") print(f" Sequence length: {args.seq_len}") print(f" Learning rate: {args.lr}") print(f" Steps: {args.steps}") print(f" Warmup: {args.warmup}") print() log = { "config": asdict(config), "train_args": vars(args), "total_params": total_params, "steps": [], } model.train() start_time = time.time() accum_loss = 0.0 best_loss = float("inf") pbar = tqdm(range(args.steps), desc="SFT") for step in pbar: lr = get_lr(step, args.warmup, args.steps, args.lr, args.lr * 0.1) for pg in optimizer.param_groups: pg["lr"] = lr optimizer.zero_grad(set_to_none=True) total_loss = 0.0 for micro_step in range(args.grad_accum): input_ids, targets, loss_mask = get_sft_batch( dataset, args.batch_size, args.seq_len, device ) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): out = model(input_ids, targets=targets, use_checkpoint=args.grad_checkpoint) loss = out["loss"] # Apply loss mask — only compute loss on assistant tokens # loss is already computed over all positions; we need to recompute # with the mask logits = out["logits"] # Recompute loss with mask if loss_mask.sum() > 0: # Shift mask to align with next-token prediction shifted_mask = loss_mask[:, 1:].contiguous() masked_logits = logits[:, :-1, :].contiguous() masked_targets = targets[:, :-1].contiguous() # Flatten and apply mask flat_logits = masked_logits.view(-1, masked_logits.size(-1)) flat_targets = masked_targets.view(-1) flat_mask = shifted_mask.view(-1).float() per_token_loss = F.cross_entropy( flat_logits, flat_targets, ignore_index=-100, reduction="none" ) masked_loss = (per_token_loss * flat_mask).sum() / flat_mask.sum().clamp(min=1) loss = masked_loss / args.grad_accum else: loss = loss / args.grad_accum loss.backward() total_loss += loss.item() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() avg_loss = total_loss accum_loss = accum_loss * 0.95 + avg_loss * 0.05 if step % args.log_every == 0 or step == args.steps - 1: elapsed = time.time() - start_time steps_per_sec = (step + 1) / elapsed vram_used = torch.cuda.max_memory_allocated() / 1e9 if device.type == "cuda" else 0 log_entry = { "step": step, "loss": avg_loss, "ema_loss": accum_loss, "lr": lr, "elapsed_s": elapsed, "steps_per_sec": steps_per_sec, "vram_gb": vram_used, } log["steps"].append(log_entry) pbar.set_postfix({ "loss": f"{avg_loss:.4f}", "ema": f"{accum_loss:.4f}", "lr": f"{lr:.2e}", "vram": f"{vram_used:.1f}G", }) if (step + 1) % args.save_every == 0 or step == args.steps - 1: ckpt_path = os.path.join(CHECKPOINT_DIR, f"sft_step_{step + 1}.pt") torch.save({ "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, ckpt_path) print(f"\n Saved checkpoint: {ckpt_path}") latest_path = os.path.join(CHECKPOINT_DIR, "sft_latest.pt") torch.save({ "model_state_dict": model.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, latest_path) if accum_loss < best_loss: best_loss = accum_loss best_path = os.path.join(CHECKPOINT_DIR, "sft_best.pt") torch.save({ "model_state_dict": model.state_dict(), "config": asdict(config), "step": step + 1, "loss": accum_loss, }, best_path) if step % 50 == 0 and device.type == "cuda": torch.cuda.reset_peak_memory_stats() # Save log log_path = os.path.join(LOGS_DIR, "sft_log.json") with open(log_path, "w") as f: json.dump(log, f, indent=2) total_time = time.time() - start_time print(f"\nSFT complete!") print(f" Total time: {total_time:.1f}s ({total_time/60:.1f} min)") print(f" Final EMA loss: {accum_loss:.4f}") print(f" Best loss: {best_loss:.4f}") def main(): parser = argparse.ArgumentParser(description="SFT the search agent") parser.add_argument("--steps", type=int, default=500, help="Total SFT steps") parser.add_argument("--batch_size", type=int, default=4, help="Micro batch size") parser.add_argument("--grad_accum", type=int, default=4, help="Gradient accumulation") parser.add_argument("--seq_len", type=int, default=768, help="Sequence length") parser.add_argument("--lr", type=float, default=5e-5, help="Peak learning rate") parser.add_argument("--warmup", type=int, default=20, help="Warmup steps") parser.add_argument("--save_every", type=int, default=100, help="Save every N steps") parser.add_argument("--log_every", type=int, default=10, help="Log every N steps") parser.add_argument("--grad_checkpoint", action="store_true", default=True) parser.add_argument("--no_grad_checkpoint", dest="grad_checkpoint", action="store_false") parser.add_argument("--use_8bit_adam", action="store_true", default=True) parser.add_argument("--no_8bit_adam", dest="use_8bit_adam", action="store_false") parser.add_argument("--resume_path", type=str, default=None, help="Pretrained checkpoint to start from") parser.add_argument("--project_dir", type=str, default=None, help="Override project directory (for Colab)") parser.add_argument("--save_dir", type=str, default=None, help="Override checkpoint save directory (for Google Drive)") args = parser.parse_args() # Override paths for Colab global PROJECT_DIR, DATA_DIR, TOKENIZER_DIR, CHECKPOINT_DIR, LOGS_DIR global TRACES_PATH, GOLD_PATH, TOKENIZER_PATH, SPECIAL_TOKENS_PATH if args.project_dir: PROJECT_DIR = args.project_dir DATA_DIR = os.path.join(PROJECT_DIR, "data") TOKENIZER_DIR = os.path.join(PROJECT_DIR, "tokenizer") LOGS_DIR = os.path.join(PROJECT_DIR, "logs") TRACES_PATH = os.path.join(DATA_DIR, "sft_traces.jsonl") GOLD_PATH = os.path.join(DATA_DIR, "gold_traces.jsonl") TOKENIZER_PATH = os.path.join(TOKENIZER_DIR, "tokenizer_agent.json") SPECIAL_TOKENS_PATH = os.path.join(TOKENIZER_DIR, "special_tokens.json") if args.save_dir: CHECKPOINT_DIR = args.save_dir train(args) if __name__ == "__main__": main()