#!/usr/bin/env python3 """ Cozet Training Script: Native SYNAXIM Pretraining =================================================== Trains the CozModel from scratch on streaming text data. Uses truncated BPTT through the M-matrix chain. Monitors C/H/R stabilization metrics from Odyssey. Usage: Local (CPU, quick test): python3 train.py --size small --max-tokens 100000 --device cpu GPU (single GPU): python3 train.py --size small --max-tokens 1000000000 --device cuda GH200 (full training): python3 train.py --size medium --max-tokens 50000000000 --device cuda --bf16 (c) 2026 GRRN Research. All rights reserved. """ import argparse import math import os import time import json import torch import torch.nn.functional as F from pathlib import Path from model import CozModel, CozConfig, COZET_SMALL, COZET_MEDIUM, COZET_LARGE # ====================================================================== # Data Pipeline # ====================================================================== class StreamingTextDataset: """ Streams tokenized text data for pretraining. Supports: - Local .bin files (pre-tokenized, uint16/uint32) - HuggingFace datasets (streamed, tokenized on-the-fly) - Synthetic data (for testing) """ def __init__(self, source: str, seq_len: int, tokenizer_name: str = "gpt2"): self.seq_len = seq_len self.source = source self._buffer = [] self._buffer_pos = 0 if source == "synthetic": self._mode = "synthetic" self._vocab_size = 32000 elif source.endswith(".bin"): self._mode = "bin" import numpy as np self._data = np.memmap(source, dtype=np.uint16, mode='r') self._pos = 0 else: self._mode = "hf" self._init_hf(source, tokenizer_name) def _init_hf(self, dataset_name: str, tokenizer_name: str): """Initialize HuggingFace streaming dataset.""" from datasets import load_dataset from transformers import AutoTokenizer print(f"[Data] Loading tokenizer: {tokenizer_name}") self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name) if self.tokenizer.pad_token is None: self.tokenizer.pad_token = self.tokenizer.eos_token print(f"[Data] Streaming dataset: {dataset_name}") self.dataset = load_dataset( dataset_name, split="train", streaming=True ) self._iter = iter(self.dataset) self._vocab_size = self.tokenizer.vocab_size def get_batch(self, batch_size: int, device: torch.device) -> torch.Tensor: """ Get a batch of token sequences. Returns: (batch_size, seq_len + 1) tensor of token IDs The +1 is for the target token at each position. """ if self._mode == "synthetic": return self._synthetic_batch(batch_size, device) elif self._mode == "bin": return self._bin_batch(batch_size, device) else: return self._hf_batch(batch_size, device) def _synthetic_batch(self, B: int, device: torch.device) -> torch.Tensor: """Generate synthetic data for architecture validation.""" # Patterns that test M-matrix retention: # Repeating sequences that the model should learn to predict import random batch = [] for _ in range(B): # Random repeating pattern of length 4-16 pat_len = random.randint(4, 16) pattern = [random.randint(1, self._vocab_size - 1) for _ in range(pat_len)] # Repeat to fill seq_len + 1 repeats = (self.seq_len + 1 + pat_len) // pat_len + 1 seq = (pattern * repeats)[:self.seq_len + 1] batch.append(seq) return torch.tensor(batch, dtype=torch.long, device=device) def _bin_batch(self, B: int, device: torch.device) -> torch.Tensor: """Read from pre-tokenized binary file.""" import numpy as np total_needed = B * (self.seq_len + 1) if self._pos + total_needed > len(self._data): self._pos = 0 # Wrap around chunk = self._data[self._pos:self._pos + total_needed].astype(np.int64) self._pos += total_needed return torch.tensor(chunk, dtype=torch.long, device=device).view(B, self.seq_len + 1) def _hf_batch(self, B: int, device: torch.device) -> torch.Tensor: """Tokenize from HuggingFace streaming dataset.""" # Fill buffer until we have enough tokens while len(self._buffer) < B * (self.seq_len + 1): try: example = next(self._iter) except StopIteration: self._iter = iter(self.dataset) example = next(self._iter) text = example.get("text", example.get("content", "")) if len(text) < 10: continue tokens = self.tokenizer.encode(text, add_special_tokens=False) self._buffer.extend(tokens) # Extract batch from buffer total = B * (self.seq_len + 1) batch_flat = self._buffer[:total] self._buffer = self._buffer[total:] return torch.tensor(batch_flat, dtype=torch.long, device=device).view(B, self.seq_len + 1) @property def vocab_size(self): return self._vocab_size # ====================================================================== # C/H/R Stabilization Metrics (from Odyssey) # ====================================================================== @torch.no_grad() def compute_chr_metrics(model: CozModel, eval_batch: torch.Tensor, device: torch.device) -> dict: """ Compute Consensus Coherence (C), Uncertainty Entropy (H), and Residual Contradiction (R) on an evaluation batch. These are the Odyssey stabilization metrics applied to the native SYNAXIM model during pretraining. """ model.eval() B, seq_len_plus1 = eval_batch.shape seq_len = seq_len_plus1 - 1 all_logits = [] all_hidden = [] # Process a subset for efficiency n_eval = min(B, 4) n_tokens = min(seq_len, 64) for b in range(n_eval): M_states = model.init_m_states(device) for t in range(n_tokens): tid = eval_batch[b, t].item() h = model.embed_tokens.weight[tid] for i, layer in enumerate(model.layers): h, M_states[i] = layer(h, M_states[i], t) h = model.final_norm(h) if model.config.tie_word_embeddings: logits = h @ model.embed_tokens.weight.T else: logits = model.lm_head(h) all_logits.append(logits) all_hidden.append(h) if not all_logits: model.train() return {"C": 0.0, "H": 0.0, "R": 0.0} logits = torch.stack(all_logits) hidden = torch.stack(all_hidden) # C: Consensus Coherence (pairwise cosine similarity) h_norm = F.normalize(hidden, dim=-1) sim = h_norm @ h_norm.T n = hidden.shape[0] if n > 1: mask = ~torch.eye(n, dtype=torch.bool, device=device) C = ((sim[mask].mean().item() + 1.0) / 2.0) else: C = 1.0 # H: Uncertainty Entropy probs = F.softmax(logits, dim=-1) log_probs = F.log_softmax(logits, dim=-1) H = -(probs * log_probs).sum(dim=-1).mean().item() # R: Residual Contradiction top1_probs = probs.max(dim=-1).values if top1_probs.mean() > 1e-8: R = (top1_probs.std() / (top1_probs.mean() + 1e-8)).item() else: R = 0.0 model.train() return {"C": C, "H": H, "R": R} # ====================================================================== # Training Loop # ====================================================================== def train(args): """Main training function.""" # ---- Config ---- configs = { "small": COZET_SMALL, "medium": COZET_MEDIUM, "large": COZET_LARGE, } config = configs[args.size] # Override vocab size if using specific tokenizer if args.dataset != "synthetic": from transformers import AutoTokenizer tok = AutoTokenizer.from_pretrained(args.tokenizer) config.vocab_size = tok.vocab_size del tok device = torch.device(args.device) dtype = torch.bfloat16 if args.bf16 and device.type == "cuda" else torch.float32 print("=" * 60) print(" COZET -- Native SYNAXIM Pretraining") print("=" * 60) print(f" Size: {args.size}") print(f" Device: {device}") print(f" Dtype: {dtype}") print(f" Max tokens: {args.max_tokens:,}") print(f" Batch size: {args.batch_size}") print(f" Chunk size: {args.chunk_size}") print(f" LR: {args.lr}") print(f" Dataset: {args.dataset}") print("=" * 60) # ---- Model ---- print("\n[1/4] Creating model...") model = CozModel(config) n_params = sum(p.numel() for p in model.parameters()) print(f" Parameters: {n_params:,}") if dtype == torch.bfloat16: model = model.to(dtype=dtype) model = model.to(device) # ---- Data ---- print("\n[2/4] Setting up data pipeline...") dataset = StreamingTextDataset( source=args.dataset, seq_len=args.chunk_size, tokenizer_name=args.tokenizer, ) print(f" Source: {args.dataset}") print(f" Seq length: {args.chunk_size}") # ---- Optimizer ---- print("\n[3/4] Configuring optimizer...") optimizer = torch.optim.AdamW( model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=0.1, eps=1e-8, ) total_steps = args.max_tokens // (args.batch_size * args.chunk_size) warmup_steps = min(2000, total_steps // 10) def lr_schedule(step): if step < warmup_steps: return step / max(warmup_steps, 1) progress = (step - warmup_steps) / max(total_steps - warmup_steps, 1) return 0.1 + 0.9 * 0.5 * (1.0 + math.cos(math.pi * min(progress, 1.0))) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_schedule) print(f" Total steps: {total_steps:,}") print(f" Warmup: {warmup_steps:,} steps") # ---- Output directory ---- output_dir = Path(args.output_dir) output_dir.mkdir(parents=True, exist_ok=True) # ---- Training ---- print("\n[4/4] Training...") model.train() tokens_processed = 0 step = 0 best_loss = float("inf") log_interval = args.log_every eval_interval = args.eval_every save_interval = args.save_every t_start = time.time() running_loss = 0.0 running_count = 0 while tokens_processed < args.max_tokens: # Get batch batch = dataset.get_batch(args.batch_size, device) if dtype == torch.bfloat16: # Token IDs stay int, but model runs in bf16 pass B = batch.shape[0] chunk_len = batch.shape[1] - 1 # -1 for targets # Forward with truncated BPTT loss = torch.tensor(0.0, device=device, dtype=dtype) n_tokens_batch = 0 for b in range(B): seq = batch[b] M_states = model.init_m_states(device) # Detach M at sequence start (no cross-sequence gradients) for t in range(chunk_len): logits, M_states = model.forward_token( seq[t].item(), M_states, t ) token_loss = F.cross_entropy( logits.unsqueeze(0).float(), seq[t + 1].unsqueeze(0), ) loss = loss + token_loss n_tokens_batch += 1 # Average loss avg_loss = loss / max(n_tokens_batch, 1) # Backward avg_loss.backward() # Gradient clipping grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # Step optimizer.step() scheduler.step() optimizer.zero_grad() # Track tokens_processed += n_tokens_batch step += 1 running_loss += avg_loss.item() running_count += 1 # ---- Logging ---- if step % log_interval == 0: avg = running_loss / max(running_count, 1) ppl = math.exp(min(avg, 20)) # Cap to prevent overflow elapsed = time.time() - t_start tok_per_sec = tokens_processed / max(elapsed, 1) lr = optimizer.param_groups[0]["lr"] print(f" Step {step:>6d} | " f"Loss {avg:.4f} | PPL {ppl:.2f} | " f"Grad {grad_norm:.2f} | " f"LR {lr:.2e} | " f"Tok/s {tok_per_sec:.0f} | " f"Tokens {tokens_processed:,}") running_loss = 0.0 running_count = 0 # ---- Evaluation (C/H/R) ---- if step % eval_interval == 0: eval_batch = dataset.get_batch(4, device) metrics = compute_chr_metrics(model, eval_batch, device) print(f" [EVAL] C={metrics['C']:.4f} | " f"H={metrics['H']:.4f} | R={metrics['R']:.4f}") # Log to file log_entry = { "step": step, "tokens": tokens_processed, "loss": avg_loss.item(), "C": metrics["C"], "H": metrics["H"], "R": metrics["R"], "lr": optimizer.param_groups[0]["lr"], } with open(output_dir / "training_log.jsonl", "a") as f: f.write(json.dumps(log_entry) + "\n") # ---- Save checkpoint ---- if step % save_interval == 0: ckpt_path = output_dir / f"checkpoint-{step}.pt" torch.save({ "step": step, "tokens_processed": tokens_processed, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "config": vars(config), "loss": avg_loss.item(), }, ckpt_path) print(f" [SAVE] {ckpt_path} ({ckpt_path.stat().st_size / 1e6:.1f} MB)") if avg_loss.item() < best_loss: best_loss = avg_loss.item() best_path = output_dir / "best_model.pt" torch.save({ "step": step, "tokens_processed": tokens_processed, "model_state_dict": model.state_dict(), "config": vars(config), "loss": best_loss, }, best_path) print(f" [BEST] New best loss: {best_loss:.4f}") # ---- Final save ---- elapsed = time.time() - t_start print(f"\n{'=' * 60}") print(f" Training complete!") print(f" Total tokens: {tokens_processed:,}") print(f" Total steps: {step:,}") print(f" Final loss: {avg_loss.item():.4f}") print(f" Best loss: {best_loss:.4f}") print(f" Time: {elapsed/3600:.2f} hours") print(f" Avg tok/s: {tokens_processed/max(elapsed,1):.0f}") print(f"{'=' * 60}") # Save final model final_path = output_dir / "final_model.pt" torch.save({ "step": step, "tokens_processed": tokens_processed, "model_state_dict": model.state_dict(), "config": vars(config), "loss": avg_loss.item(), }, final_path) print(f" Final model saved: {final_path}") return model # ====================================================================== # Push to HuggingFace # ====================================================================== def push_to_hf(output_dir: str, repo_id: str = "GRRNMAKER/Cozet"): """Push trained checkpoint and logs to HuggingFace.""" from huggingface_hub import HfApi token = os.environ.get("HF_TOKEN") if not token: print("[WARN] HF_TOKEN not set. Skipping push.") return api = HfApi(token=token) output_path = Path(output_dir) files_to_push = [] for f in output_path.iterdir(): if f.suffix in (".pt", ".jsonl", ".json", ".md"): files_to_push.append(f) for f in sorted(files_to_push): size_mb = f.stat().st_size / 1e6 print(f" Uploading {f.name} ({size_mb:.1f} MB)...") api.upload_file( path_or_fileobj=str(f), path_in_repo=f"checkpoints/{f.name}", repo_id=repo_id, repo_type="model", ) print(f" Pushed {len(files_to_push)} files to {repo_id}") # ====================================================================== # CLI # ====================================================================== if __name__ == "__main__": parser = argparse.ArgumentParser(description="Cozet: Native SYNAXIM Pretraining") # Model parser.add_argument("--size", choices=["small", "medium", "large"], default="small", help="Model size preset") # Data parser.add_argument("--dataset", default="synthetic", help="Dataset: 'synthetic', path to .bin, or HF dataset name") parser.add_argument("--tokenizer", default="gpt2", help="Tokenizer for HF datasets") # Training parser.add_argument("--max-tokens", type=int, default=1_000_000, help="Total tokens to train on") parser.add_argument("--batch-size", type=int, default=2, help="Sequences per batch") parser.add_argument("--chunk-size", type=int, default=128, help="Tokens per sequence (truncated BPTT window)") parser.add_argument("--lr", type=float, default=3e-4, help="Peak learning rate") parser.add_argument("--bf16", action="store_true", help="Use bfloat16 training") parser.add_argument("--device", default="cpu", help="Device: cpu or cuda") # Output parser.add_argument("--output-dir", default="./cozet-checkpoints", help="Directory for checkpoints and logs") parser.add_argument("--log-every", type=int, default=10, help="Log every N steps") parser.add_argument("--eval-every", type=int, default=50, help="Evaluate C/H/R every N steps") parser.add_argument("--save-every", type=int, default=500, help="Save checkpoint every N steps") # HuggingFace parser.add_argument("--push", action="store_true", help="Push checkpoints to HuggingFace after training") parser.add_argument("--hf-repo", default="GRRNMAKER/Cozet", help="HuggingFace repo ID") args = parser.parse_args() model = train(args) if args.push: push_to_hf(args.output_dir, args.hf_repo)