Nova / scripts /train_sft.py
kings1's picture
Upload folder using huggingface_hub
23ea6bd verified
Raw History Blame Contribute Delete
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()