Cozet / train.py
GRRNMAKER's picture
Upload train.py with huggingface_hub
0ece874 verified
Raw
History Blame Contribute Delete
19.4 kB
#!/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)