ARK-65M / ark65m.py
ThingsAI's picture
Upload folder using huggingface_hub
2a31ec5 verified
Raw History Blame Contribute Delete
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
@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}")