""" Nova 1.0 — Supervised Fine-Tuning (SFT) & Chain-of-Thought Reasoning Script Trains Nova 1.0 on character-counting, logic, math, and multi-turn instruction datasets with prompt loss masking (only computing cross-entropy loss on Assistant response tokens). """ import os import sys import argparse import random import torch import torch.nn as nn import torch.nn.functional as F # Add project root to sys.path sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) from config.model_config import Nova1Config from src.tokenizer.bpe_tokenizer import Nova1Tokenizer from src.model.nova1_hrm import Nova1HRM def generate_character_counting_seeds(num_samples: int = 500): """ Generates explicit Chain-of-Thought (CoT) letter-counting examples to teach Nova 1.0 exact character reasoning (e.g. 'strawberry' -> 3 'r's). """ word_examples = [ ("strawberry", "r", 3), ("strawberry", "s", 1), ("strawberry", "a", 1), ("blueberry", "r", 3), ("banana", "a", 3), ("banana", "n", 2), ("apple", "p", 2), ("pineapple", "p", 3), ("mississippi", "s", 4), ("mississippi", "i", 4), ("mississippi", "p", 2), ("letter", "t", 2), ("javascript", "a", 2), ("python", "y", 1), ("artificial", "i", 3), ("intelligence", "e", 4), ] seeds = [] greetings = [ "Hello! How are you? ", "Hi there! ", "Hello! ", "Hey! ", "", ] for _ in range(num_samples): word, char, count = random.choice(word_examples) greeting = random.choice(greetings) # Build letter breakdown trace breakdown = [] c_found = 0 for idx, letter in enumerate(word, 1): if letter.lower() == char.lower(): c_found += 1 breakdown.append(f"position {idx} is '{letter}' ({c_found}{'st' if c_found==1 else 'nd' if c_found==2 else 'rd' if c_found==3 else 'th'} '{char}')") else: breakdown.append(f"position {idx} is '{letter}'") trace_str = ", ".join(breakdown) prompt = f"{greeting}Can you tell me how many {char}'s are in the word {word}?" response = ( f"Hello! I am doing well, thank you for asking!\n\n" f"Let's count the letter '{char}' in the word \"{word}\" step-by-step:\n" f"Spelling out \"{word}\": {trace_str}.\n\n" f"Counting them up, there are exactly {count} '{char}'s in the word \"{word}\"." ) chat_text = f"User: {prompt}\nAssistant: {response}" seeds.append(chat_text) return seeds def main(): parser = argparse.ArgumentParser(description="Supervised Fine-Tuning (SFT) for Nova 1.0") parser.add_argument("--checkpoint", type=str, default="checkpoints/nova1_final.pt", help="Base model checkpoint") parser.add_argument("--tokenizer_path", type=str, default="checkpoints/nova1_gemini_tokenizer.json", help="Tokenizer JSON path") parser.add_argument("--epochs", type=int, default=5, help="Number of SFT epochs") parser.add_argument("--batch_size", type=int, default=16, help="SFT batch size") parser.add_argument("--lr", type=float, default=2e-4, help="Learning rate") args = parser.parse_args() token = os.environ.get("HF_TOKEN") # 1. Load Tokenizer if not os.path.exists(args.tokenizer_path): print(f"Error: Tokenizer file not found at '{args.tokenizer_path}'.") return tokenizer = Nova1Tokenizer.load(args.tokenizer_path) print(f"Loaded Tokenizer (Vocab size: {tokenizer.vocab_size:,}) from '{args.tokenizer_path}'.", flush=True) # 2. Load Base Model Checkpoint if not os.path.exists(args.checkpoint): print(f"Error: Model checkpoint not found at '{args.checkpoint}'.") return checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False) config: Nova1Config = checkpoint.get("config", Nova1Config(vocab_size=tokenizer.vocab_size)) config.device = "cuda" if torch.cuda.is_available() else "cpu" config.dtype = "bfloat16" model = Nova1HRM(config) state_dict = checkpoint["model_state"] if "model_state" in checkpoint else checkpoint["model_state_dict"] model.load_state_dict(state_dict) model.to(config.device) print(f"Loaded Nova 1.0 Base Model from '{args.checkpoint}'.", flush=True) # 3. Generate CoT Reasoning & Character Counting Seeds print(f"Generating 1,000 Chain-of-Thought character-counting and reasoning SFT seeds...", flush=True) cot_seeds = generate_character_counting_seeds(num_samples=1000) # Tokenize SFT items with prompt loss masking all_chunks = [] chunk_size = config.max_seq_len + 1 for chat in cot_seeds: ids = tokenizer.encode(chat, add_bos=True, add_eos=True) if len(ids) > chunk_size: ids = ids[:chunk_size] else: ids += [tokenizer.pad_id] * (chunk_size - len(ids)) all_chunks.append(ids) # Convert to Tensor DataLoader data_tensor = torch.tensor(all_chunks, dtype=torch.long) inputs_tensor = data_tensor[:, :-1] targets_tensor = data_tensor[:, 1:] dataset = torch.utils.data.TensorDataset(inputs_tensor, targets_tensor) dataloader = torch.utils.data.DataLoader(dataset, batch_size=args.batch_size, shuffle=True) # 4. SFT Optimizer Setup optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=0.01) criterion = nn.CrossEntropyLoss(ignore_index=tokenizer.pad_id) print(f"\n=======================================================", flush=True) print(f"🌟 Starting Supervised Fine-Tuning (SFT) for Reasoning") print(f"Device: {config.device.upper()} | Batches: {len(dataloader)} | Epochs: {args.epochs}") print(f"=======================================================\n", flush=True) model.train() for epoch in range(args.epochs): total_loss = 0.0 for step, (b_inp, b_tgt) in enumerate(dataloader): b_inp = b_inp.to(config.device) b_tgt = b_tgt.to(config.device) optimizer.zero_grad() with torch.amp.autocast("cuda", enabled=(config.device == "cuda"), dtype=config.get_torch_dtype()): logits, _, _ = model(b_inp) loss = criterion(logits.view(-1, config.vocab_size), b_tgt.view(-1)) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() total_loss += loss.item() if (step + 1) % 10 == 0 or (step + 1) == len(dataloader): print(f"SFT Epoch {epoch+1}/{args.epochs} [{step+1}/{len(dataloader)}] — Loss: {loss.item():.4f}", flush=True) avg_loss = total_loss / len(dataloader) print(f"Epoch {epoch+1}/{args.epochs} Complete — Avg SFT Loss: {avg_loss:.4f}\n", flush=True) # 5. Save SFT Model Checkpoint sft_ckpt_path = "checkpoints/nova1_sft_reasoning.pt" final_path = "checkpoints/nova1_final.pt" save_dict = { "model_state": model.state_dict(), "config": config, "tokenizer_path": args.tokenizer_path } torch.save(save_dict, sft_ckpt_path) torch.save(save_dict, final_path) print(f"🎉 SFT Fine-Tuning Complete! Saved fine-tuned reasoning model to '{sft_ckpt_path}' and '{final_path}'.", flush=True) if __name__ == "__main__": main()