#!/usr/bin/env python3 """ ARK-65M — ModotAI Usage: python ark65m.py train \ --data-config data_config.json \ --bin-dir pretokenized \ --tokenizer ThingAI/msqark-tokenizer \ --output-dir checkpoints-ark65m \ --batch-size 4 --grad-accum 16 --lr 3e-4 \ --max-steps 49591 --save-every 500 --compile """ import os, sys, json, time, argparse import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import Dataset, DataLoader from torch.utils.checkpoint import checkpoint from transformers import AutoTokenizer from dataclasses import dataclass @dataclass class ARKConfig: vocab_size: int = 32768 d_model: int = 576 n_heads: int = 8 n_kv_heads: int = 2 n_layers: int = 14 d_ff: int = 1536 max_seq_len: int = 2048 rope_theta: float = 500000.0 rms_eps: float = 1e-5 dropout: float = 0.0 router_start_layer: int = 7 gradient_checkpointing: bool = True @property def head_dim(self): return self.d_model // self.n_heads class RMSNorm(nn.Module): def __init__(self, dim, eps=1e-5): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, x): norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt() return (x.float() * norm).type_as(x) * self.weight class RotaryEmbedding(nn.Module): def __init__(self, dim, theta=500000.0, max_seq_len=2048): super().__init__() inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer("inv_freq", inv_freq, persistent=False) t = torch.arange(max_seq_len, dtype=self.inv_freq.dtype, device=self.inv_freq.device) freqs = torch.outer(t, self.inv_freq) emb = torch.cat([freqs, freqs], dim=-1) self.register_buffer("cos_cache", emb.cos(), persistent=False) self.register_buffer("sin_cache", emb.sin(), persistent=False) def forward(self, positions): positions = positions.clamp(0, self.cos_cache.shape[0] - 1) return self.cos_cache[positions], self.sin_cache[positions] def rotate_half(x): x1, x2 = x.chunk(2, dim=-1) return torch.cat((-x2, x1), dim=-1) def apply_rotary_pos_emb(q, k, cos, sin): cos, sin = cos.unsqueeze(1), sin.unsqueeze(1) return (q * cos) + (rotate_half(q) * sin), (k * cos) + (rotate_half(k) * sin) class SwiGLU(nn.Module): def __init__(self, d_model, d_ff): super().__init__() self.gate_proj = nn.Linear(d_model, d_ff, bias=False) self.up_proj = nn.Linear(d_model, d_ff, bias=False) self.down_proj = nn.Linear(d_ff, d_model, bias=False) def forward(self, x): return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) class GQAttention(nn.Module): def __init__(self, config): super().__init__() self.n_heads = config.n_heads self.n_kv_heads = config.n_kv_heads self.head_dim = config.head_dim self.n_rep = self.n_heads // self.n_kv_heads self.q_proj = nn.Linear(config.d_model, config.n_heads * self.head_dim, bias=False) self.k_proj = nn.Linear(config.d_model, config.n_kv_heads * self.head_dim, bias=False) self.v_proj = nn.Linear(config.d_model, config.n_kv_heads * self.head_dim, bias=False) self.o_proj = nn.Linear(config.n_heads * self.head_dim, config.d_model, bias=False) self.q_norm = RMSNorm(self.head_dim, eps=config.rms_eps) self.k_norm = RMSNorm(self.head_dim, eps=config.rms_eps) def forward(self, x, cos, sin, mask=None): B, L, _ = x.shape q = self.q_proj(x).view(B, L, self.n_heads, self.head_dim).transpose(1, 2) k = self.k_proj(x).view(B, L, self.n_kv_heads, self.head_dim).transpose(1, 2) v = self.v_proj(x).view(B, L, self.n_kv_heads, self.head_dim).transpose(1, 2) q, k = self.q_norm(q), self.k_norm(k) q, k = apply_rotary_pos_emb(q, k, cos, sin) k = k.repeat_interleave(self.n_rep, dim=1) v = v.repeat_interleave(self.n_rep, dim=1) out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, is_causal=(mask is None)) return self.o_proj(out.transpose(1, 2).contiguous().view(B, L, -1)) class MSAAttention(nn.Module): def __init__(self, config, layer_idx): super().__init__() self.n_heads = config.n_heads self.n_kv_heads = config.n_kv_heads self.head_dim = config.head_dim self.n_rep = self.n_heads // self.n_kv_heads self.window_size = 512 self.n_global = 64 self.q_proj = nn.Linear(config.d_model, config.n_heads * self.head_dim, bias=False) self.k_proj = nn.Linear(config.d_model, config.n_kv_heads * self.head_dim, bias=False) self.v_proj = nn.Linear(config.d_model, config.n_kv_heads * self.head_dim, bias=False) self.o_proj = nn.Linear(config.n_heads * self.head_dim, config.d_model, bias=False) self.q_norm = RMSNorm(self.head_dim, eps=config.rms_eps) self.k_norm = RMSNorm(self.head_dim, eps=config.rms_eps) self.router_q = nn.Linear(config.d_model, self.head_dim, bias=False) self.router_k = nn.Linear(config.d_model, self.head_dim, bias=False) def forward(self, x, cos, sin, mask=None): B, L, _ = x.shape q = self.q_proj(x).view(B, L, self.n_heads, self.head_dim).transpose(1, 2) k = self.k_proj(x).view(B, L, self.n_kv_heads, self.head_dim).transpose(1, 2) v = self.v_proj(x).view(B, L, self.n_kv_heads, self.head_dim).transpose(1, 2) q, k = self.q_norm(q), self.k_norm(k) q, k = apply_rotary_pos_emb(q, k, cos, sin) k = k.repeat_interleave(self.n_rep, dim=1) v = v.repeat_interleave(self.n_rep, dim=1) if L <= self.window_size + self.n_global: out = F.scaled_dot_product_attention(q, k, v, is_causal=True) else: out = self._sparse_attention(q, k, v, L) return self.o_proj(out.transpose(1, 2).contiguous().view(B, L, -1)) def _sparse_attention(self, q, k, v, L): output = torch.zeros_like(q) for i in range(0, L, self.window_size): end = min(i + self.window_size, L) q_chunk = q[:, :, i:end] local_start = max(0, i - self.window_size) k_local = k[:, :, local_start:end] v_local = v[:, :, local_start:end] q_pos = torch.arange(i, end, device=q.device) if i > self.n_global: k_combined = torch.cat([k[:, :, :self.n_global], k_local], dim=2) v_combined = torch.cat([v[:, :, :self.n_global], v_local], dim=2) k_pos = torch.cat([ torch.arange(self.n_global, device=q.device), torch.arange(local_start, end, device=q.device), ]) else: k_combined, v_combined = k_local, v_local k_pos = torch.arange(local_start, end, device=q.device) causal = q_pos.unsqueeze(1) >= k_pos.unsqueeze(0) attn_mask = torch.where(causal, 0.0, float('-inf')).unsqueeze(0).unsqueeze(0) output[:, :, i:end] = F.scaled_dot_product_attention( q_chunk, k_combined, v_combined, attn_mask=attn_mask ) return output class TransformerBlock(nn.Module): def __init__(self, config, layer_idx): super().__init__() self.attn = MSAAttention(config, layer_idx) if layer_idx >= config.router_start_layer else GQAttention(config) self.ffn = SwiGLU(config.d_model, config.d_ff) self.norm1 = RMSNorm(config.d_model, eps=config.rms_eps) self.norm2 = RMSNorm(config.d_model, eps=config.rms_eps) def forward(self, x, cos, sin, mask=None): x = x + self.attn(self.norm1(x), cos, sin, mask) x = x + self.ffn(self.norm2(x)) return x class ARK65M(nn.Module): def __init__(self, config: ARKConfig): super().__init__() self.config = config self.embed = nn.Embedding(config.vocab_size, config.d_model) self.rope = RotaryEmbedding(config.head_dim, config.rope_theta, config.max_seq_len) self.layers = nn.ModuleList([TransformerBlock(config, i) for i in range(config.n_layers)]) self.norm = RMSNorm(config.d_model, eps=config.rms_eps) self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False) self.lm_head.weight = self.embed.weight self.apply(self._init_weights) def _init_weights(self, module): if isinstance(module, nn.Linear): nn.init.normal_(module.weight, mean=0.0, std=0.02) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, input_ids, labels=None): B, L = input_ids.shape x = self.embed(input_ids) positions = torch.arange(L, device=input_ids.device).unsqueeze(0).expand(B, -1) cos, sin = self.rope(positions) for layer in self.layers: if self.config.gradient_checkpointing and self.training: x = checkpoint(layer, x, cos, sin, None, use_reentrant=False) else: x = layer(x, cos, sin) logits = self.lm_head(self.norm(x)) loss = None if labels is not None: loss = F.cross_entropy( logits[:, :-1].contiguous().view(-1, self.config.vocab_size), labels[:, 1:].contiguous().view(-1), ignore_index=-100, ) return {"logits": logits, "loss": loss} def count_parameters(self): total = sum(p.numel() for p in self.parameters()) embed = self.embed.weight.numel() return { "total": total, "embedding": embed, "embed_pct": f"{embed/total*100:.1f}%", "transformer": total - embed, } def _unwrap(model): return model._orig_mod if hasattr(model, '_orig_mod') else model def _safe_filename(name): return name.replace(' ', '_').replace('/', '_').replace('\\', '_') class PreTokenizedDataset(Dataset): def __init__(self, data_config, bin_dir, max_len=2048): super().__init__() self.max_len = max_len with open(data_config, 'r', encoding='utf-8') as f: sources_cfg = json.load(f) self.sources = [] for src in sources_cfg: safe_name = _safe_filename(src["name"]) bin_path = os.path.join(bin_dir, f"{safe_name}.bin") if not os.path.exists(bin_path): print(f"Warning: {bin_path} not found, skipping {src['name']}") continue file_bytes = os.path.getsize(bin_path) dtype = np.uint16 if file_bytes % 2 == 0 else np.uint32 memmap = np.memmap(bin_path, dtype=dtype, mode='r') if len(memmap) <= max_len: continue self.sources.append({ "name": src["name"], "weight": float(src["weight"]), "memmap": memmap, "length": len(memmap), }) if not self.sources: raise ValueError("No valid sources!") weights = np.array([s["weight"] for s in self.sources], dtype=np.float64) weights /= weights.sum() self.cum_weights = np.cumsum(weights) self._len = int(sum(s["length"] for s in self.sources) // max_len * 2) def __len__(self): return self._len def __getitem__(self, idx): rng = np.random.RandomState(idx) source = self.sources[int(np.searchsorted(self.cum_weights, rng.rand()))] start = rng.randint(0, source["length"] - self.max_len) tokens = source["memmap"][start: start + self.max_len] input_ids = torch.from_numpy(tokens.astype(np.int64)).long() return {"input_ids": input_ids, "labels": input_ids.clone()} def train(): parser = argparse.ArgumentParser() parser.add_argument("--data-config", required=True) parser.add_argument("--bin-dir", default="pretokenized") parser.add_argument("--tokenizer", default="ThingAI/msqark-tokenizer") parser.add_argument("--output-dir", default="checkpoints-ark65m") parser.add_argument("--batch-size", type=int, default=4) parser.add_argument("--grad-accum", type=int, default=16) parser.add_argument("--lr", type=float, default=3e-4) parser.add_argument("--max-len", type=int, default=2048) parser.add_argument("--warmup-steps", type=int, default=500) parser.add_argument("--max-steps", type=int, default=0) parser.add_argument("--save-every", type=int, default=500) parser.add_argument("--num-workers", type=int, default=4) parser.add_argument("--resume", type=str, default=None) parser.add_argument("--compile", action="store_true") parser.add_argument("--no-checkpoint", action="store_true") args = parser.parse_args() torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True torch.set_float32_matmul_precision('high') device = torch.device("cuda" if torch.cuda.is_available() else "cpu") os.makedirs(args.output_dir, exist_ok=True) tokenizer = AutoTokenizer.from_pretrained(args.tokenizer, trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token config = ARKConfig( vocab_size=max(tokenizer.vocab_size, 32768), max_seq_len=args.max_len, gradient_checkpointing=not args.no_checkpoint, ) model = ARK65M(config).to(device) if args.compile: try: print("Compiling with torch.compile...") model = torch.compile(model, mode="default") print("✅ Compiled.") except Exception as e: print(f"⚠️ torch.compile failed: {e}") params = _unwrap(model).count_parameters() print(f"\n{'═'*55}") print(f" ARK-65M — ModotAI") print(f"{'═'*55}") print(f" Parameters: {params['total']:,}") print(f" Embedding: {params['embed_pct']}") print(f" Transformer: {params['transformer']:,}") print(f" Context: {config.max_seq_len:,} tokens") print(f" Layers: {config.n_layers} ({config.router_start_layer} GQA + {config.n_layers - config.router_start_layer} MSA)") print(f" d_model: {config.d_model}") print(f" Batch eff.: {args.batch_size} x {args.grad_accum} = {args.batch_size * args.grad_accum}") print(f" LR: {args.lr}") print(f" Grad ckpt: {config.gradient_checkpointing}") print(f" torch.compile: {args.compile}") print(f" Device: {device}") if torch.cuda.is_available(): print(f" GPU: {torch.cuda.get_device_name()}") print(f" VRAM: {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB") print(f"{'═'*55}\n") dataset = PreTokenizedDataset(args.data_config, args.bin_dir, max_len=args.max_len) loader = DataLoader( dataset, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers, persistent_workers=args.num_workers > 0, pin_memory=True, drop_last=True, ) print(f"Dataset: {len(dataset):,} samples") print(f"Sources loaded: {len(dataset.sources)}") for s in dataset.sources: print(f" • {s['name']:.<35} {s['length']:>12,} tokens") print() optimizer = torch.optim.AdamW( model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=0.1, ) def lr_lambda(step): if step < args.warmup_steps: return step / max(args.warmup_steps, 1) return 1.0 scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda) scaler = torch.amp.GradScaler("cuda") step = 0 tokens_seen = 0 if args.resume: print(f"🔄 Resuming from {args.resume}") ckpt = torch.load(args.resume, map_location=device, weights_only=False) _unwrap(model).load_state_dict(ckpt["model_state_dict"]) optimizer.load_state_dict(ckpt["optimizer_state_dict"]) scheduler.load_state_dict(ckpt["scheduler_state_dict"]) scaler.load_state_dict(ckpt["scaler_state_dict"]) step = ckpt.get("step", 0) tokens_seen = ckpt.get("tokens_seen", 0) print(f" ✅ Resumed at step {step:,} | {tokens_seen/1e9:.3f}B tokens | LR {scheduler.get_last_lr()[0]:.2e}\n") model.train() loss_accum = 0.0 log_steps = 0 start_time = time.time() last_log = start_time tokens_inst = 0 done = False print("Starting training...\n") while not done: for batch_idx, batch in enumerate(loader): input_ids = batch["input_ids"].to(device, non_blocking=True) labels = batch["labels"].to(device, non_blocking=True) with torch.amp.autocast("cuda", dtype=torch.bfloat16): out = model(input_ids, labels=labels) loss = out["loss"] / args.grad_accum scaler.scale(loss).backward() loss_accum += loss.item() batch_tokens = input_ids.numel() tokens_seen += batch_tokens tokens_inst += batch_tokens if (batch_idx + 1) % args.grad_accum == 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad() scheduler.step() step += 1 log_steps += 1 if step % 10 == 0: now = time.time() elapsed_total = now - start_time elapsed_inst = now - last_log avg_loss = loss_accum / log_steps tok_avg = tokens_seen / elapsed_total if elapsed_total > 0 else 0 tok_inst = tokens_inst / elapsed_inst if elapsed_inst > 0 else 0 lr_now = scheduler.get_last_lr()[0] vram = torch.cuda.memory_allocated() / 1e9 print( f"step {step:>6} │ loss {avg_loss:.4f} │ lr {lr_now:.2e} │ " f"avg {tok_avg:>7,.0f} tok/s │ inst {tok_inst:>7,.0f} tok/s │ " f"{tokens_seen/1e9:.3f}B tok │ VRAM {vram:.1f}GB │ {elapsed_total:.0f}s" ) loss_accum = 0.0 log_steps = 0 tokens_inst = 0 last_log = now if step % args.save_every == 0: path = os.path.join(args.output_dir, f"step-{step}.pt") torch.save({ "model_state_dict": _unwrap(model).state_dict(), "optimizer_state_dict": optimizer.state_dict(), "scheduler_state_dict": scheduler.state_dict(), "scaler_state_dict": scaler.state_dict(), "config": config.__dict__, "step": step, "tokens_seen": tokens_seen, }, path) print(f" 💾 Saved: {path}") if args.max_steps > 0 and step >= args.max_steps: done = True break if done: break final = os.path.join(args.output_dir, "final.pt") torch.save({ "model_state_dict": _unwrap(model).state_dict(), "config": config.__dict__, "step": step, "tokens_seen": tokens_seen, }, final) elapsed = time.time() - start_time print(f"\n{'═'*55}") print(f" Done!") print(f" Steps: {step:,}") print(f" Tokens: {tokens_seen/1e9:.3f}B") print(f" Time: {elapsed/3600:.1f}h") print(f" Saved: {final}") print(f"{'═'*55}") if __name__ == "__main__": if len(sys.argv) > 1 and sys.argv[1] == "train": sys.argv.pop(1) train() else: config = ARKConfig(gradient_checkpointing=False) model = ARK65M(config) params = model.count_parameters() print(f"{'═'*45}") print(f" ARK-65M — ModotAI") print(f"{'═'*45}") for k, v in params.items(): if isinstance(v, int): print(f" {k:.<25} {v:>12,}") else: print(f" {k:.<25} {v:>12}") x = torch.randint(0, config.vocab_size, (2, 128)) out = model(x, labels=x) print(f"\n Test forward OK") print(f" Loss: {out['loss'].item():.4f}") print(f"{'═'*45}")