""" Hebrew DINT Transformer — A from-scratch Hebrew LLM Architecture based on DINT Transformer (arxiv:2501.17486) which extends Differential Attention (arxiv:2410.05258) with integral global context. Key innovations: - Differential Attention: cancels attention noise via dual softmax subtraction - DINT integral term: adds global token importance awareness - Pre-RMSNorm, SwiGLU FFN, RoPE (LLaMA-style macro layout) - Per-head RMSNorm after differential attention for gradient stability """ import math from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F # ─── RMSNorm ──────────────────────────────────────────────────────────────── class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: norm = torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps) return (x.float() * norm).type_as(x) * self.weight # ─── Rotary Position Embeddings ───────────────────────────────────────────── class RotaryEmbedding(nn.Module): def __init__(self, dim: int, max_seq_len: int = 4096, theta: float = 10000.0): super().__init__() inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer("inv_freq", inv_freq, persistent=False) self.max_seq_len = max_seq_len def forward(self, seq_len: int, device: torch.device, dtype: torch.dtype): t = torch.arange(seq_len, device=device, dtype=torch.float32) freqs = torch.outer(t, self.inv_freq.to(device)) cos = freqs.cos().to(dtype) sin = freqs.sin().to(dtype) return cos, sin def apply_rotary_emb(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: """Apply rotary embeddings. x: [B, S, H, D], cos/sin: [S, D/2]""" rotary_dim = cos.shape[-1] * 2 x_rot = x[..., :rotary_dim] x_pass = x[..., rotary_dim:] x1, x2 = x_rot[..., ::2], x_rot[..., 1::2] cos = cos.unsqueeze(0).unsqueeze(2) # [1, S, 1, D/2] sin = sin.unsqueeze(0).unsqueeze(2) rot_x1 = x1 * cos - x2 * sin rot_x2 = x1 * sin + x2 * cos rot_x = torch.stack([rot_x1, rot_x2], dim=-1).reshape_as(x_rot) return torch.cat([rot_x, x_pass], dim=-1) # ─── DINT Attention ───────────────────────────────────────────────────────── def lambda_init_fn(depth: int) -> float: """Per-layer λ initialization schedule from DIFF Transformer paper.""" return 0.8 - 0.6 * math.exp(-0.3 * depth) class DINTAttention(nn.Module): """ DINT (Differential + INTegral) Attention. DIFF: A = softmax(Q1·K1^T) - λ·softmax(Q2·K2^T) (noise cancellation) DINT: adds global integral term G = mean(A_pos, dim=0) broadcast back Final: A_diff + γ·G (where γ = λ for row-normalization stability) Uses SDPA (Flash Attention compatible) for efficiency. """ def __init__( self, embed_dim: int, depth: int, num_heads: int, num_kv_heads: Optional[int] = None, ): super().__init__() self.embed_dim = embed_dim self.num_heads = num_heads self.num_kv_heads = num_kv_heads or num_heads self.n_rep = self.num_heads // self.num_kv_heads # DINT uses half the head dim (pair of heads for diff) self.head_dim = embed_dim // self.num_heads // 2 assert self.head_dim * self.num_heads * 2 == embed_dim # Projections self.q_proj = nn.Linear(embed_dim, embed_dim, bias=False) self.k_proj = nn.Linear(embed_dim, embed_dim // self.n_rep, bias=False) self.v_proj = nn.Linear(embed_dim, embed_dim // self.n_rep, bias=False) self.out_proj = nn.Linear(embed_dim, embed_dim, bias=False) # λ parameters (learnable, shared across heads in this layer) self.lambda_init = lambda_init_fn(depth) self.lambda_q1 = nn.Parameter(torch.randn(self.head_dim) * 0.1) self.lambda_k1 = nn.Parameter(torch.randn(self.head_dim) * 0.1) self.lambda_q2 = nn.Parameter(torch.randn(self.head_dim) * 0.1) self.lambda_k2 = nn.Parameter(torch.randn(self.head_dim) * 0.1) # Per-head LayerNorm (GroupNorm over heads) self.subln = nn.LayerNorm(2 * self.head_dim, eps=1e-5) def _repeat_kv(self, x: torch.Tensor) -> torch.Tensor: """Repeat K/V heads for GQA.""" bs, n_kv_heads, slen, head_dim = x.shape if self.n_rep == 1: return x return ( x[:, :, None, :, :] .expand(bs, n_kv_heads, self.n_rep, slen, head_dim) .reshape(bs, n_kv_heads * self.n_rep, slen, head_dim) ) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, ) -> torch.Tensor: bsz, seq_len, _ = x.size() # Project q = self.q_proj(x) k = self.k_proj(x) v = self.v_proj(x) # Reshape into paired heads q = q.view(bsz, seq_len, 2 * self.num_heads, self.head_dim) k = k.view(bsz, seq_len, 2 * self.num_kv_heads, self.head_dim) v = v.view(bsz, seq_len, self.num_kv_heads, 2 * self.head_dim) # Apply RoPE q = apply_rotary_emb(q, cos, sin) k = apply_rotary_emb(k, cos, sin) # Prepare for attention [B, H, S, D] q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) # GQA repeat k = self._repeat_kv(k) v = self._repeat_kv(v) # Split into positive/negative head pairs q_pairs = q.view(bsz, 2, self.num_heads, seq_len, self.head_dim).permute(0, 2, 1, 3, 4) k_pairs = k.view(bsz, 2, self.num_heads, seq_len, self.head_dim).permute(0, 2, 1, 3, 4) q_pos, q_neg = q_pairs[:, :, 0], q_pairs[:, :, 1] # [B, H, S, D] k_pos, k_neg = k_pairs[:, :, 0], k_pairs[:, :, 1] # Compute λ scalar lambda_1 = torch.exp(torch.sum(self.lambda_q1 * self.lambda_k1)).type_as(q_pos) lambda_2 = torch.exp(torch.sum(self.lambda_q2 * self.lambda_k2)).type_as(q_pos) lambda_full = lambda_1 - lambda_2 + self.lambda_init # ── DIFF: Two SDPA calls ── ctx_pos = F.scaled_dot_product_attention(q_pos, k_pos, v, is_causal=True) ctx_neg = F.scaled_dot_product_attention(q_neg, k_neg, v, is_causal=True) # Differential attention: noise cancellation attn_out = ctx_pos - lambda_full * ctx_neg # [B, H, S, 2D] # ── DINT: Add integral (global importance) term ── # Global importance = mean attention output across sequence positions # This gives each position access to a "summary" of global context global_ctx = ctx_pos.mean(dim=2, keepdim=True).expand_as(ctx_pos) # [B, H, S, 2D] attn_out = attn_out + lambda_full * global_ctx # Per-head LayerNorm + residual scaling attn_out = self.subln(attn_out) * (1.0 - self.lambda_init) # Reshape and project attn_out = attn_out.transpose(1, 2).reshape(bsz, seq_len, self.embed_dim) return self.out_proj(attn_out) # ─── SwiGLU Feed-Forward ──────────────────────────────────────────────────── class SwiGLUFFN(nn.Module): """SwiGLU feed-forward network (LLaMA-style).""" def __init__(self, embed_dim: int, ffn_dim: int): super().__init__() self.gate_proj = nn.Linear(embed_dim, ffn_dim, bias=False) self.up_proj = nn.Linear(embed_dim, ffn_dim, bias=False) self.down_proj = nn.Linear(ffn_dim, embed_dim, bias=False) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) # ─── DINT Transformer Block ──────────────────────────────────────────────── class DINTTransformerBlock(nn.Module): def __init__(self, embed_dim: int, num_heads: int, ffn_dim: int, depth: int): super().__init__() self.attn = DINTAttention(embed_dim, depth, num_heads) self.ffn = SwiGLUFFN(embed_dim, ffn_dim) self.attn_norm = RMSNorm(embed_dim) self.ffn_norm = RMSNorm(embed_dim) def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: # Pre-norm residual x = x + self.attn(self.attn_norm(x), cos, sin) x = x + self.ffn(self.ffn_norm(x)) return x # ─── Full DINT Transformer LLM ───────────────────────────────────────────── class DINTTransformerLM(nn.Module): """ DINT Transformer Language Model for Hebrew. Architecture: Pre-RMSNorm + DINT Attention + SwiGLU FFN + RoPE Based on: arxiv:2501.17486 (DINT) + arxiv:2410.05258 (DIFF Transformer) """ def __init__( self, vocab_size: int = 32000, embed_dim: int = 2048, num_layers: int = 24, num_heads: int = 16, ffn_dim: int = 5504, max_seq_len: int = 2048, tie_embeddings: bool = True, ): super().__init__() self.vocab_size = vocab_size self.embed_dim = embed_dim self.max_seq_len = max_seq_len # Token embeddings self.token_emb = nn.Embedding(vocab_size, embed_dim) # Rotary embeddings head_dim = embed_dim // num_heads // 2 # DINT half head dim self.rotary = RotaryEmbedding(head_dim, max_seq_len) # Transformer layers self.layers = nn.ModuleList([ DINTTransformerBlock(embed_dim, num_heads, ffn_dim, depth=i) for i in range(num_layers) ]) # Output self.norm = RMSNorm(embed_dim) self.lm_head = nn.Linear(embed_dim, vocab_size, bias=False) # Weight tying if tie_embeddings: self.lm_head.weight = self.token_emb.weight # Initialize weights self.apply(self._init_weights) def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward( self, input_ids: torch.Tensor, labels: Optional[torch.Tensor] = None, ) -> dict: bsz, seq_len = input_ids.shape assert seq_len <= self.max_seq_len, f"Sequence length {seq_len} exceeds max {self.max_seq_len}" # Embeddings x = self.token_emb(input_ids) # RoPE cos, sin = self.rotary(seq_len, x.device, x.dtype) # Transformer layers for layer in self.layers: x = layer(x, cos, sin) # Output x = self.norm(x) logits = self.lm_head(x) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss = F.cross_entropy( shift_logits.view(-1, self.vocab_size), shift_labels.view(-1), ignore_index=-100, ) return {"loss": loss, "logits": logits} def count_parameters(self) -> int: return sum(p.numel() for p in self.parameters() if p.requires_grad) # ─── Model configurations ────────────────────────────────────────────────── def create_hebrew_dint_1_5b(vocab_size: int = 32000) -> DINTTransformerLM: """~1.5B parameter DINT Transformer for Hebrew.""" return DINTTransformerLM( vocab_size=vocab_size, embed_dim=2048, num_layers=24, num_heads=16, ffn_dim=5504, max_seq_len=2048, ) def create_hebrew_dint_400m(vocab_size: int = 32000) -> DINTTransformerLM: """~400M parameter DINT Transformer for Hebrew (faster iteration).""" return DINTTransformerLM( vocab_size=vocab_size, embed_dim=1024, num_layers=20, num_heads=8, ffn_dim=2816, max_seq_len=2048, ) # ═══════════════════════════════════════════════════════════════════════════════ # TRAINING SCRIPT # ═══════════════════════════════════════════════════════════════════════════════ """ Hebrew DINT Transformer — Pretraining Script Trains a DINT Transformer (Differential + Integral Attention) from scratch on Hebrew text data (HeDC4 + OzLabs Wikipedia + Ben Yehuda + Military + Wiktionary). Architecture: arxiv:2501.17486 (DINT) + arxiv:2410.05258 (DIFF Transformer) Training recipe: Informed by DictaLM 2.0 (arxiv:2407.07080) hyperparameters """ import os import sys import math import json import time import argparse from pathlib import Path import torch import torch.nn.functional as F from torch.utils.data import DataLoader, Dataset, IterableDataset from torch.cuda.amp import GradScaler from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR # Import our model # model classes defined above (inline) # ─── Config ───────────────────────────────────────────────────────────────── def get_args(): parser = argparse.ArgumentParser(description="Train Hebrew DINT Transformer") # Model parser.add_argument("--model_size", type=str, default="400m", choices=["400m", "1.5b"]) parser.add_argument("--vocab_size", type=int, default=32000) parser.add_argument("--max_seq_len", type=int, default=2048) # Training parser.add_argument("--batch_size", type=int, default=8) parser.add_argument("--grad_accum", type=int, default=8) parser.add_argument("--lr", type=float, default=3e-4) parser.add_argument("--min_lr", type=float, default=3e-5) parser.add_argument("--weight_decay", type=float, default=0.1) parser.add_argument("--warmup_steps", type=int, default=1000) parser.add_argument("--max_steps", type=int, default=50000) parser.add_argument("--bf16", action="store_true", default=True) parser.add_argument("--gradient_checkpointing", action="store_true", default=True) # Data parser.add_argument("--tokenizer_path", type=str, default="./tokenizer") # Logging & Saving parser.add_argument("--output_dir", type=str, default="./hebrew-dint-transformer") parser.add_argument("--hub_model_id", type=str, default="guychuk/hebrew-dint-transformer") parser.add_argument("--log_every", type=int, default=10) parser.add_argument("--save_every", type=int, default=2000) parser.add_argument("--eval_every", type=int, default=500) return parser.parse_args() # ─── Hebrew Tokenizer Training ────────────────────────────────────────────── def train_hebrew_tokenizer(vocab_size: int, save_path: str): """Train a BPE tokenizer on Hebrew text data.""" from tokenizers import Tokenizer, models, trainers, pre_tokenizers, decoders, processors from datasets import load_dataset print("Loading Hebrew datasets for tokenizer training...") texts = [] # Load HeDC4 (stream a subset for tokenizer training) try: hedc4 = load_dataset("HeNLP/HeDC4", split="train", streaming=True) for i, row in enumerate(hedc4): if i >= 100000: break text = row.get("text", "") if text and len(text.strip()) > 10: texts.append(text) print(f" HeDC4: {len(texts)} docs") except Exception as e: print(f" HeDC4 error: {e}") # Load OzLabs Wikipedia try: wiki = load_dataset("OzLabs/hebrew-wiki-articles", split="train", streaming=True) wiki_count = 0 for row in wiki: if wiki_count >= 50000: break text = row.get("text", "") if text and len(text.strip()) > 10: texts.append(text) wiki_count += 1 print(f" Wiki: {wiki_count} docs") except Exception as e: print(f" Wiki error: {e}") # Load Ben Yehuda try: benyehuda = load_dataset("OzLabs/hebrew-project-benyehuda", split="train", streaming=True) by_count = 0 for row in benyehuda: text = row.get("text", "") if text and len(text.strip()) > 10: texts.append(text) by_count += 1 print(f" Ben Yehuda: {by_count} docs") except Exception as e: print(f" Ben Yehuda error: {e}") # Filter out None/empty texts texts = [t for t in texts if t and isinstance(t, str) and len(t.strip()) > 0] print(f"Total texts for tokenizer: {len(texts)}") # Train BPE tokenizer tokenizer = Tokenizer(models.BPE()) tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False) tokenizer.decoder = decoders.ByteLevel() tokenizer.post_processor = processors.ByteLevel(trim_offsets=False) bpe_trainer = trainers.BpeTrainer( vocab_size=vocab_size, special_tokens=["<|pad|>", "<|bos|>", "<|eos|>", "<|unk|>"], min_frequency=2, show_progress=True, ) tokenizer.train_from_iterator(iter(texts), bpe_trainer, length=len(texts)) # Save os.makedirs(save_path, exist_ok=True) tokenizer.save(os.path.join(save_path, "tokenizer.json")) # Save config config = { "vocab_size": tokenizer.get_vocab_size(), "pad_id": tokenizer.token_to_id("<|pad|>"), "bos_id": tokenizer.token_to_id("<|bos|>"), "eos_id": tokenizer.token_to_id("<|eos|>"), } with open(os.path.join(save_path, "config.json"), "w") as f: json.dump(config, f) print(f"Tokenizer saved to {save_path} (vocab size: {config['vocab_size']})") # Test encoding test = "שלום עולם, זוהי בדיקת הטוקנייזר העברי" encoded = tokenizer.encode(test) decoded = tokenizer.decode(encoded.ids) print(f"Test encode/decode: '{test}' -> {len(encoded.ids)} tokens -> '{decoded}'") return tokenizer, config def load_tokenizer(path: str): from tokenizers import Tokenizer tokenizer = Tokenizer.from_file(os.path.join(path, "tokenizer.json")) with open(os.path.join(path, "config.json")) as f: config = json.load(f) return tokenizer, config # ─── Dataset ──────────────────────────────────────────────────────────────── class HebrewPretrainDataset(IterableDataset): """ Streaming dataset that concatenates Hebrew text into chunks of max_seq_len. Packs documents together with EOS separators for maximum throughput. """ def __init__(self, tokenizer, tok_config, max_seq_len: int = 2048, seed: int = 42): self.tokenizer = tokenizer self.max_seq_len = max_seq_len self.eos_id = tok_config["eos_id"] self.pad_id = tok_config["pad_id"] self.seed = seed def _text_iterator(self): """Yields text from all Hebrew sources, streaming.""" from datasets import load_dataset # HeDC4 — largest source (~2.4GB) try: ds = load_dataset("HeNLP/HeDC4", split="train", streaming=True) for row in ds: text = row.get("text", "") if text and len(text) > 50: yield text except Exception as e: print(f"HeDC4 stream error: {e}") # OzLabs Wikipedia (~1.5GB) try: ds = load_dataset("OzLabs/hebrew-wiki-articles", split="train", streaming=True) for row in ds: text = row.get("text", "") if text and len(text) > 50: yield text except Exception as e: print(f"Wiki stream error: {e}") # OzLabs Ben Yehuda (~250MB) try: ds = load_dataset("OzLabs/hebrew-project-benyehuda", split="train", streaming=True) for row in ds: text = row.get("text", "") if text and len(text) > 50: yield text except Exception as e: print(f"BenYehuda stream error: {e}") # OzLabs Military (~93MB) try: ds = load_dataset("OzLabs/hebrew-military-documents", split="train", streaming=True) for row in ds: text = row.get("text", "") if text and len(text) > 50: yield text except Exception as e: print(f"Military stream error: {e}") # OzLabs Wiktionary (~10MB) try: ds = load_dataset("OzLabs/hebrew-wiktionary-articles", split="train", streaming=True) for row in ds: text = row.get("text", "") if text and len(text) > 20: yield text except Exception as e: print(f"Wiktionary stream error: {e}") def __iter__(self): """Pack documents into fixed-length sequences.""" buffer = [] for text in self._text_iterator(): # Tokenize encoded = self.tokenizer.encode(text) tokens = encoded.ids + [self.eos_id] buffer.extend(tokens) # Yield full sequences while len(buffer) >= self.max_seq_len + 1: chunk = buffer[:self.max_seq_len + 1] buffer = buffer[self.max_seq_len:] # Overlap by 1 for labels input_ids = torch.tensor(chunk[:-1], dtype=torch.long) labels = torch.tensor(chunk[1:], dtype=torch.long) yield {"input_ids": input_ids, "labels": labels} # ─── Learning Rate Schedule ───────────────────────────────────────────────── def get_lr(step: int, warmup_steps: int, max_steps: int, lr: float, min_lr: float) -> float: """Cosine schedule with linear warmup.""" if step < warmup_steps: return lr * (step + 1) / warmup_steps if step >= max_steps: return min_lr progress = (step - warmup_steps) / (max_steps - warmup_steps) return min_lr + 0.5 * (lr - min_lr) * (1 + math.cos(math.pi * progress)) # ─── Training Loop ────────────────────────────────────────────────────────── def train(args): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device}") if device.type == "cuda": print(f"GPU: {torch.cuda.get_device_name()}") print(f"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB") # ── Tokenizer ── tokenizer_path = args.tokenizer_path if not os.path.exists(os.path.join(tokenizer_path, "tokenizer.json")): print("Training Hebrew tokenizer...") tokenizer, tok_config = train_hebrew_tokenizer(args.vocab_size, tokenizer_path) else: print("Loading existing tokenizer...") tokenizer, tok_config = load_tokenizer(tokenizer_path) actual_vocab_size = tok_config["vocab_size"] print(f"Vocab size: {actual_vocab_size}") # ── Model ── print(f"Creating DINT Transformer ({args.model_size})...") if args.model_size == "1.5b": model = create_hebrew_dint_1_5b(vocab_size=actual_vocab_size) else: model = create_hebrew_dint_400m(vocab_size=actual_vocab_size) n_params = model.count_parameters() print(f"Model parameters: {n_params:,} ({n_params/1e9:.2f}B)") model = model.to(device) if args.bf16 and device.type == "cuda": model = model.to(torch.bfloat16) # ── Dataset ── print("Setting up streaming Hebrew dataset...") dataset = HebrewPretrainDataset( tokenizer=tokenizer, tok_config=tok_config, max_seq_len=args.max_seq_len, ) dataloader = DataLoader( dataset, batch_size=args.batch_size, num_workers=0, # IterableDataset with streaming - no multiprocess pin_memory=True if device.type == "cuda" else False, ) # ── Optimizer ── # Separate weight decay for non-bias, non-norm params (AdamW) decay_params = [] no_decay_params = [] for name, param in model.named_parameters(): if param.requires_grad: if "norm" in name or "bias" in name or "lambda" in name: no_decay_params.append(param) else: decay_params.append(param) optimizer = AdamW([ {"params": decay_params, "weight_decay": args.weight_decay}, {"params": no_decay_params, "weight_decay": 0.0}, ], lr=args.lr, betas=(0.9, 0.95), eps=1e-8) print(f"Optimizer: AdamW (lr={args.lr}, betas=(0.9, 0.95), wd={args.weight_decay})") print(f"Schedule: Cosine with {args.warmup_steps} warmup steps") print(f"Effective batch size: {args.batch_size * args.grad_accum}") print(f"Max steps: {args.max_steps}") # ── Tracking ── try: import trackio trackio.init( project="hebrew-dint-transformer", name=f"pretrain-{args.model_size}", ) use_trackio = True print("Trackio monitoring enabled") except Exception as e: print(f"Trackio not available: {e}") use_trackio = False # ── Training ── os.makedirs(args.output_dir, exist_ok=True) model.train() step = 0 accum_loss = 0.0 tokens_processed = 0 start_time = time.time() best_loss = float("inf") print("\n" + "="*60) print("Starting Hebrew DINT Transformer pretraining!") print("="*60 + "\n") data_iter = iter(dataloader) while step < args.max_steps: optimizer.zero_grad() # Gradient accumulation for micro_step in range(args.grad_accum): try: batch = next(data_iter) except StopIteration: # Reset data iterator (loop over data) data_iter = iter(dataloader) batch = next(data_iter) input_ids = batch["input_ids"].to(device) labels = batch["labels"].to(device) if args.bf16 and device.type == "cuda": with torch.amp.autocast("cuda", dtype=torch.bfloat16): output = model(input_ids, labels=labels) loss = output["loss"] / args.grad_accum else: output = model(input_ids, labels=labels) loss = output["loss"] / args.grad_accum loss.backward() accum_loss += loss.item() tokens_processed += input_ids.numel() # Gradient clipping grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # Update learning rate lr = get_lr(step, args.warmup_steps, args.max_steps, args.lr, args.min_lr) for param_group in optimizer.param_groups: param_group["lr"] = lr optimizer.step() step += 1 # ── Logging ── if step % args.log_every == 0: elapsed = time.time() - start_time tokens_per_sec = tokens_processed / elapsed current_loss = accum_loss / args.log_every print( f"step={step:6d} | " f"loss={current_loss:.4f} | " f"lr={lr:.2e} | " f"grad_norm={grad_norm:.2f} | " f"tok/s={tokens_per_sec:.0f} | " f"tokens={tokens_processed:,}" ) if use_trackio: trackio.log({ "train/loss": current_loss, "train/lr": lr, "train/grad_norm": grad_norm.item() if torch.is_tensor(grad_norm) else grad_norm, "train/tokens_per_sec": tokens_per_sec, "train/tokens_total": tokens_processed, "train/step": step, }) if current_loss < best_loss: best_loss = current_loss accum_loss = 0.0 # ── Save checkpoint ── if step % args.save_every == 0: ckpt_path = os.path.join(args.output_dir, f"checkpoint-{step}") os.makedirs(ckpt_path, exist_ok=True) torch.save({ "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "step": step, "best_loss": best_loss, "tokens_processed": tokens_processed, "args": vars(args), }, os.path.join(ckpt_path, "checkpoint.pt")) print(f" → Checkpoint saved to {ckpt_path}") # ── Final save & push to hub ── print("\nTraining complete! Saving final model...") final_path = os.path.join(args.output_dir, "final") os.makedirs(final_path, exist_ok=True) torch.save(model.state_dict(), os.path.join(final_path, "model.pt")) # Save model config model_config = { "architecture": "DINTTransformer", "model_size": args.model_size, "vocab_size": actual_vocab_size, "embed_dim": model.embed_dim, "num_layers": len(model.layers), "max_seq_len": model.max_seq_len, "total_params": n_params, "total_tokens_trained": tokens_processed, "best_loss": best_loss, } with open(os.path.join(final_path, "config.json"), "w") as f: json.dump(model_config, f, indent=2) # Copy tokenizer import shutil for fname in ["tokenizer.json", "config.json"]: src = os.path.join(tokenizer_path, fname) if os.path.exists(src): shutil.copy2(src, os.path.join(final_path, fname.replace("config", "tok_config") if fname == "config.json" else fname)) # Push to HF Hub print(f"Pushing to HuggingFace Hub: {args.hub_model_id}") try: from huggingface_hub import HfApi, upload_folder api = HfApi() api.create_repo(args.hub_model_id, exist_ok=True, private=False) # Upload final model + tokenizer + config upload_folder( folder_path=final_path, repo_id=args.hub_model_id, commit_message=f"Hebrew DINT Transformer ({args.model_size}) — step {step}, loss {best_loss:.4f}", ) # Upload model.py for reproducibility api.upload_file( path_or_fileobj="model.py", path_in_repo="model.py", repo_id=args.hub_model_id, ) api.upload_file( path_or_fileobj="train.py", path_in_repo="train.py", repo_id=args.hub_model_id, ) print(f"✅ Model pushed to https://huggingface.co/{args.hub_model_id}") except Exception as e: print(f"Hub push error: {e}") print("Model saved locally at:", final_path) print(f"\n{'='*60}") print(f"Hebrew DINT Transformer Pretraining Complete!") print(f" Model size: {args.model_size} ({n_params/1e9:.2f}B params)") print(f" Steps: {step}") print(f" Tokens: {tokens_processed:,}") print(f" Best loss: {best_loss:.4f}") print(f"{'='*60}") if __name__ == "__main__": args = get_args() train(args)