Download scripts/train_sft.py from kings1/Nova: direct link, hf CLI and curl.
- Browser
- Download file 7.53 kB
-
https://huggingface.co/kings1/Nova/resolve/main/scripts/train_sft.py
- Command line
-
hf download hf://kings1/Nova/scripts/train_sft.py
-
curl -L -o train_sft.py https://huggingface.co/kings1/Nova/resolve/main/scripts/train_sft.py
7.53 kB
| """ | |
| 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() | |