Download src/sft.py from Reizxn/makeitwork1: direct link, hf CLI and curl.
- Browser
- Download file 22.3 kB
-
https://huggingface.co/Reizxn/makeitwork1/resolve/main/src/sft.py
- Command line
-
hf download hf://Reizxn/makeitwork1/src/sft.py
-
curl -L -o sft.py https://huggingface.co/Reizxn/makeitwork1/resolve/main/src/sft.py
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() | |