Saturday-AI / scripts /train.py
Tejas123we's picture
Deploy Saturday-1.1B AI with Iridescent Glassmorphic GUI
8bf1a8c
Raw History Blame Contribute Delete
7.62 kB
#!/usr/bin/env python3
"""
Production Training CLI for Saturday LLM on Custom Dataset Files.
Train Saturday on real datasets like C:\\Users\\ojastejas\\anthropic_data.txt (~157 MB).
Usage:
python scripts/train.py --data_path C:\\Users\\ojastejas\\anthropic_data.txt --config configs/saturday_100m.yaml --steps 2000
"""
import sys
import os
import argparse
import time
import re
import numpy as np
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from rich.console import Console
from rich.panel import Panel
from rich.table import Table
from rich.progress import Progress, SpinnerColumn, TextColumn, BarColumn, TimeRemainingColumn
from saturday_numpy.config import SaturdayConfig
from saturday_numpy.tokenizer.word_tokenizer import WordTokenizer
from saturday_numpy.model.saturday import SaturdayModel
from saturday_numpy.training.loss import cross_entropy_loss
from saturday_numpy.training.optimizer import AdamW
from saturday_numpy.utils.checkpoint import save_checkpoint, load_checkpoint
console = Console()
def load_file_token_batches(file_path: str, tokenizer: WordTokenizer, batch_size: int, seq_len: int, max_tokens: int = 10_000_000):
"""Loads and tokenizes text from a dataset file into batch arrays."""
console.print(f"[bold white]Reading dataset from [yellow]{file_path}[/yellow]...[/bold white]")
with open(file_path, "r", encoding="utf-8", errors="ignore") as f:
text = f.read(max_tokens * 6) # Read initial chunk for fast loading
console.print(f" [dim]Text loaded ({len(text):,} characters). Tokenizing words...[/dim]")
token_ids = np.array(tokenizer.encode(text), dtype=np.int32)
console.print(f" [green][OK] Tokenized {len(token_ids):,} total tokens![/green]")
# Create sequence batches
batches = []
chunk_size = seq_len + 1
total_tokens_available = len(token_ids)
for i in range(0, total_tokens_available - chunk_size, seq_len):
seq = token_ids[i : i + chunk_size]
batches.append(seq)
if len(batches) >= batch_size * 2000:
break
batches = np.array(batches, dtype=np.int32)
return token_ids, batches
def main():
parser = argparse.ArgumentParser(description="Train Saturday LLM on Custom Dataset")
parser.add_argument("--data_path", type=str, default=r"C:\Users\ojastejas\anthropic_data.txt", help="Path to text dataset file")
parser.add_argument("--config", type=str, default="configs/saturday_100m.yaml", help="Path to YAML config file")
parser.add_argument("--steps", type=int, default=1000, help="Total training steps")
parser.add_argument("--lr", type=float, default=3e-4, help="Learning rate")
parser.add_argument("--checkpoint_dir", type=str, default="checkpoints", help="Directory to save checkpoints")
args = parser.parse_args()
console.clear()
console.print("=" * 60)
console.print(" Saturday LLM Dataset Training Pipeline")
console.print("=" * 60)
if not os.path.exists(args.data_path):
console.print(f"[bold red]Error:[/bold red] Dataset file not found at {args.data_path}")
return
# 1. Load Config
config = SaturdayConfig.from_yaml(args.config)
# 2. Build Vocabulary & Load Data
with open(args.data_path, "r", encoding="utf-8", errors="ignore") as f:
sample_text = f.read(500_000)
tokenizer = WordTokenizer.build_from_text(sample_text)
# Update config vocab_size to match dataset tokenizer
config.vocab_size = tokenizer.vocab_size
breakdown = config.count_parameters()
token_ids, batches = load_file_token_batches(
file_path=args.data_path,
tokenizer=tokenizer,
batch_size=config.batch_size,
seq_len=min(128, config.max_sequence_length)
)
console.print(f"\n[bold white]Dataset Info:[/bold white] [yellow]{args.data_path}[/yellow]")
console.print(f" Vocabulary Size: [bold cyan]{tokenizer.vocab_size:,} words[/bold cyan]")
console.print(f" Model Parameters: [bold green]{breakdown['total']:,}[/bold green]")
console.print(f" Available Sequence Batches: [cyan]{len(batches):,}[/cyan]\n")
# 3. Instantiate Model
console.print("[bold white]Initializing Model & Optimizer...[/bold white]")
model = SaturdayModel(config)
optimizer = AdamW(model=model, learning_rate=args.lr, weight_decay=config.weight_decay)
# 4. Training Loop
console.print(f"\n[bold bright_green]Starting Training Loop on Anthropic Dataset...[/bold bright_green]\n")
seq_len = min(128, config.max_sequence_length)
batch_size = config.batch_size
num_batches_available = len(batches)
start_training_time = time.time()
tokens_processed = 0
loss = 0.0
with Progress(
SpinnerColumn("dots", style="bright_red"),
TextColumn("[progress.description]{task.description}"),
BarColumn(bar_width=25, style="dim white", complete_style="bright_red"),
TextColumn("[bold yellow]{task.fields[loss]}[/bold yellow]"),
TextColumn("[cyan]{task.fields[tok_sec]}[/cyan]"),
TimeRemainingColumn(),
console=console,
) as progress:
task = progress.add_task(f"[bright_red]Training on anthropic_data.txt...", total=args.steps, loss="Loss: --", tok_sec="0 tok/s")
step_start_time = time.time()
for step in range(1, args.steps + 1):
batch_idx = (step * batch_size) % (num_batches_available - batch_size)
batch = batches[batch_idx : batch_idx + batch_size]
inputs = batch[:, :-1]
targets = batch[:, 1:]
logits = model.forward(inputs)
loss, d_logits = cross_entropy_loss(logits, targets)
model.backward(d_logits)
optimizer.step()
tokens_processed += inputs.size
if step % 10 == 0 or step == args.steps:
elapsed_step = time.time() - step_start_time
tok_sec = (10 * inputs.size) / max(1e-5, elapsed_step)
step_start_time = time.time()
progress.update(
task,
advance=10 if step > 10 else step,
loss=f"Loss: {loss:.4f}",
tok_sec=f"{tok_sec:,.0f} tok/s"
)
if step % 200 == 0 or step == args.steps:
ckpt_file = os.path.join(args.checkpoint_dir, f"saturday_anthropic_step_{step}.pkl")
save_checkpoint(
model=model,
optimizer=optimizer,
config=config,
tokenizer=tokenizer,
step=step,
train_tokens=tokens_processed,
val_loss=float(loss),
path=ckpt_file,
)
total_time = time.time() - start_training_time
avg_tok_sec = tokens_processed / total_time
console.print("\n" + "=" * 60)
console.print(f"[bold bright_green][OK] Anthropic Dataset Training Completed in {total_time:.2f}s![/bold bright_green]")
console.print(f" Total Processed Tokens: [bold yellow]{tokens_processed:,}[/bold yellow]")
console.print(f" Average Throughput: [cyan]{avg_tok_sec:,.1f} tokens/second[/cyan]")
console.print(f" Final Loss: [bold green]{loss:.4f}[/bold green] (Perplexity: {np.exp(loss):.2f})")
console.print(f" Saved Checkpoint: [yellow]checkpoints/saturday_anthropic_step_{args.steps}.pkl[/yellow]")
console.print("=" * 60 + "\n")
if __name__ == "__main__":
main()