#!/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()