| """
|
| 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: <tool_call>...<|end|><tool_call>...<|end|><tool_call>...<|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:
|
| <tool_call>system_prompt<|end|>
|
| <tool_call>user_query<|end|>
|
| <tool_call>assistant_turn_1<|end|> ← loss computed here
|
| [result injected by harness]
|
| <tool_call>result_content<|end|>
|
| <tool_call>assistant_turn_2<|end|> ← loss computed here
|
| ...
|
| <tool_call>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
|
|
|
|
|
| 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")
|
|
|
|
|
|
|
| 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("<tool_call>", 32000)
|
| USER_ID = ids.get("<tool_call>", 32001)
|
| ASSISTANT_ID = ids.get("<tool_call>", 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)
|
|
|
|
|
|
|
|
|
| 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":
|
|
|
| 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":
|
|
|
| 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":
|
|
|
| role_id = ASSISTANT_ID
|
|
|
|
|
| tokens = [role_id] + tokenizer.encode(content).ids + [END_ID]
|
| all_tokens.extend(tokens)
|
| loss_mask.extend([1] * len(tokens))
|
|
|
| elif role == "result":
|
|
|
|
|
| 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))
|
|
|
|
|
| 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:
|
| 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
|
| """
|
|
|
| indices = np.random.randint(0, len(dataset), size=batch_size)
|
|
|
| input_ids_list = []
|
| loss_mask_list = []
|
|
|
| for idx in indices:
|
| ids, mask = dataset[idx]
|
|
|
| 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 = torch.cat([input_ids[:, 1:], torch.zeros_like(input_ids[:, :1])], dim=1)
|
|
|
| return input_ids, targets, loss_mask
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
| old_weight = model.token_embedding.weight.data
|
| new_embedding = nn.Embedding(new_vocab_size, d_model)
|
| new_embedding.weight.data[:old_size] = old_weight
|
|
|
| nn.init.normal_(new_embedding.weight.data[old_size:], mean=0.0, std=0.02)
|
| model.token_embedding = new_embedding
|
|
|
|
|
|
|
| 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)
|
|
|
|
|
| print("Loading agent tokenizer...")
|
| tokenizer = Tokenizer.from_file(TOKENIZER_PATH)
|
| vocab_size = tokenizer.get_vocab_size()
|
| print(f"Vocab size: {vocab_size}")
|
|
|
|
|
| dataset = load_sft_dataset(TRACES_PATH, tokenizer, args.seq_len)
|
|
|
|
|
| gold_dataset = load_sft_dataset(GOLD_PATH, tokenizer, args.seq_len)
|
| dataset.extend(gold_dataset)
|
| print(f" Total (with gold): {len(dataset):,}")
|
|
|
|
|
| config = ModelConfig(
|
| vocab_size=vocab_size,
|
| 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)
|
|
|
|
|
| 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"])
|
|
|
|
|
| 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...")
|
|
|
| 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 = setup_optimizer(model, args.lr, use_8bit=args.use_8bit_adam)
|
|
|
|
|
| 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"]
|
|
|
|
|
|
|
|
|
| logits = out["logits"]
|
|
|
| if loss_mask.sum() > 0:
|
|
|
| shifted_mask = loss_mask[:, 1:].contiguous()
|
| masked_logits = logits[:, :-1, :].contiguous()
|
| masked_targets = targets[:, :-1].contiguous()
|
|
|
|
|
| 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()
|
|
|
|
|
| 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()
|
|
|
|
|
| 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()
|
|
|