makeitwork1 / src /sft.py
Reizxn's picture
Upload folder using huggingface_hub
803b5e8 verified
Raw
History Blame Contribute Delete
22.3 kB
"""
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
# ─── 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("<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)
# ─── 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: <tool_call>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: <tool_call>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: <tool_call>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()