| |
| """ |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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.""" |
| |
| |
| import random |
| batch = [] |
| for _ in range(B): |
| |
| pat_len = random.randint(4, 16) |
| pattern = [random.randint(1, self._vocab_size - 1) for _ in range(pat_len)] |
| |
| 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 |
| |
| 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.""" |
| |
| 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) |
| |
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| @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 = [] |
| |
| |
| 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) |
| |
| |
| 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 |
| |
| |
| probs = F.softmax(logits, dim=-1) |
| log_probs = F.log_softmax(logits, dim=-1) |
| H = -(probs * log_probs).sum(dim=-1).mean().item() |
| |
| |
| 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} |
|
|
|
|
| |
| |
| |
|
|
| def train(args): |
| """Main training function.""" |
| |
| |
| configs = { |
| "small": COZET_SMALL, |
| "medium": COZET_MEDIUM, |
| "large": COZET_LARGE, |
| } |
| config = configs[args.size] |
| |
| |
| 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) |
| |
| |
| 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) |
| |
| |
| 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}") |
| |
| |
| 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_dir = Path(args.output_dir) |
| output_dir.mkdir(parents=True, exist_ok=True) |
| |
| |
| 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: |
| |
| batch = dataset.get_batch(args.batch_size, device) |
| if dtype == torch.bfloat16: |
| |
| pass |
| |
| B = batch.shape[0] |
| chunk_len = batch.shape[1] - 1 |
| |
| |
| 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) |
| |
| |
| 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 |
| |
| |
| avg_loss = loss / max(n_tokens_batch, 1) |
| |
| |
| avg_loss.backward() |
| |
| |
| grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) |
| |
| |
| optimizer.step() |
| scheduler.step() |
| optimizer.zero_grad() |
| |
| |
| tokens_processed += n_tokens_batch |
| step += 1 |
| running_loss += avg_loss.item() |
| running_count += 1 |
| |
| |
| if step % log_interval == 0: |
| avg = running_loss / max(running_count, 1) |
| ppl = math.exp(min(avg, 20)) |
| 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 |
| |
| |
| 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_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") |
| |
| |
| 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}") |
| |
| |
| 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}") |
| |
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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}") |
|
|
|
|
| |
| |
| |
|
|
| if __name__ == "__main__": |
| parser = argparse.ArgumentParser(description="Cozet: Native SYNAXIM Pretraining") |
| |
| |
| parser.add_argument("--size", choices=["small", "medium", "large"], |
| default="small", help="Model size preset") |
| |
| |
| 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") |
| |
| |
| 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") |
| |
| |
| 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") |
| |
| |
| 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) |
|
|