Download ark65m.py from ThingAI/ARK-65M: direct link, hf CLI and curl.
- Browser
- Download file 21.1 kB
-
https://huggingface.co/ThingAI/ARK-65M/resolve/main/ark65m.py
- Command line
-
hf download hf://ThingAI/ARK-65M/ark65m.py
-
curl -L -o ark65m.py https://huggingface.co/ThingAI/ARK-65M/resolve/main/ark65m.py
21.1 kB
| #!/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 | |
| 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 | |
| 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}") | |