diff --git "a/bytefalcon_fast60m.py" "b/bytefalcon_fast60m.py" new file mode 100644--- /dev/null +++ "b/bytefalcon_fast60m.py" @@ -0,0 +1,3935 @@ +#!/usr/bin/env python3 +""" +bytefalcon_fast60m.py + +One CLI for: + * building the byte-fallback + universal-special tokenizer, + * initializing a deeper ~60M-parameter hybrid local-attention model from scratch, + * atomically appending/deduplicating/shuffling new rewrite batches, + * packing rewrite.jsonl into 4096-token byte streams, + * training/resuming on ROCm, + * running quick validation and generation. + +Expected JSONL schema: + {"instruction": "...", "text": "...", "output": "..."} +""" + +from __future__ import annotations +from collections import Counter +import argparse +import contextlib +import gc +import hashlib +import inspect +import json +import math +import os +import random +import shutil +import sqlite3 +import sys +import tempfile +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Iterable, Iterator, Mapping, Sequence + + +SCRIPT_VERSION = "3.1.1-fast60m-rocm-compile-safe" +PROJECT_DIR = Path(__file__).resolve().parent +DEFAULT_INVENTORY = PROJECT_DIR / "special_tokens.json" + +CONTROL_TOKENS = [ + "", + "", + "", + "", + "", + "", + "", + "", + "", + "", +] + +DEFAULT_ARCHITECTURE = { + "target_parameters": 60_000_000, + "hidden_size": 512, + "embedding_size": 256, + "ffn_latent_size": 256, + "num_hidden_layers": 24, + "num_attention_heads": 8, + "num_key_value_heads": 2, + "attention_every": 4, + "window_size": 512, + "conv_kernel_size": 4, + "memory_size": 128, + "memory_heads": 4, + "attention_residual_group_size": 4, + "mtp_loss_weight": 0.20, + "max_position_embeddings": 4096, +} + +# Runtime defaults for ROCm/PyTorch. +os.environ.setdefault("USE_HUB_KERNELS", "NO") +os.environ.setdefault("PYTORCH_ALLOC_CONF", "expandable_segments:True") +os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") +os.environ.setdefault("USE_ROCM_CK_GEMM", "1") +os.environ.pop("PYTORCH_HIP_ALLOC_CONF", None) + + +# --------------------------------------------------------------------------- +# Generic utilities +# --------------------------------------------------------------------------- + +def atomic_json(path: Path, value: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + file_descriptor, temporary_name = tempfile.mkstemp( + prefix=path.name + ".", + suffix=".tmp", + dir=path.parent, + ) + try: + with os.fdopen( + file_descriptor, + "w", + encoding="utf-8", + ) as handle: + json.dump( + value, + handle, + ensure_ascii=False, + indent=2, + sort_keys=True, + ) + handle.write("\n") + os.replace(temporary_name, path) + finally: + with contextlib.suppress(FileNotFoundError): + os.unlink(temporary_name) + + +def sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter( + lambda: handle.read(8 * 1024 * 1024), + b"", + ): + digest.update(block) + return digest.hexdigest() + + +def sha256_text(text: str) -> str: + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + +def now_iso() -> str: + import datetime as dt + + return dt.datetime.now(dt.timezone.utc).isoformat() + + +def configure_torch_runtime(torch: Any) -> None: + with contextlib.suppress(Exception): + torch.set_float32_matmul_precision("high") + with contextlib.suppress(Exception): + torch.backends.cuda.matmul.allow_tf32 = True + with contextlib.suppress(Exception): + torch.backends.cudnn.benchmark = True + # Prefer fused SDPA kernels on CUDA/ROCm, while retaining the math fallback. + with contextlib.suppress(Exception): + torch.backends.cuda.enable_flash_sdp(True) + with contextlib.suppress(Exception): + torch.backends.cuda.enable_mem_efficient_sdp(True) + with contextlib.suppress(Exception): + torch.backends.cuda.enable_math_sdp(True) + with contextlib.suppress(Exception): + torch._dynamo.config.cache_size_limit = 64 + + +def clear_memory(torch: Any | None = None) -> None: + gc.collect() + if torch is not None and torch.cuda.is_available(): + torch.cuda.empty_cache() + with contextlib.suppress(Exception): + torch.cuda.ipc_collect() + + +# --------------------------------------------------------------------------- +# Rewrite records and atomic corpus expansion +# --------------------------------------------------------------------------- + +@dataclass(frozen=True) +class RewriteRecord: + instruction: str + text: str + output: str + + @property + def digest(self) -> str: + value = ( + self.instruction + + "\x1f" + + self.text + + "\x1f" + + self.output + ) + return sha256_text(value) + + def to_dict(self) -> dict[str, str]: + return { + "instruction": self.instruction, + "text": self.text, + "output": self.output, + } + + +def normalize_record(value: Mapping[str, Any]) -> RewriteRecord: + missing = [ + key + for key in ("instruction", "text", "output") + if key not in value + ] + if missing: + raise ValueError( + "Missing rewrite keys: " + ", ".join(missing) + ) + + return RewriteRecord( + instruction=str(value["instruction"] or ""), + text=str(value["text"] or ""), + output=str(value["output"] or ""), + ) + + +def iter_jsonl(path: Path) -> Iterator[RewriteRecord]: + with path.open("r", encoding="utf-8-sig") as handle: + for line_number, raw_line in enumerate(handle, start=1): + line = raw_line.strip() + if not line: + continue + try: + value = json.loads(line) + except json.JSONDecodeError as error: + raise ValueError( + f"{path}:{line_number}: invalid JSON: {error}" + ) from error + if not isinstance(value, dict): + raise ValueError( + f"{path}:{line_number}: expected a JSON object." + ) + try: + yield normalize_record(value) + except ValueError as error: + raise ValueError( + f"{path}:{line_number}: {error}" + ) from error + + +def iter_json_file(path: Path) -> Iterator[RewriteRecord]: + value = json.loads(path.read_text(encoding="utf-8-sig")) + if isinstance(value, dict) and isinstance(value.get("data"), list): + value = value["data"] + if not isinstance(value, list): + raise ValueError( + f"{path}: expected a JSON list or a {{'data': [...]}} object." + ) + for index, item in enumerate(value): + if not isinstance(item, dict): + raise ValueError( + f"{path}: item {index} is not an object." + ) + yield normalize_record(item) + + +def iter_records(path: Path) -> Iterator[RewriteRecord]: + suffix = path.suffix.lower() + if suffix in {".jsonl", ".ndjson"}: + yield from iter_jsonl(path) + elif suffix == ".json": + yield from iter_json_file(path) + else: + raise ValueError(f"Unsupported batch type: {path}") + + +def discover_batches( + inbox: Path, + *, + recursive: bool, +) -> list[Path]: + patterns = ("*.jsonl", "*.ndjson", "*.json") + found: set[Path] = set() + for pattern in patterns: + iterator = ( + inbox.rglob(pattern) + if recursive + else inbox.glob(pattern) + ) + found.update( + path.resolve() + for path in iterator + if path.is_file() + ) + return sorted(found) + + +def deterministic_sort_key(seed: int, digest: str) -> str: + return sha256_text(f"{seed}:{digest}") + + +def sync_dataset(args: argparse.Namespace) -> dict[str, Any]: + base = args.data.resolve() + inbox = args.inbox.resolve() + archive = args.archive.resolve() if args.archive else None + + if not base.is_file(): + raise FileNotFoundError(f"Base dataset does not exist: {base}") + if not inbox.is_dir(): + raise FileNotFoundError(f"Inbox directory does not exist: {inbox}") + + batches = [ + path + for path in discover_batches( + inbox, + recursive=args.recursive, + ) + if path != base + ] + + work_dir = args.work_dir.resolve() + work_dir.mkdir(parents=True, exist_ok=True) + database_path = work_dir / "dataset-sync.sqlite3" + database_path.unlink(missing_ok=True) + + connection = sqlite3.connect(database_path) + connection.execute("PRAGMA journal_mode=WAL") + connection.execute("PRAGMA synchronous=NORMAL") + connection.execute("PRAGMA temp_store=FILE") + connection.execute( + """ + CREATE TABLE records ( + digest TEXT PRIMARY KEY, + sort_key TEXT NOT NULL, + instruction TEXT NOT NULL, + text_value TEXT NOT NULL, + output_value TEXT NOT NULL, + source TEXT NOT NULL + ) + """ + ) + + stats = { + "base_rows_seen": 0, + "new_rows_seen": 0, + "unique_rows": 0, + "duplicates": 0, + "invalid_files": [], + "batch_files": [str(path) for path in batches], + } + + def insert_record(record: RewriteRecord, source: str) -> None: + cursor = connection.execute( + """ + INSERT OR IGNORE INTO records + (digest, sort_key, instruction, text_value, output_value, source) + VALUES (?, ?, ?, ?, ?, ?) + """, + ( + record.digest, + deterministic_sort_key(args.seed, record.digest), + record.instruction, + record.text, + record.output, + source, + ), + ) + if cursor.rowcount == 0: + stats["duplicates"] += 1 + + with connection: + for record in iter_jsonl(base): + stats["base_rows_seen"] += 1 + insert_record(record, str(base)) + + for batch in batches: + try: + for record in iter_records(batch): + stats["new_rows_seen"] += 1 + insert_record(record, str(batch)) + except Exception as error: + stats["invalid_files"].append( + { + "path": str(batch), + "error": f"{type(error).__name__}: {error}", + } + ) + if not args.skip_invalid_files: + connection.close() + database_path.unlink(missing_ok=True) + raise + + stats["unique_rows"] = int( + connection.execute("SELECT COUNT(*) FROM records").fetchone()[0] + ) + + temporary = base.with_suffix(base.suffix + ".sync.tmp") + with temporary.open("w", encoding="utf-8") as handle: + cursor = connection.execute( + """ + SELECT instruction, text_value, output_value + FROM records + ORDER BY sort_key, digest + """ + ) + for instruction, text_value, output_value in cursor: + handle.write( + json.dumps( + { + "instruction": instruction, + "text": text_value, + "output": output_value, + }, + ensure_ascii=False, + separators=(",", ":"), + ) + ) + handle.write("\n") + handle.flush() + os.fsync(handle.fileno()) + + connection.close() + + backup = None + if args.backup: + backup = base.with_name( + f"{base.name}.before-sync-{int(time.time())}" + ) + shutil.copy2(base, backup) + + os.replace(temporary, base) + + archived = [] + if archive is not None: + archive.mkdir(parents=True, exist_ok=True) + for batch in batches: + if not batch.exists(): + continue + destination = archive / batch.name + if destination.exists(): + destination = archive / ( + f"{batch.stem}-{int(time.time())}{batch.suffix}" + ) + shutil.move(str(batch), str(destination)) + archived.append(str(destination)) + + database_path.unlink(missing_ok=True) + stats.update( + { + "data": str(base), + "sha256": sha256_file(base), + "seed": args.seed, + "backup": str(backup) if backup else None, + "archived": archived, + "completed_at": now_iso(), + } + ) + + audit = args.audit or base.with_suffix(".sync.json") + atomic_json(audit, stats) + print(json.dumps(stats, ensure_ascii=False, indent=2)) + return stats + + +# --------------------------------------------------------------------------- +# Byte-level tokenizer with universal special atoms +# --------------------------------------------------------------------------- + +def bytes_to_unicode() -> dict[int, str]: + """ + GPT-2/ByteLevel's reversible byte-to-Unicode alphabet. + """ + byte_values = ( + list(range(ord("!"), ord("~") + 1)) + + list(range(ord("¡"), ord("¬") + 1)) + + list(range(ord("®"), ord("ÿ") + 1)) + ) + unicode_values = list(byte_values) + extra = 0 + for byte_value in range(256): + if byte_value not in byte_values: + byte_values.append(byte_value) + unicode_values.append(256 + extra) + extra += 1 + return { + byte_value: chr(codepoint) + for byte_value, codepoint in zip( + byte_values, + unicode_values, + strict=True, + ) + } + + +def build_tokenizer(args: argparse.Namespace) -> dict[str, Any]: + try: + from tokenizers import AddedToken, Tokenizer, decoders, models + from tokenizers import pre_tokenizers + from transformers import PreTrainedTokenizerFast + except ImportError as error: + raise RuntimeError( + "Tokenizer construction requires tokenizers and transformers." + ) from error + + inventory_path = args.inventory.resolve() + inventory = json.loads( + inventory_path.read_text(encoding="utf-8") + ) + output_dir = args.output.resolve() + output_dir.mkdir(parents=True, exist_ok=True) + + byte_alphabet = bytes_to_unicode() + vocab: dict[str, int] = {} + + for token in CONTROL_TOKENS: + vocab[token] = len(vocab) + + byte_ids: dict[str, int] = {} + for byte_value in range(256): + token = byte_alphabet[byte_value] + vocab[token] = len(vocab) + byte_ids[f"{byte_value:02X}"] = vocab[token] + + backend = Tokenizer( + models.BPE( + vocab=vocab, + merges=[], + unk_token="", + byte_fallback=False, + ) + ) + backend.pre_tokenizer = pre_tokenizers.ByteLevel( + add_prefix_space=False, + use_regex=False, + ) + backend.decoder = decoders.ByteLevel() + + added_tokens = [] + for entry in inventory["tokens"]: + surface = entry["token"] + if surface in CONTROL_TOKENS: + continue + added_tokens.append( + AddedToken( + surface, + single_word=(entry["mode"] == "word"), + normalized=False, + lstrip=False, + rstrip=False, + special=True, + ) + ) + + backend.add_special_tokens(added_tokens) + + universal_surfaces = [ + entry["token"] + for entry in inventory["tokens"] + if entry["token"] not in CONTROL_TOKENS + ] + tokenizer = PreTrainedTokenizerFast( + tokenizer_object=backend, + bos_token="", + eos_token="", + unk_token="", + pad_token="", + additional_special_tokens=[ + *CONTROL_TOKENS[4:], + *universal_surfaces, + ], + clean_up_tokenization_spaces=False, + model_max_length=args.context_length, + ) + tokenizer.padding_side = "right" + tokenizer.truncation_side = "right" + tokenizer.save_pretrained(output_dir) + + samples = [ + "Hello, byte world.", + "0.003 + 15 = 15.003", + "encode tokens and matrices", + "😀 👍🏽 🇩🇴 👩‍💻", + "line one\nline two\tend", + "UTF-8: café, 日本語, العربية", + ] + audits = [] + for sample in samples: + ids = tokenizer.encode( + sample, + add_special_tokens=False, + ) + decoded = tokenizer.decode( + ids, + skip_special_tokens=False, + clean_up_tokenization_spaces=False, + ) + if decoded != sample: + raise RuntimeError( + f"Tokenizer round-trip failed: {sample!r} -> {decoded!r}" + ) + audits.append( + { + "text": sample, + "tokens": len(ids), + "ids": ids[:64], + } + ) + + # A string deliberately absent from the universal inventory must fall back + # to one token per UTF-8 byte. + fallback_sample = "qxjv" + fallback_ids = tokenizer.encode( + fallback_sample, + add_special_tokens=False, + ) + expected_bytes = len(fallback_sample.encode("utf-8")) + if len(fallback_ids) != expected_bytes: + raise RuntimeError( + "Strict byte fallback audit failed for qxjv: " + f"{len(fallback_ids)} != {expected_bytes}" + ) + + special_ids = set(tokenizer.all_special_ids) + universal_atomic = 0 + for surface in universal_surfaces: + ids = tokenizer.encode( + surface, + add_special_tokens=False, + ) + if len(ids) == 1 and ids[0] in special_ids: + universal_atomic += 1 + + report = { + "version": 1, + "inventory": str(inventory_path), + "inventory_sha256": sha256_file(inventory_path), + "vocab_size": len(tokenizer), + "byte_rows": 256, + "control_tokens": CONTROL_TOKENS, + "universal_special_surfaces": len(universal_surfaces), + "universal_specials_atomic": universal_atomic, + "all_special_ids_count": len(tokenizer.all_special_ids), + "context_length": args.context_length, + "byte_id_map": byte_ids, + "roundtrip_audits": audits, + "fallback_audit": { + "text": fallback_sample, + "utf8_bytes": expected_bytes, + "token_count": len(fallback_ids), + }, + "warning": ( + "Do not decode with skip_special_tokens=True: universal lexical " + "and emoji atoms are intentionally registered as special." + ), + "created_at": now_iso(), + } + atomic_json(output_dir / "byte_tokenizer_report.json", report) + + print(json.dumps(report, ensure_ascii=False, indent=2)) + return report + + +def load_tokenizer(path: Path): + from transformers import AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained( + path, + use_fast=True, + ) + tokenizer.model_max_length = 4096 + return tokenizer + + +def control_token_id_map(tokenizer: Any) -> dict[str, int]: + result: dict[str, int] = {} + for token in CONTROL_TOKENS: + token_id = tokenizer.convert_tokens_to_ids(token) + if token_id is None: + continue + token_id = int(token_id) + if token_id < 0: + continue + result[token] = token_id + return result + + +def blocked_generation_token_ids(tokenizer: Any) -> list[int]: + """ + Reserved control tokens are structural, not normal text-generation targets. + + EOS remains allowed. Lexical/emoji atoms are deliberately *not* blocked, + even though the tokenizer registers them as special tokens. + """ + allowed = {""} + mapping = control_token_id_map(tokenizer) + return sorted( + { + token_id + for token, token_id in mapping.items() + if token not in allowed + } + ) + + +def audit_packed_dataset(args: argparse.Namespace) -> dict[str, Any]: + try: + import numpy as np + except ImportError as error: + raise RuntimeError("Packed auditing requires NumPy.") from error + + tokenizer = load_tokenizer(args.tokenizer.resolve()) + packed_dir = args.packed.resolve() + mapping = control_token_id_map(tokenizer) + + report: dict[str, Any] = { + "packed": str(packed_dir), + "tokenizer": str(args.tokenizer.resolve()), + "vocab_size": len(tokenizer), + "control_ids": mapping, + "splits": {}, + } + + for split in ("train", "validation"): + path = packed_dir / f"{split}.bin" + if not path.is_file(): + continue + + values = np.memmap(path, mode="r", dtype=np.uint16) + counts = { + token: int(np.count_nonzero(values == token_id)) + for token, token_id in mapping.items() + } + invalid = int(np.count_nonzero(values >= len(tokenizer))) + report["splits"][split] = { + "path": str(path), + "tokens": int(values.size), + "control_token_counts": counts, + "invalid_token_ids": invalid, + } + + dangerous = {} + for split, details in report["splits"].items(): + hits = { + token: count + for token, count in details["control_token_counts"].items() + if token in {"", "", ""} and count > 0 + } + if hits: + dangerous[split] = hits + + report["dangerous_reserved_tokens"] = dangerous + report["healthy"] = not dangerous and all( + details["invalid_token_ids"] == 0 + for details in report["splits"].values() + ) + print(json.dumps(report, indent=2)) + return report + + +# --------------------------------------------------------------------------- +# Packing +# --------------------------------------------------------------------------- + +def format_rewrite(record: RewriteRecord) -> str: + instruction = record.instruction.strip() + quoted_text = '"' + record.text + '"' + quoted_output = '"' + record.output + '"' + if instruction: + return ( + instruction + + "\n\n" + + quoted_text + + "\n\n" + + quoted_output + ) + return quoted_text + "\n\n" + quoted_output + + +def stable_validation_record( + record: RewriteRecord, + ratio: float, +) -> bool: + threshold = int(ratio * (2**64)) + value = int(record.digest[:16], 16) + return value < threshold + + +def pack_dataset(args: argparse.Namespace) -> dict[str, Any]: + try: + import numpy as np + except ImportError as error: + raise RuntimeError("Packing requires NumPy.") from error + + data_path = args.data.resolve() + tokenizer_dir = args.tokenizer.resolve() + output_dir = args.output.resolve() + output_dir.mkdir(parents=True, exist_ok=True) + + tokenizer = load_tokenizer(tokenizer_dir) + if len(tokenizer) >= 65536: + raise RuntimeError( + "Tokenizer is too large for uint16 packing." + ) + + fingerprint = { + "data_sha256": sha256_file(data_path), + "tokenizer_sha256": sha256_file( + tokenizer_dir / "tokenizer.json" + ), + "context_length": args.context_length, + "validation_ratio": args.validation_ratio, + "format": 'instruction\\n\\n"text"\\n\\n"output"', + "packing_version": 2, + } + + manifest_path = output_dir / "packed_manifest.json" + if manifest_path.is_file() and not args.force: + existing = json.loads( + manifest_path.read_text(encoding="utf-8") + ) + if existing.get("fingerprint") == fingerprint: + print("Packed cache is current:", output_dir) + print(json.dumps(existing, indent=2)) + return existing + + temporary_dir = Path( + tempfile.mkdtemp( + prefix=output_dir.name + ".packing.", + dir=output_dir.parent, + ) + ) + train_path = temporary_dir / "train.bin" + validation_path = temporary_dir / "validation.bin" + + train_handle = train_path.open("wb") + validation_handle = validation_path.open("wb") + + buffers = { + "train": [], + "validation": [], + } + token_counts = Counter() + record_counts = Counter() + control_token_counts = { + "train": Counter(), + "validation": Counter(), + } + control_ids = control_token_id_map(tokenizer) + forbidden_control_tokens = {"", "", ""} + max_buffer = 1_000_000 + + def flush(split: str, force: bool = False) -> None: + buffer = buffers[split] + if not buffer: + return + if len(buffer) < max_buffer and not force: + return + array = np.asarray(buffer, dtype=np.uint16) + target = ( + train_handle if split == "train" else validation_handle + ) + array.tofile(target) + buffer.clear() + + eos_id = int(tokenizer.eos_token_id) + + for record in iter_jsonl(data_path): + split = ( + "validation" + if stable_validation_record( + record, + args.validation_ratio, + ) + else "train" + ) + text = format_rewrite(record) + ids = tokenizer.encode( + text, + add_special_tokens=False, + ) + + id_counts = Counter(ids) + forbidden_hits = {} + for control_token, control_id in control_ids.items(): + occurrences = int(id_counts.get(control_id, 0)) + if occurrences: + control_token_counts[split][control_token] += occurrences + if control_token in forbidden_control_tokens: + forbidden_hits[control_token] = occurrences + + if forbidden_hits: + raise RuntimeError( + "Reserved control token text was found in rewrite.jsonl. " + f"record_digest={record.digest}, hits={forbidden_hits}. " + "Remove or escape literal , , and strings " + "before packing; these tokens must never become training text." + ) + + ids.append(eos_id) + buffers[split].extend(ids) + token_counts[split] += len(ids) + record_counts[split] += 1 + flush(split) + + for split in ("train", "validation"): + flush(split, force=True) + + train_handle.flush() + validation_handle.flush() + os.fsync(train_handle.fileno()) + os.fsync(validation_handle.fileno()) + train_handle.close() + validation_handle.close() + + if token_counts["train"] <= args.context_length: + raise RuntimeError("Not enough training tokens for one block.") + if token_counts["validation"] <= args.context_length: + print( + "WARNING: validation split contains fewer than one full block." + ) + + manifest = { + "fingerprint": fingerprint, + "data": str(data_path), + "tokenizer": str(tokenizer_dir), + "dtype": "uint16", + "train_records": record_counts["train"], + "validation_records": record_counts["validation"], + "train_tokens": token_counts["train"], + "validation_tokens": token_counts["validation"], + "train_blocks": max( + 0, + (token_counts["train"] - 1) // args.context_length, + ), + "validation_blocks": max( + 0, + (token_counts["validation"] - 1) + // args.context_length, + ), + "control_token_counts": { + split: dict(counts) + for split, counts in control_token_counts.items() + }, + "created_at": now_iso(), + } + atomic_json(temporary_dir / "packed_manifest.json", manifest) + + for name in ("train.bin", "validation.bin", "packed_manifest.json"): + os.replace(temporary_dir / name, output_dir / name) + temporary_dir.rmdir() + + print(json.dumps(manifest, indent=2)) + return manifest + + +# --------------------------------------------------------------------------- +# Deep speed-first hybrid language model (~60M) +# --------------------------------------------------------------------------- + + +def import_training_stack(): + try: + import numpy as np + import torch + import torch.nn as nn + import torch.nn.functional as F + from torch.utils.data import DataLoader, Dataset + except ImportError as error: + raise RuntimeError( + "Training requires NumPy and a ROCm-enabled PyTorch build." + ) from error + + configure_torch_runtime(torch) + return np, torch, nn, F, DataLoader, Dataset + + +@dataclass +class Fast60MConfig: + vocab_size: int + padded_vocab_size: int + hidden_size: int = 512 + embedding_size: int = 256 + ffn_latent_size: int = 256 + intermediate_size: int = 1792 + num_hidden_layers: int = 24 + num_attention_heads: int = 8 + num_key_value_heads: int = 2 + attention_every: int = 4 + window_size: int = 512 + conv_kernel_size: int = 4 + memory_size: int = 128 + memory_heads: int = 4 + attention_residual_group_size: int = 4 + mtp_loss_weight: float = 0.20 + max_position_embeddings: int = 4096 + rope_theta: float = 10_000.0 + rms_norm_eps: float = 1e-5 + initializer_range: float = 0.02 + pad_token_id: int = 0 + bos_token_id: int = 1 + eos_token_id: int = 2 + model_type: str = "byte-deep-hybrid" + architecture: str = "FastDeepHybridLM" + + @property + def head_dim(self) -> int: + return self.hidden_size // self.num_attention_heads + + @property + def kv_width(self) -> int: + return self.num_key_value_heads * self.head_dim + + @property + def attention_layer_count(self) -> int: + return sum( + 1 + for index in range(self.num_hidden_layers) + if (index + 1) % self.attention_every == 0 + ) + + @property + def convolution_layer_count(self) -> int: + return self.num_hidden_layers - self.attention_layer_count + + def to_dict(self) -> dict[str, Any]: + return dict(self.__dict__) + + @classmethod + def from_dict(cls, value: Mapping[str, Any]) -> "Fast60MConfig": + fields = cls.__dataclass_fields__ + return cls(**{key: value[key] for key in fields if key in value}) + + +def round_to_multiple(value: float, multiple: int) -> int: + return max(multiple, int(round(value / multiple)) * multiple) + + +def _fixed_parameter_count(config: Fast60MConfig) -> int: + """Count every parameter except the expandable latent FFN matrices.""" + d_model = config.hidden_size + d_embed = config.embedding_size + d_latent = config.ffn_latent_size + memory = config.memory_size + kv_width = config.kv_width + + total = config.padded_vocab_size * d_embed + total += 2 * d_model * d_embed + if config.mtp_loss_weight > 0: + total += d_model * d_embed + total += d_model # final RMSNorm + + for index in range(config.num_hidden_layers): + has_attention = (index + 1) % config.attention_every == 0 + total += 1 # AttnRes-lite scalar gate. + total += d_model # mixer RMSNorm. + total += d_model # FFN RMSNorm. + total += 2 * d_model * d_latent # FFN latent down/up projections. + + if has_attention: + total += 2 * d_model * d_model + total += 2 * d_model * kv_width + total += d_model # summary-memory RMSNorm. + total += 2 * d_model * memory + 4 * memory * memory + else: + total += 3 * d_model * d_model + total += d_model * config.conv_kernel_size + + return total + + +def estimate_parameter_count(config: Fast60MConfig) -> int: + expandable = ( + config.num_hidden_layers + * 3 + * config.ffn_latent_size + * config.intermediate_size + ) + return _fixed_parameter_count(config) + expandable + + +def build_fast_config( + tokenizer: Any, + *, + target_parameters: int = 60_000_000, + hidden_size: int = 512, + embedding_size: int = 256, + ffn_latent_size: int = 256, + num_hidden_layers: int = 24, + num_attention_heads: int = 8, + num_key_value_heads: int = 2, + attention_every: int = 4, + window_size: int = 512, + conv_kernel_size: int = 4, + memory_size: int = 128, + memory_heads: int = 4, + attention_residual_group_size: int = 4, + mtp_loss_weight: float = 0.20, + context_length: int = 4096, +) -> Fast60MConfig: + if hidden_size % num_attention_heads != 0: + raise ValueError("hidden_size must be divisible by num_attention_heads.") + if num_attention_heads % num_key_value_heads != 0: + raise ValueError( + "num_attention_heads must be divisible by num_key_value_heads." + ) + if context_length % window_size != 0: + raise ValueError("context_length must be divisible by window_size.") + if memory_size % memory_heads != 0: + raise ValueError("memory_size must be divisible by memory_heads.") + if attention_every <= 0: + raise ValueError("attention_every must be positive.") + if ffn_latent_size <= 0 or ffn_latent_size > hidden_size: + raise ValueError("ffn_latent_size must be in (0, hidden_size].") + if conv_kernel_size <= 0: + raise ValueError("conv_kernel_size must be positive.") + if attention_residual_group_size <= 0: + raise ValueError("attention_residual_group_size must be positive.") + + vocab_size = len(tokenizer) + padded_vocab_size = int(math.ceil(vocab_size / 64) * 64) + provisional = Fast60MConfig( + vocab_size=vocab_size, + padded_vocab_size=padded_vocab_size, + hidden_size=hidden_size, + embedding_size=embedding_size, + ffn_latent_size=ffn_latent_size, + intermediate_size=64, + num_hidden_layers=num_hidden_layers, + num_attention_heads=num_attention_heads, + num_key_value_heads=num_key_value_heads, + attention_every=attention_every, + window_size=window_size, + conv_kernel_size=conv_kernel_size, + memory_size=memory_size, + memory_heads=memory_heads, + attention_residual_group_size=attention_residual_group_size, + mtp_loss_weight=mtp_loss_weight, + max_position_embeddings=context_length, + pad_token_id=int(tokenizer.pad_token_id), + bos_token_id=int(tokenizer.bos_token_id), + eos_token_id=int(tokenizer.eos_token_id), + ) + + fixed = _fixed_parameter_count(provisional) + denominator = num_hidden_layers * 3 * ffn_latent_size + raw_intermediate = (target_parameters - fixed) / max(1, denominator) + intermediate_size = round_to_multiple(raw_intermediate, 64) + intermediate_size = max(512, min(4096, intermediate_size)) + provisional.intermediate_size = intermediate_size + return provisional + + +def create_model_classes(torch: Any, nn: Any, F: Any): + class RMSNorm(nn.Module): + def __init__(self, width: int, eps: float): + super().__init__() + self.weight = nn.Parameter(torch.ones(width)) + self.eps = eps + + def forward(self, hidden_states): + # ROCm's fused RMSNorm requires input and weight to share a dtype. + # The cast remains differentiable, so FP32 master weights still + # receive gradients while BF16 activations use the fused kernel. + weight = self.weight + if weight.dtype != hidden_states.dtype: + weight = weight.to(dtype=hidden_states.dtype) + return F.rms_norm( + hidden_states, + (hidden_states.shape[-1],), + weight, + self.eps, + ) + + def rotate_half(value): + even = value[..., 0::2] + odd = value[..., 1::2] + return torch.stack((-odd, even), dim=-1).flatten(-2) + + class GroupedQueryWindowAttention(nn.Module): + """Windowed causal attention with cheap grouped K/V projections. + + The K/V heads are repeated only inside each local window. This keeps the + stable PyTorch SDPA path on ROCm while reducing projection parameters and + projection FLOPs relative to full multi-head QKV. + """ + + def __init__(self, config: Fast60MConfig, shifted: bool): + super().__init__() + self.hidden_size = config.hidden_size + self.num_heads = config.num_attention_heads + self.num_kv_heads = config.num_key_value_heads + self.kv_repeat = self.num_heads // self.num_kv_heads + self.head_dim = config.head_dim + self.kv_width = config.kv_width + self.window_size = config.window_size + self.shift_size = config.window_size // 2 if shifted else 0 + + self.q_proj = nn.Linear( + config.hidden_size, + config.hidden_size, + bias=False, + ) + self.k_proj = nn.Linear( + config.hidden_size, + self.kv_width, + bias=False, + ) + self.v_proj = nn.Linear( + config.hidden_size, + self.kv_width, + bias=False, + ) + self.out_proj = nn.Linear( + config.hidden_size, + config.hidden_size, + bias=False, + ) + + def _attend_segment(self, query, key, value): + batch, query_heads, length, head_dim = query.shape + if length == 0: + return query + + padding = (-length) % self.window_size + if padding: + query = F.pad(query, (0, 0, 0, padding)) + key = F.pad(key, (0, 0, 0, padding)) + value = F.pad(value, (0, 0, 0, padding)) + + padded_length = query.shape[-2] + windows = padded_length // self.window_size + + def partition(tensor, heads): + return ( + tensor.reshape( + batch, + heads, + windows, + self.window_size, + head_dim, + ) + .permute(0, 2, 1, 3, 4) + .reshape( + batch * windows, + heads, + self.window_size, + head_dim, + ) + ) + + query_windows = partition(query, query_heads) + key_windows = partition(key, self.num_kv_heads) + value_windows = partition(value, self.num_kv_heads) + if self.kv_repeat > 1: + key_windows = key_windows.repeat_interleave( + self.kv_repeat, + dim=1, + ) + value_windows = value_windows.repeat_interleave( + self.kv_repeat, + dim=1, + ) + + output = F.scaled_dot_product_attention( + query_windows, + key_windows, + value_windows, + dropout_p=0.0, + is_causal=True, + ) + output = ( + output.reshape( + batch, + windows, + query_heads, + self.window_size, + head_dim, + ) + .permute(0, 2, 1, 3, 4) + .reshape(batch, query_heads, padded_length, head_dim) + ) + return output[:, :, :length, :] + + def forward(self, hidden_states, cos, sin): + batch, length, _ = hidden_states.shape + query = self.q_proj(hidden_states).view( + batch, + length, + self.num_heads, + self.head_dim, + ).transpose(1, 2) + key = self.k_proj(hidden_states).view( + batch, + length, + self.num_kv_heads, + self.head_dim, + ).transpose(1, 2) + value = self.v_proj(hidden_states).view( + batch, + length, + self.num_kv_heads, + self.head_dim, + ).transpose(1, 2) + + query = query * cos + rotate_half(query) * sin + key = key * cos + rotate_half(key) * sin + + if self.shift_size and length > self.shift_size: + prefix = self.shift_size + first = self._attend_segment( + query[:, :, :prefix], + key[:, :, :prefix], + value[:, :, :prefix], + ) + rest = self._attend_segment( + query[:, :, prefix:], + key[:, :, prefix:], + value[:, :, prefix:], + ) + output = torch.cat((first, rest), dim=-2) + else: + output = self._attend_segment(query, key, value) + + output = output.transpose(1, 2).contiguous().view( + batch, + length, + self.hidden_size, + ) + return self.out_proj(output) + + class CausalShortConvMixer(nn.Module): + """KDA-inspired short causal path for non-attention layers. + + This is deliberately not a literal Kimi Delta Attention port: exact KDA + needs custom recurrent kernels to be fast. The short depthwise convolution + keeps local high-frequency mixing at O(sequence) cost using stock ROCm ops. + """ + + def __init__(self, config: Fast60MConfig): + super().__init__() + self.hidden_size = config.hidden_size + self.kernel_size = config.conv_kernel_size + self.in_proj = nn.Linear( + config.hidden_size, + 2 * config.hidden_size, + bias=False, + ) + self.depthwise_weight = nn.Parameter( + torch.empty(config.hidden_size, 1, self.kernel_size) + ) + self.out_proj = nn.Linear( + config.hidden_size, + config.hidden_size, + bias=False, + ) + nn.init.normal_( + self.depthwise_weight, + mean=0.0, + std=config.initializer_range, + ) + + def forward(self, hidden_states): + length = hidden_states.shape[1] + value, gate = self.in_proj(hidden_states).chunk(2, dim=-1) + value = F.conv1d( + value.transpose(1, 2), + self.depthwise_weight, + padding=self.kernel_size - 1, + groups=self.hidden_size, + )[..., :length].transpose(1, 2) + return self.out_proj(F.silu(value) * torch.sigmoid(gate)) + + class SummaryMemoryMixer(nn.Module): + """Cheap causal communication across completed local windows.""" + + def __init__(self, config: Fast60MConfig): + super().__init__() + self.hidden_size = config.hidden_size + self.memory_size = config.memory_size + self.memory_heads = config.memory_heads + self.memory_head_dim = ( + config.memory_size // config.memory_heads + ) + self.window_size = config.window_size + self.down = nn.Linear( + config.hidden_size, + config.memory_size, + bias=False, + ) + self.qkv = nn.Linear( + config.memory_size, + 3 * config.memory_size, + bias=False, + ) + self.out = nn.Linear( + config.memory_size, + config.memory_size, + bias=False, + ) + self.up = nn.Linear( + config.memory_size, + config.hidden_size, + bias=False, + ) + + def forward(self, hidden_states): + batch, length, width = hidden_states.shape + padding = (-length) % self.window_size + padded = ( + F.pad(hidden_states, (0, 0, 0, padding)) + if padding + else hidden_states + ) + windows = padded.view( + batch, + padded.shape[1] // self.window_size, + self.window_size, + width, + ) + summaries = windows[:, :, -1, :] + if padding: + summaries = torch.cat( + (summaries[:, :-1], hidden_states[:, -1:, :]), + dim=1, + ) + + summaries = self.down(summaries) + query, key, value = self.qkv(summaries).chunk(3, dim=-1) + window_count = summaries.shape[1] + + def split_heads(tensor): + return tensor.view( + batch, + window_count, + self.memory_heads, + self.memory_head_dim, + ).transpose(1, 2) + + query = split_heads(query) + key = split_heads(key) + value = split_heads(value) + memory = F.scaled_dot_product_attention( + query, + key, + value, + dropout_p=0.0, + is_causal=True, + ) + memory = memory.transpose(1, 2).contiguous().view( + batch, + window_count, + self.memory_size, + ) + memory = self.up(self.out(memory)) + + previous_memory = torch.cat( + (torch.zeros_like(memory[:, :1]), memory[:, :-1]), + dim=1, + ) + broadcast = ( + previous_memory[:, :, None, :] + .expand(-1, -1, self.window_size, -1) + .reshape(batch, padded.shape[1], width) + ) + return broadcast[:, :length] + + class LatentSwiGLU(nn.Module): + """Stable-LatentMoE-inspired dense FFN bottleneck. + + All tokens use the same dense FFN, but its expensive expansion operates + at ffn_latent_size instead of the full residual width. This is much more + single-GPU friendly than sparse MoE while preserving the latent-compute idea. + """ + + def __init__(self, config: Fast60MConfig): + super().__init__() + self.down_in = nn.Linear( + config.hidden_size, + config.ffn_latent_size, + bias=False, + ) + self.gate_up = nn.Linear( + config.ffn_latent_size, + 2 * config.intermediate_size, + bias=False, + ) + self.down = nn.Linear( + config.intermediate_size, + config.ffn_latent_size, + bias=False, + ) + self.up_out = nn.Linear( + config.ffn_latent_size, + config.hidden_size, + bias=False, + ) + + def forward(self, hidden_states): + latent = self.down_in(hidden_states) + gate, up = self.gate_up(latent).chunk(2, dim=-1) + latent = self.down(F.silu(gate) * up) + return self.up_out(latent) + + class FastBlock(nn.Module): + def __init__(self, config: Fast60MConfig, index: int): + super().__init__() + self.index = index + self.has_attention = ( + (index + 1) % config.attention_every == 0 + ) + attention_rank = index // config.attention_every + self.depth_residual_gate = nn.Parameter(torch.zeros(())) + self.mixer_norm = RMSNorm( + config.hidden_size, + config.rms_norm_eps, + ) + if self.has_attention: + self.mixer = GroupedQueryWindowAttention( + config, + shifted=(attention_rank % 2 == 1), + ) + self.memory_norm = RMSNorm( + config.hidden_size, + config.rms_norm_eps, + ) + self.memory_mixer = SummaryMemoryMixer(config) + else: + self.mixer = CausalShortConvMixer(config) + self.memory_norm = None + self.memory_mixer = None + self.ffn_norm = RMSNorm( + config.hidden_size, + config.rms_norm_eps, + ) + self.feed_forward = LatentSwiGLU(config) + + def forward(self, hidden_states, cos, sin, depth_anchor): + # AttnRes-lite: each layer can retrieve its group's earlier residual + # stream through one learned scalar, initialized as an exact no-op. + mixer_source = hidden_states + torch.tanh( + self.depth_residual_gate + ) * depth_anchor + normalized = self.mixer_norm(mixer_source) + if self.has_attention: + hidden_states = hidden_states + self.mixer( + normalized, + cos, + sin, + ) + hidden_states = hidden_states + self.memory_mixer( + self.memory_norm(hidden_states) + ) + else: + hidden_states = hidden_states + self.mixer(normalized) + hidden_states = hidden_states + self.feed_forward( + self.ffn_norm(hidden_states) + ) + return hidden_states + + @dataclass + class FastLMOutput: + loss: Any | None = None + logits: Any | None = None + main_loss: Any | None = None + mtp_loss: Any | None = None + + class FastDeepHybridLM(nn.Module): + def __init__(self, config: Fast60MConfig): + super().__init__() + self.config = config + self.token_embedding = nn.Embedding( + config.padded_vocab_size, + config.embedding_size, + ) + self.embedding_projection = nn.Linear( + config.embedding_size, + config.hidden_size, + bias=False, + ) + self.blocks = nn.ModuleList( + FastBlock(config, index) + for index in range(config.num_hidden_layers) + ) + self.final_norm = RMSNorm( + config.hidden_size, + config.rms_norm_eps, + ) + self.output_projection = nn.Linear( + config.hidden_size, + config.embedding_size, + bias=False, + ) + self.mtp_projection = ( + nn.Linear( + config.hidden_size, + config.embedding_size, + bias=False, + ) + if config.mtp_loss_weight > 0 + else None + ) + + inverse_frequency = 1.0 / ( + config.rope_theta + ** ( + torch.arange(0, config.head_dim, 2).float() + / config.head_dim + ) + ) + positions = torch.arange( + config.max_position_embeddings, + dtype=torch.float32, + ) + frequencies = torch.outer(positions, inverse_frequency) + embedding = torch.repeat_interleave(frequencies, 2, dim=-1) + self.register_buffer( + "rope_cos", + embedding.cos()[None, None, :, :], + persistent=False, + ) + self.register_buffer( + "rope_sin", + embedding.sin()[None, None, :, :], + persistent=False, + ) + self.apply(self._initialize_weights) + residual_std = config.initializer_range / math.sqrt( + 2 * config.num_hidden_layers + ) + for block in self.blocks: + if block.has_attention: + nn.init.normal_( + block.mixer.out_proj.weight, + mean=0.0, + std=residual_std, + ) + nn.init.normal_( + block.memory_mixer.up.weight, + mean=0.0, + std=residual_std, + ) + else: + nn.init.normal_( + block.mixer.out_proj.weight, + mean=0.0, + std=residual_std, + ) + nn.init.normal_( + block.feed_forward.up_out.weight, + mean=0.0, + std=residual_std, + ) + + def _initialize_weights(self, module): + if isinstance(module, (nn.Linear, nn.Embedding)): + nn.init.normal_( + module.weight, + mean=0.0, + std=self.config.initializer_range, + ) + + def get_input_embeddings(self): + return self.token_embedding + + def _project_logits(self, hidden_states, projection=None): + active_projection = ( + self.output_projection if projection is None else projection + ) + vocabulary_states = active_projection(hidden_states) + logits = F.linear( + vocabulary_states, + self.token_embedding.weight, + ) + return logits[..., : self.config.vocab_size] + + def forward( + self, + input_ids, + labels=None, + return_last_logits: bool = False, + use_mtp: bool = True, + ): + if input_ids.ndim != 2: + raise ValueError("input_ids must have shape [batch, sequence].") + sequence_length = input_ids.shape[1] + if sequence_length > self.config.max_position_embeddings: + raise ValueError( + f"Sequence length {sequence_length} exceeds " + f"{self.config.max_position_embeddings}." + ) + + hidden_states = self.embedding_projection( + self.token_embedding(input_ids) + ) + cos = self.rope_cos[:, :, :sequence_length].to( + dtype=hidden_states.dtype + ) + sin = self.rope_sin[:, :, :sequence_length].to( + dtype=hidden_states.dtype + ) + depth_anchor = hidden_states + group_size = self.config.attention_residual_group_size + for index, block in enumerate(self.blocks): + if index % group_size == 0: + depth_anchor = hidden_states + hidden_states = block( + hidden_states, + cos, + sin, + depth_anchor, + ) + hidden_states = self.final_norm(hidden_states) + + if labels is not None: + logits = self._project_logits(hidden_states[:, :-1]) + main_loss = F.cross_entropy( + logits.reshape(-1, self.config.vocab_size), + labels[:, 1:].reshape(-1), + ) + mtp_loss = None + loss = main_loss + if ( + use_mtp + and self.mtp_projection is not None + and sequence_length > 2 + ): + mtp_logits = self._project_logits( + hidden_states[:, :-2], + self.mtp_projection, + ) + mtp_loss = F.cross_entropy( + mtp_logits.reshape(-1, self.config.vocab_size), + labels[:, 2:].reshape(-1), + ) + loss = loss + self.config.mtp_loss_weight * mtp_loss + return FastLMOutput( + loss=loss, + logits=None, + main_loss=main_loss, + mtp_loss=mtp_loss, + ) + + if return_last_logits: + hidden_states = hidden_states[:, -1:, :] + logits = self._project_logits(hidden_states) + return FastLMOutput(loss=None, logits=logits) + + return FastDeepHybridLM, FastLMOutput + + +def count_parameters(model: Any) -> dict[str, int]: + total = sum(parameter.numel() for parameter in model.parameters()) + trainable = sum( + parameter.numel() + for parameter in model.parameters() + if parameter.requires_grad + ) + embedding = model.get_input_embeddings().weight.numel() + return { + "total": total, + "trainable": trainable, + "embedding": embedding, + "non_embedding": total - embedding, + } + + +def save_model_bundle( + model: Any, + tokenizer: Any, + output_dir: Path, + torch: Any, +) -> None: + output_dir.mkdir(parents=True, exist_ok=True) + atomic_json(output_dir / "config.json", model.config.to_dict()) + torch.save(model.state_dict(), output_dir / "model.pt") + tokenizer.save_pretrained(output_dir) + + +def load_model_bundle(path: Path, torch: Any, nn: Any, F: Any): + config = Fast60MConfig.from_dict( + json.loads((path / "config.json").read_text(encoding="utf-8")) + ) + model_class, _ = create_model_classes(torch, nn, F) + model = model_class(config) + try: + state = torch.load( + path / "model.pt", + map_location="cpu", + weights_only=True, + ) + except TypeError: + state = torch.load(path / "model.pt", map_location="cpu") + try: + model.load_state_dict(state, strict=True) + except RuntimeError as error: + raise RuntimeError( + "Checkpoint is not architecture-compatible with fast60m-hybrid. " + "Start a new run or use a checkpoint created by this script." + ) from error + return model + + +def _config_from_args(tokenizer: Any, args: argparse.Namespace, context: int): + return build_fast_config( + tokenizer, + target_parameters=args.target_parameters, + hidden_size=args.hidden_size, + embedding_size=args.embedding_size, + ffn_latent_size=args.ffn_latent_size, + num_hidden_layers=args.layers, + num_attention_heads=args.heads, + num_key_value_heads=args.kv_heads, + attention_every=args.attention_every, + window_size=args.window_size, + conv_kernel_size=args.conv_kernel_size, + memory_size=args.memory_size, + memory_heads=args.memory_heads, + attention_residual_group_size=args.attention_residual_group_size, + mtp_loss_weight=args.mtp_loss_weight, + context_length=context, + ) + + +def initialize_model(args: argparse.Namespace) -> dict[str, Any]: + np, torch, nn, F, DataLoader, Dataset = import_training_stack() + del np, DataLoader, Dataset + + tokenizer = load_tokenizer(args.tokenizer.resolve()) + config = _config_from_args(tokenizer, args, args.context_length) + model_class, _ = create_model_classes(torch, nn, F) + model = model_class(config) + parameters = count_parameters(model) + + output_dir = args.output.resolve() + save_model_bundle(model, tokenizer, output_dir, torch) + + full_attention_projection = 4 * config.hidden_size * config.hidden_size + gqa_projection = ( + 2 * config.hidden_size * config.hidden_size + + 2 * config.hidden_size * config.kv_width + ) + dense_ffn = 3 * config.hidden_size * config.intermediate_size + latent_ffn = ( + 2 * config.hidden_size * config.ffn_latent_size + + 3 * config.ffn_latent_size * config.intermediate_size + ) + report = { + "parameters": parameters, + "estimated_parameters": estimate_parameter_count(config), + "parameters_millions": parameters["total"] / 1_000_000, + "config": config.to_dict(), + "speed_design": { + "depth": config.num_hidden_layers, + "attention_layers": config.attention_layer_count, + "linear_conv_layers": config.convolution_layer_count, + "attention_fraction": ( + config.attention_layer_count / config.num_hidden_layers + ), + "attention_window": config.window_size, + "full_context": config.max_position_embeddings, + "attention_pair_fraction_vs_full": ( + config.window_size / config.max_position_embeddings + ), + "gqa_projection_fraction_vs_mha": ( + gqa_projection / full_attention_projection + ), + "latent_ffn_parameter_fraction_vs_full": ( + latent_ffn / dense_ffn + ), + "factorized_embedding_head": True, + "causal_summary_memory_on_attention_layers_only": True, + "attention_residuals_lite": True, + "multi_token_prediction": config.mtp_loss_weight > 0, + "external_custom_kernels_required": False, + }, + "created_at": now_iso(), + } + atomic_json(output_dir / "initialization_report.json", report) + + lower = int(args.target_parameters * 0.90) + upper = int(args.target_parameters * 1.10) + if not (lower <= parameters["total"] <= upper): + raise RuntimeError( + f"Model is outside the requested ~{args.target_parameters / 1e6:.0f}M " + f"range: {parameters['total']:,}. Adjust width, depth, or target." + ) + + print(json.dumps(report, indent=2)) + return report + + +# --------------------------------------------------------------------------- +# Training helpers +# --------------------------------------------------------------------------- + + +def find_latest_checkpoint(output_dir: Path) -> Path | None: + checkpoint_root = output_dir / "checkpoints" + if not checkpoint_root.is_dir(): + return None + candidates = sorted( + ( + path + for path in checkpoint_root.glob("step-*") + if path.is_dir() + ), + key=lambda path: int(path.name.split("-")[-1]), + ) + return candidates[-1] if candidates else None + + +def checkpoint_step(path: Path | None) -> int: + if path is None: + return 0 + return int(path.name.split("-")[-1]) + + +def prune_checkpoints(root: Path, keep: int) -> None: + candidates = sorted( + ( + path + for path in root.glob("step-*") + if path.is_dir() + ), + key=lambda path: int(path.name.split("-")[-1]), + ) + for path in candidates[:-keep]: + shutil.rmtree(path) + + +def load_training_state(torch: Any, path: Path) -> dict[str, Any]: + try: + return torch.load( + path, + map_location="cpu", + weights_only=False, + ) + except TypeError: + return torch.load(path, map_location="cpu") + + +def build_adamw(torch: Any, model: Any, args: argparse.Namespace): + common = dict( + params=model.parameters(), + lr=args.learning_rate, + betas=(args.beta1, args.beta2), + eps=args.adam_epsilon, + weight_decay=args.weight_decay, + ) + if args.fused_optimizer: + try: + optimizer = torch.optim.AdamW(**common, fused=True) + return optimizer, "fused" + except (TypeError, RuntimeError) as error: + print( + "Fused AdamW unavailable; falling back to foreach AdamW:", + error, + ) + try: + return torch.optim.AdamW(**common, foreach=True), "foreach" + except (TypeError, RuntimeError): + return torch.optim.AdamW(**common), "single-tensor" + + +def make_scheduler( + torch: Any, + optimizer: Any, + *, + warmup_steps: int, + total_steps: int, + minimum_ratio: float, +): + def multiplier(step: int) -> float: + if step < warmup_steps: + return max(1e-8, float(step + 1) / max(1, warmup_steps)) + progress = ( + float(step - warmup_steps) + / max(1, total_steps - warmup_steps) + ) + progress = min(1.0, max(0.0, progress)) + cosine = 0.5 * (1.0 + math.cos(math.pi * progress)) + return minimum_ratio + (1.0 - minimum_ratio) * cosine + + return torch.optim.lr_scheduler.LambdaLR(optimizer, multiplier) + + +def _atomic_replace_directory( + temporary: Path, + destination: Path, +) -> None: + previous = destination.with_name( + destination.name + f".previous-{os.getpid()}" + ) + if previous.exists(): + shutil.rmtree(previous) + + if destination.exists(): + os.replace(destination, previous) + + try: + os.replace(temporary, destination) + except Exception: + if previous.exists() and not destination.exists(): + os.replace(previous, destination) + raise + else: + if previous.exists(): + shutil.rmtree(previous) + + +def save_named_training_checkpoint( + *, + model: Any, + tokenizer: Any, + optimizer: Any, + scheduler: Any, + torch: Any, + destination: Path, + state: dict[str, Any], + metadata: Mapping[str, Any] | None = None, +) -> Path: + destination.parent.mkdir(parents=True, exist_ok=True) + temporary = destination.with_name( + destination.name + f".tmp-{os.getpid()}" + ) + if temporary.exists(): + shutil.rmtree(temporary) + temporary.mkdir(parents=True) + + save_model_bundle(model, tokenizer, temporary, torch) + torch.save( + { + "optimizer": optimizer.state_dict(), + "scheduler": scheduler.state_dict(), + "state": state, + "torch_rng": torch.get_rng_state(), + "cuda_rng": ( + torch.cuda.get_rng_state_all() + if torch.cuda.is_available() + else None + ), + "python_rng": random.getstate(), + }, + temporary / "training_state.pt", + ) + atomic_json(temporary / "training_state.json", state) + if metadata is not None: + atomic_json( + temporary / "checkpoint_metadata.json", + dict(metadata), + ) + + required = ( + temporary / "config.json", + temporary / "model.pt", + temporary / "training_state.pt", + ) + missing = [str(path) for path in required if not path.is_file()] + if missing: + shutil.rmtree(temporary, ignore_errors=True) + raise RuntimeError( + f"Checkpoint write was incomplete; missing={missing}" + ) + + _atomic_replace_directory(temporary, destination) + return destination + + +def save_checkpoint( + *, + model: Any, + tokenizer: Any, + optimizer: Any, + scheduler: Any, + torch: Any, + output_dir: Path, + state: dict[str, Any], + keep: int, +) -> Path: + checkpoint_root = output_dir / "checkpoints" + destination = checkpoint_root / f"step-{state['global_step']:08d}" + save_named_training_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + destination=destination, + state=state, + metadata={ + "kind": "regular", + "global_step": state["global_step"], + "saved_at": now_iso(), + }, + ) + prune_checkpoints(checkpoint_root, keep) + return destination + + +def restore_training_checkpoint_in_place( + *, + checkpoint: Path, + model: Any, + optimizer: Any, + scheduler: Any, + state: dict[str, Any], + torch: Any, + nn: Any, + F: Any, +) -> None: + restored_model = load_model_bundle( + checkpoint, + torch, + nn, + F, + ) + model.load_state_dict( + restored_model.state_dict(), + strict=True, + ) + del restored_model + + saved = load_training_state( + torch, + checkpoint / "training_state.pt", + ) + optimizer.load_state_dict(saved["optimizer"]) + scheduler.load_state_dict(saved["scheduler"]) + state.clear() + state.update(saved["state"]) + + if saved.get("torch_rng") is not None: + torch.set_rng_state(saved["torch_rng"]) + if ( + torch.cuda.is_available() + and saved.get("cuda_rng") is not None + ): + torch.cuda.set_rng_state_all(saved["cuda_rng"]) + if saved.get("python_rng") is not None: + random.setstate(saved["python_rng"]) + + +def backoff_learning_rate( + optimizer: Any, + scheduler: Any, + *, + factor: float, + minimum: float, +) -> list[float]: + updated = [] + for group in optimizer.param_groups: + new_lr = max( + minimum, + float(group["lr"]) * factor, + ) + group["lr"] = new_lr + group["initial_lr"] = min( + float(group.get("initial_lr", new_lr)), + new_lr, + ) + updated.append(new_lr) + + if hasattr(scheduler, "base_lrs"): + scheduler.base_lrs = [ + max(minimum, float(value) * factor) + for value in scheduler.base_lrs + ] + if hasattr(scheduler, "_last_lr"): + scheduler._last_lr = list(updated) + + return updated + + +def model_parameters_are_finite( + torch: Any, + model: Any, +) -> bool: + with torch.no_grad(): + for parameter in model.parameters(): + if not bool(torch.isfinite(parameter).all().item()): + return False + return True + + +def evaluate_loss( + *, + model: Any, + loader: Any, + torch: Any, + device: Any, + dtype_name: str, + max_batches: int, +) -> float | None: + """ + Evaluate with mixed precision first. If a batch becomes non-finite, retry + that batch in FP32 before declaring the checkpoint unhealthy. + """ + model.eval() + total = 0.0 + count = 0 + autocast_dtype = ( + torch.bfloat16 if dtype_name == "bf16" else torch.float16 + ) + autocast_enabled = dtype_name in {"bf16", "fp16"} + + try: + with torch.no_grad(): + for batch_index, batch in enumerate(loader): + if batch_index >= max_batches: + break + + input_ids = batch.to( + device, + non_blocking=True, + ) + with torch.autocast( + device_type="cuda", + dtype=autocast_dtype, + enabled=autocast_enabled, + ): + output = model( + input_ids=input_ids, + labels=input_ids, + use_mtp=False, + ) + loss = output.loss.detach() + + if not bool(torch.isfinite(loss).item()): + print( + "Validation loss was non-finite under autocast; " + f"retrying batch {batch_index} in FP32." + ) + with torch.autocast( + device_type="cuda", + enabled=False, + ): + output = model( + input_ids=input_ids, + labels=input_ids, + use_mtp=False, + ) + loss = output.loss.detach().float() + + if not bool(torch.isfinite(loss).item()): + return float("nan") + + total += float(loss.item()) + count += 1 + finally: + if torch.cuda.is_available(): + torch.cuda.synchronize() + model.train() + + if count == 0: + return None + return total / count + + + +def train_model(args: argparse.Namespace) -> dict[str, Any]: + np, torch, nn, F, DataLoader, Dataset = import_training_stack() + + if not torch.cuda.is_available(): + raise RuntimeError( + "ROCm PyTorch did not expose the AMD GPU through torch.cuda." + ) + + device = torch.device("cuda") + torch.manual_seed(args.seed) + random.seed(args.seed) + torch.cuda.manual_seed_all(args.seed) + + tokenizer = load_tokenizer(args.tokenizer.resolve()) + packed_dir = args.packed.resolve() + manifest = json.loads( + (packed_dir / "packed_manifest.json").read_text(encoding="utf-8") + ) + context_length = int(manifest["fingerprint"]["context_length"]) + if context_length != 4096: + raise RuntimeError( + f"This project expects 4096-token blocks, got {context_length}." + ) + + class TokenBlocks(Dataset): + def __init__(self, path: Path, context: int): + self.tokens = np.memmap(path, mode="r", dtype=np.uint16) + self.context = context + self.blocks = max(0, (len(self.tokens) - 1) // context) + + def __len__(self): + return self.blocks + + def __getitem__(self, index): + start = index * self.context + values = np.asarray( + self.tokens[start : start + self.context], + dtype=np.int64, + ).copy() + return torch.from_numpy(values) + + train_dataset = TokenBlocks( + packed_dir / "train.bin", + context_length, + ) + validation_dataset = TokenBlocks( + packed_dir / "validation.bin", + context_length, + ) + if len(train_dataset) == 0: + raise RuntimeError("Packed training dataset has zero blocks.") + + output_dir = args.output.resolve() + output_dir.mkdir(parents=True, exist_ok=True) + checkpoint = ( + find_latest_checkpoint(output_dir) + if args.resume == "auto" + else ( + Path(args.resume).resolve() + if args.resume != "none" + else None + ) + ) + + if checkpoint is not None: + print("Resuming checkpoint:", checkpoint) + model = load_model_bundle(checkpoint, torch, nn, F) + if model.config.vocab_size != len(tokenizer): + raise RuntimeError( + "Checkpoint tokenizer size does not match --tokenizer." + ) + else: + config = _config_from_args(tokenizer, args, context_length) + model_class, _ = create_model_classes(torch, nn, F) + model = model_class(config) + + model.to(device) + model.train() + parameters = count_parameters(model) + print( + f"Architecture: {model.config.architecture}; " + f"parameters={parameters['total']:,} " + f"({parameters['total'] / 1e6:.3f}M); " + f"layers={model.config.num_hidden_layers}; " + f"attention_layers={model.config.attention_layer_count}; " + f"conv_layers={model.config.convolution_layer_count}; " + f"window={model.config.window_size}; " + f"ffn_latent={model.config.ffn_latent_size}" + ) + + optimizer, optimizer_backend = build_adamw(torch, model, args) + + updates_per_epoch = math.ceil( + len(train_dataset) + / max(1, args.batch_size * args.gradient_accumulation) + ) + run_target_steps = ( + args.max_steps + if args.max_steps > 0 + else max(1, args.epochs * updates_per_epoch) + ) + schedule_steps = ( + args.lr_decay_steps + if args.lr_decay_steps > 0 + else run_target_steps + ) + schedule_steps = max(schedule_steps, run_target_steps) + warmup_steps = ( + args.warmup_steps + if args.warmup_steps >= 0 + else int(schedule_steps * args.warmup_ratio) + ) + scheduler = make_scheduler( + torch, + optimizer, + warmup_steps=warmup_steps, + total_steps=schedule_steps, + minimum_ratio=args.minimum_lr_ratio, + ) + + current_packed_fingerprint = manifest["fingerprint"] + state = { + "global_step": 0, + "epoch": 0, + "batch_in_epoch": 0, + "tokens_seen": 0, + "best_validation_loss": None, + "best_checkpoint": None, + "nonfinite_events": 0, + "last_finite_step": 0, + "packed_fingerprint": current_packed_fingerprint, + "lr_decay_steps": schedule_steps, + "warmup_steps": warmup_steps, + "started_at": now_iso(), + } + + if checkpoint is not None: + saved = load_training_state( + torch, + checkpoint / "training_state.pt", + ) + optimizer.load_state_dict(saved["optimizer"]) + scheduler.load_state_dict(saved["scheduler"]) + saved_state = dict(saved["state"]) + previous_fingerprint = saved_state.get("packed_fingerprint") + state.update(saved_state) + + if previous_fingerprint != current_packed_fingerprint: + print( + "Packed dataset changed; resetting epoch/batch cursor while " + "preserving model, optimizer, scheduler, and global step." + ) + state["epoch"] = 0 + state["batch_in_epoch"] = 0 + state["best_validation_loss"] = None + state["packed_fingerprint"] = current_packed_fingerprint + + state["lr_decay_steps"] = schedule_steps + state["warmup_steps"] = warmup_steps + torch.set_rng_state(saved["torch_rng"]) + if saved.get("cuda_rng") is not None: + torch.cuda.set_rng_state_all(saved["cuda_rng"]) + random.setstate(saved["python_rng"]) + + def loader_for_epoch(epoch: int): + generator = torch.Generator() + generator.manual_seed(args.seed + epoch) + loader_kwargs = dict( + dataset=train_dataset, + batch_size=args.batch_size, + shuffle=True, + generator=generator, + num_workers=args.num_workers, + pin_memory=args.pin_memory, + drop_last=True, + persistent_workers=( + args.num_workers > 0 and args.persistent_workers + ), + ) + if args.num_workers > 0: + loader_kwargs["prefetch_factor"] = args.prefetch_factor + return DataLoader(**loader_kwargs) + + validation_loader = DataLoader( + validation_dataset, + batch_size=args.batch_size, + shuffle=False, + num_workers=0, + pin_memory=args.pin_memory, + drop_last=False, + ) + + best_dir = output_dir / "best" + recovery_dir = output_dir / "recovery" + + # A recovery checkpoint always exists before the first optimizer update. + save_named_training_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + destination=recovery_dir, + state=state, + metadata={ + "kind": "recovery", + "global_step": state["global_step"], + "saved_at": now_iso(), + }, + ) + + # Keep a persistent best candidate even before the first finite validation. + if not ( + (best_dir / "config.json").is_file() + and (best_dir / "model.pt").is_file() + ): + save_named_training_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + destination=best_dir, + state=state, + metadata={ + "kind": "best", + "provisional": True, + "validation_loss": state.get("best_validation_loss"), + "global_step": state["global_step"], + "saved_at": now_iso(), + }, + ) + state["best_checkpoint"] = str(best_dir) + + if ( + args.eval_at_start + and len(validation_dataset) > 0 + ): + starting_validation_loss = evaluate_loss( + model=model, + loader=validation_loader, + torch=torch, + device=device, + dtype_name=args.dtype, + max_batches=args.eval_batches, + ) + print( + "starting validation " + f"step={state['global_step']:,} " + f"loss={starting_validation_loss}" + ) + if ( + starting_validation_loss is not None + and math.isfinite(starting_validation_loss) + and ( + state["best_validation_loss"] is None + or starting_validation_loss + < state["best_validation_loss"] + ) + ): + state["best_validation_loss"] = starting_validation_loss + state["best_checkpoint"] = str(best_dir) + save_named_training_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + destination=best_dir, + state=state, + metadata={ + "kind": "best", + "provisional": False, + "validation_loss": starting_validation_loss, + "global_step": state["global_step"], + "saved_at": now_iso(), + }, + ) + + training_model = model + compile_status = "disabled" + effective_compile_mode = args.compile_mode + compile_uses_cudagraphs = False + + if args.compile: + is_rocm = getattr(torch.version, "hip", None) is not None + + # reduce-overhead explicitly relies on CUDA Graphs. With gradient + # accumulation, repeated compiled forward/backward calls can overwrite + # graph-owned outputs before autograd has finished consuming them. + if is_rocm and effective_compile_mode == "reduce-overhead": + effective_compile_mode = "default" + print( + "ROCm safety: replacing compile mode 'reduce-overhead' " + "with 'default' to avoid CUDAGraph output reuse." + ) + + compile_options = None + if is_rocm: + compile_options = {"triton.cudagraphs": False} + + compile_kwargs = { + "mode": effective_compile_mode, + "fullgraph": args.compile_fullgraph, + "dynamic": False, + } + if compile_options is not None: + compile_kwargs["options"] = compile_options + + try: + training_model = torch.compile( + model, + **compile_kwargs, + ) + compile_status = ( + f"enabled:{effective_compile_mode}:cudagraphs-disabled" + if is_rocm + else f"enabled:{effective_compile_mode}" + ) + print( + "torch.compile enabled:", + effective_compile_mode, + "(CUDAGraphs disabled on ROCm)" + if is_rocm + else "", + ) + except (TypeError, RuntimeError) as option_error: + # Older builds may reject the explicit option. Retry with the + # default mode, which does not request reduce-overhead graphs. + if compile_options is not None: + try: + training_model = torch.compile( + model, + mode="default", + fullgraph=args.compile_fullgraph, + dynamic=False, + ) + effective_compile_mode = "default" + compile_status = ( + "enabled:default:option-fallback" + ) + print( + "torch.compile option fallback enabled in default " + "mode after:", + option_error, + ) + except Exception as error: + compile_status = ( + f"setup-failed:{type(error).__name__}" + ) + training_model = model + print( + "torch.compile setup failed; using eager mode:", + error, + ) + else: + compile_status = ( + f"setup-failed:{type(option_error).__name__}" + ) + training_model = model + print( + "torch.compile setup failed; using eager mode:", + option_error, + ) + except Exception as error: + compile_status = f"setup-failed:{type(error).__name__}" + training_model = model + print("torch.compile setup failed; using eager mode:", error) + + autocast_dtype = ( + torch.bfloat16 if args.dtype == "bf16" else torch.float16 + ) + autocast_enabled = args.dtype in {"bf16", "fp16"} + scaler = None + if args.dtype == "fp16": + scaler = torch.amp.GradScaler("cuda") + + optimizer.zero_grad(set_to_none=True) + accumulation = 0 + running_loss = torch.zeros((), device=device) + running_microbatches = 0 + nonfinite_loss_seen = torch.zeros( + (), + device=device, + dtype=torch.bool, + ) + last_log_time = time.perf_counter() + last_log_tokens = state["tokens_seen"] + stop = False + + def recover_from_nonfinite( + reason: str, + batch_index: int, + ) -> None: + nonlocal training_model + nonlocal compile_status + nonlocal accumulation + nonlocal running_microbatches + nonlocal nonfinite_loss_seen + + event_count = int(state.get("nonfinite_events", 0)) + 1 + print( + f"NON-FINITE TRAINING EVENT #{event_count}: {reason}" + ) + + optimizer.zero_grad(set_to_none=True) + accumulation = 0 + running_loss.zero_() + running_microbatches = 0 + nonfinite_loss_seen.zero_() + + if scaler is not None: + current_scale = float(scaler.get_scale()) + with contextlib.suppress(Exception): + scaler.update(max(1.0, current_scale * args.nan_lr_factor)) + + if args.nan_action == "stop": + raise FloatingPointError( + f"Stopping after non-finite training state: {reason}" + ) + + if args.nan_action == "rollback": + restore_training_checkpoint_in_place( + checkpoint=recovery_dir, + model=model, + optimizer=optimizer, + scheduler=scheduler, + state=state, + torch=torch, + nn=nn, + F=F, + ) + training_model = model + if compile_status.startswith("enabled"): + compile_status = "disabled-after-nonfinite" + print( + "Rolled back to recovery checkpoint:", + recovery_dir, + ) + + state["nonfinite_events"] = event_count + state["last_nonfinite_reason"] = reason + state["batch_in_epoch"] = batch_index + 1 + + new_lrs = backoff_learning_rate( + optimizer, + scheduler, + factor=args.nan_lr_factor, + minimum=args.min_learning_rate, + ) + print("Learning-rate fallback:", new_lrs) + + save_named_training_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + destination=recovery_dir, + state=state, + metadata={ + "kind": "recovery", + "reason": reason, + "nonfinite_events": event_count, + "global_step": state["global_step"], + "saved_at": now_iso(), + }, + ) + clear_memory(torch) + + if event_count > args.max_nan_recoveries: + raise FloatingPointError( + "Exceeded --max-nan-recoveries=" + f"{args.max_nan_recoveries}." + ) + + while not stop: + epoch = int(state["epoch"]) + if args.max_steps <= 0 and epoch >= args.epochs: + break + + loader = loader_for_epoch(epoch) + resume_batch = int(state["batch_in_epoch"]) + + for batch_index, batch in enumerate(loader): + if batch_index < resume_batch: + continue + + input_ids = batch.to(device, non_blocking=True) + + def forward_backward(active_model): + # This marker is harmless when CUDAGraphs are disabled, and + # protects compatible compiled modes that still use graph + # iteration tracking internally. + if active_model is not model: + marker = getattr( + getattr(torch, "compiler", None), + "cudagraph_mark_step_begin", + None, + ) + if marker is not None: + marker() + + with torch.autocast( + device_type="cuda", + dtype=autocast_dtype, + enabled=autocast_enabled, + ): + output = active_model( + input_ids=input_ids, + labels=input_ids, + ) + scaled_loss = output.loss / args.gradient_accumulation + if scaler is None: + scaled_loss.backward() + else: + scaler.scale(scaled_loss).backward() + return output.loss.detach() + + try: + detached_loss = forward_backward(training_model) + except Exception as error: + if training_model is not model: + print( + "torch.compile failed during training; discarding " + "the current accumulation window and continuing in " + "eager mode:", + f"{type(error).__name__}: {error}", + ) + optimizer.zero_grad(set_to_none=True) + accumulation = 0 + running_loss.zero_() + running_microbatches = 0 + nonfinite_loss_seen.zero_() + training_model = model + compile_status = ( + f"runtime-failed:{type(error).__name__}:eager-fallback" + ) + with contextlib.suppress(Exception): + torch._dynamo.reset() + clear_memory(torch) + detached_loss = forward_backward(model) + else: + raise + + nonfinite_loss_seen.logical_or_( + ~torch.isfinite(detached_loss) + ) + accumulation += 1 + running_loss += torch.nan_to_num( + detached_loss, + nan=0.0, + posinf=0.0, + neginf=0.0, + ) + running_microbatches += 1 + state["tokens_seen"] += int(input_ids.numel()) + state["batch_in_epoch"] = batch_index + 1 + + if accumulation < args.gradient_accumulation: + continue + + if scaler is not None: + scaler.unscale_(optimizer) + + grad_norm = torch.nn.utils.clip_grad_norm_( + model.parameters(), + ( + args.max_grad_norm + if args.max_grad_norm > 0 + else float("inf") + ), + error_if_nonfinite=False, + ) + loss_was_nonfinite = bool( + nonfinite_loss_seen.item() + ) + grad_norm_value = float( + grad_norm.detach().float().item() + ) + + if ( + loss_was_nonfinite + or not math.isfinite(grad_norm_value) + ): + recover_from_nonfinite( + ( + "non-finite loss" + if loss_was_nonfinite + else f"non-finite grad norm={grad_norm_value}" + ), + batch_index, + ) + continue + + if scaler is None: + optimizer.step() + else: + scaler.step(optimizer) + scaler.update() + scheduler.step() + optimizer.zero_grad(set_to_none=True) + accumulation = 0 + nonfinite_loss_seen.zero_() + + prospective_step = int(state["global_step"]) + 1 + if ( + args.finite_check_every > 0 + and prospective_step % args.finite_check_every == 0 + and not model_parameters_are_finite(torch, model) + ): + recover_from_nonfinite( + "non-finite model parameters after optimizer.step()", + batch_index, + ) + continue + + state["global_step"] = prospective_step + state["last_finite_step"] = prospective_step + step = prospective_step + if step % args.log_every == 0: + torch.cuda.synchronize() + now = time.perf_counter() + elapsed = max(1e-9, now - last_log_time) + delta_tokens = state["tokens_seen"] - last_log_tokens + tokens_per_second = delta_tokens / elapsed + mean_loss = float( + (running_loss / max(1, running_microbatches)).item() + ) + memory = torch.cuda.max_memory_allocated() / (1024**3) + print( + f"step={step:,} " + f"loss={mean_loss:.5f} " + f"lr={scheduler.get_last_lr()[0]:.3e} " + f"tok/s={tokens_per_second:,.0f} " + f"tokens={state['tokens_seen']:,} " + f"peak_gib={memory:.2f}" + ) + running_loss.zero_() + running_microbatches = 0 + last_log_time = now + last_log_tokens = state["tokens_seen"] + torch.cuda.reset_peak_memory_stats() + + if ( + args.eval_every > 0 + and step % args.eval_every == 0 + and len(validation_dataset) > 0 + ): + validation_loss = evaluate_loss( + model=model, + loader=validation_loader, + torch=torch, + device=device, + dtype_name=args.dtype, + max_batches=args.eval_batches, + ) + print( + f"validation step={step:,} loss={validation_loss}" + ) + if ( + validation_loss is not None + and not math.isfinite(validation_loss) + ): + recover_from_nonfinite( + "validation remained non-finite after FP32 retry", + batch_index, + ) + continue + + if ( + validation_loss is not None + and math.isfinite(validation_loss) + and ( + state["best_validation_loss"] is None + or validation_loss + < state["best_validation_loss"] + ) + ): + state["best_validation_loss"] = validation_loss + state["best_checkpoint"] = str(best_dir) + save_named_training_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + destination=best_dir, + state=state, + metadata={ + "kind": "best", + "provisional": False, + "validation_loss": validation_loss, + "global_step": step, + "saved_at": now_iso(), + }, + ) + print( + "New best checkpoint:", + best_dir, + f"validation_loss={validation_loss}", + ) + + if args.save_every > 0 and step % args.save_every == 0: + destination = save_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + output_dir=output_dir, + state=state, + keep=args.keep_checkpoints, + ) + print("Saved:", destination) + save_named_training_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + destination=recovery_dir, + state=state, + metadata={ + "kind": "recovery", + "source_checkpoint": str(destination), + "global_step": step, + "saved_at": now_iso(), + }, + ) + + if args.max_steps > 0 and step >= args.max_steps: + stop = True + break + + if stop: + break + state["epoch"] = epoch + 1 + state["batch_in_epoch"] = 0 + + final_checkpoint = save_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + output_dir=output_dir, + state=state, + keep=args.keep_checkpoints, + ) + save_named_training_checkpoint( + model=model, + tokenizer=tokenizer, + optimizer=optimizer, + scheduler=scheduler, + torch=torch, + destination=recovery_dir, + state=state, + metadata={ + "kind": "recovery", + "source_checkpoint": str(final_checkpoint), + "global_step": state["global_step"], + "saved_at": now_iso(), + }, + ) + + final_dir = output_dir / "final" + if final_dir.exists(): + shutil.rmtree(final_dir) + save_model_bundle(model, tokenizer, final_dir, torch) + + result = { + "state": state, + "parameters": parameters, + "architecture": model.config.to_dict(), + "optimizer_backend": optimizer_backend, + "compile_status": compile_status, + "final_checkpoint": str(final_checkpoint), + "final_model": str(final_dir), + "packed_manifest": manifest, + "schedule": { + "run_target_steps": run_target_steps, + "lr_decay_steps": schedule_steps, + "warmup_steps": warmup_steps, + "minimum_lr_ratio": args.minimum_lr_ratio, + }, + "completed_at": now_iso(), + } + atomic_json(output_dir / "training_result.json", result) + print(json.dumps(result, indent=2)) + return result + + +# --------------------------------------------------------------------------- +# Generation, benchmarking, and diagnostics +# --------------------------------------------------------------------------- + + +def doctor(args: argparse.Namespace) -> dict[str, Any]: + report: dict[str, Any] = { + "python": sys.version, + "script_version": SCRIPT_VERSION, + "environment": { + "PYTORCH_ALLOC_CONF": os.environ.get("PYTORCH_ALLOC_CONF"), + "TOKENIZERS_PARALLELISM": os.environ.get( + "TOKENIZERS_PARALLELISM" + ), + "USE_ROCM_CK_GEMM": os.environ.get("USE_ROCM_CK_GEMM"), + }, + } + try: + np, torch, nn, F, DataLoader, Dataset = import_training_stack() + del np, DataLoader, Dataset + report["torch"] = { + "version": torch.__version__, + "hip": getattr(torch.version, "hip", None), + "cuda_available": torch.cuda.is_available(), + "device_count": torch.cuda.device_count(), + "device_name": ( + torch.cuda.get_device_name(0) + if torch.cuda.is_available() + else None + ), + "bf16_supported": ( + torch.cuda.is_bf16_supported() + if torch.cuda.is_available() + else False + ), + "compile_available": hasattr(torch, "compile"), + } + if torch.cuda.is_available(): + query = torch.randn( + 1, + 8, + 128, + 64, + device="cuda", + dtype=torch.bfloat16, + ) + with torch.no_grad(): + output = F.scaled_dot_product_attention( + query, + query, + query, + is_causal=True, + ) + torch.cuda.synchronize() + report["sdpa_probe"] = { + "ok": True, + "shape": list(output.shape), + } + del query, output + clear_memory(torch) + except Exception as error: + report["error"] = f"{type(error).__name__}: {error}" + + print(json.dumps(report, indent=2)) + return report + + +def inspect_project(args: argparse.Namespace) -> dict[str, Any]: + tokenizer = load_tokenizer(args.tokenizer.resolve()) + report = { + "tokenizer_vocab_size": len(tokenizer), + "special_ids": len(tokenizer.all_special_ids), + "model_max_length": tokenizer.model_max_length, + "default_architecture": DEFAULT_ARCHITECTURE, + } + model_path = args.model.resolve() if args.model else None + if model_path and (model_path / "config.json").is_file(): + report["saved_model_config"] = json.loads( + (model_path / "config.json").read_text(encoding="utf-8") + ) + print(json.dumps(report, indent=2)) + return report + + +def benchmark_model(args: argparse.Namespace) -> dict[str, Any]: + np, torch, nn, F, DataLoader, Dataset = import_training_stack() + del np, DataLoader, Dataset + if not torch.cuda.is_available(): + raise RuntimeError("ROCm GPU is unavailable.") + + device = torch.device("cuda") + if args.model: + model = load_model_bundle(args.model.resolve(), torch, nn, F) + else: + tokenizer = load_tokenizer(args.tokenizer.resolve()) + config = _config_from_args(tokenizer, args, args.context_length) + model_class, _ = create_model_classes(torch, nn, F) + model = model_class(config) + + model.to(device).train() + active_model = model + compile_status = "disabled" + if args.compile: + active_model = torch.compile( + model, + mode=args.compile_mode, + fullgraph=args.compile_fullgraph, + dynamic=False, + ) + compile_status = f"enabled:{args.compile_mode}" + + input_ids = torch.randint( + 0, + model.config.vocab_size, + (args.batch_size, args.context_length), + device=device, + ) + optimizer, optimizer_backend = build_adamw(torch, model, args) + autocast_dtype = ( + torch.bfloat16 if args.dtype == "bf16" else torch.float16 + ) + autocast_enabled = args.dtype in {"bf16", "fp16"} + + def iteration(): + optimizer.zero_grad(set_to_none=True) + with torch.autocast( + device_type="cuda", + dtype=autocast_dtype, + enabled=autocast_enabled, + ): + output = active_model(input_ids=input_ids, labels=input_ids) + output.loss.backward() + optimizer.step() + return output.loss + + for _ in range(args.warmup): + iteration() + torch.cuda.synchronize() + torch.cuda.reset_peak_memory_stats() + started = time.perf_counter() + last_loss = None + for _ in range(args.steps): + last_loss = iteration() + torch.cuda.synchronize() + elapsed = time.perf_counter() - started + tokens = args.steps * args.batch_size * args.context_length + + result = { + "tokens_per_second": tokens / elapsed, + "seconds": elapsed, + "steps": args.steps, + "batch_size": args.batch_size, + "context_length": args.context_length, + "loss": ( + float(last_loss.detach().item()) + if last_loss is not None + else None + ), + "peak_gib": torch.cuda.max_memory_allocated() / (1024**3), + "parameters": count_parameters(model), + "config": model.config.to_dict(), + "compile_status": compile_status, + "optimizer_backend": optimizer_backend, + } + print(json.dumps(result, indent=2)) + return result + + +def generate_text(args: argparse.Namespace) -> str: + np, torch, nn, F, DataLoader, Dataset = import_training_stack() + del np, DataLoader, Dataset + if not torch.cuda.is_available(): + raise RuntimeError("ROCm GPU is unavailable.") + + model_path = args.model.resolve() + tokenizer = load_tokenizer(model_path) + model = load_model_bundle(model_path, torch, nn, F).to("cuda") + model.eval() + + encoded = tokenizer( + args.prompt, + add_special_tokens=False, + return_tensors="pt", + return_token_type_ids=False, + ) + input_ids = encoded.input_ids.to("cuda") + prompt_length = int(input_ids.shape[1]) + + blocked_ids = ( + [] + if args.allow_control_tokens + else blocked_generation_token_ids(tokenizer) + ) + blocked_tensor = ( + torch.tensor( + blocked_ids, + device="cuda", + dtype=torch.long, + ) + if blocked_ids + else None + ) + + generated: list[int] = [] + with torch.no_grad(): + for generation_step in range(args.max_new_tokens): + model_input = input_ids[ + :, -model.config.max_position_embeddings : + ] + + with torch.autocast( + device_type="cuda", + dtype=torch.bfloat16, + enabled=True, + ): + logits = model( + input_ids=model_input, + return_last_logits=True, + ).logits[:, -1, :] + + if not bool(torch.isfinite(logits).all().item()): + print( + "Non-finite generation logits under BF16; " + "retrying this token in FP32.", + file=sys.stderr, + ) + with torch.autocast( + device_type="cuda", + enabled=False, + ): + logits = model( + input_ids=model_input, + return_last_logits=True, + ).logits[:, -1, :].float() + + if not bool(torch.isfinite(logits).all().item()): + logits = torch.nan_to_num( + logits, + nan=-float("inf"), + posinf=1e4, + neginf=-1e4, + ) + + if args.show_top_tokens > 0: + top_values, top_indices = torch.topk( + logits, + min(args.show_top_tokens, logits.shape[-1]), + dim=-1, + ) + decoded = [ + { + "id": int(token_id), + "token": tokenizer.decode( + [int(token_id)], + skip_special_tokens=False, + clean_up_tokenization_spaces=False, + ), + "logit": float(value), + } + for token_id, value in zip( + top_indices[0].tolist(), + top_values[0].float().tolist(), + ) + ] + print( + f"raw top tokens at generation step {generation_step}: " + + json.dumps(decoded, ensure_ascii=False), + file=sys.stderr, + ) + + # Structural control IDs must not compete with real text. + # EOS remains available and ends generation normally. + if blocked_tensor is not None: + logits.index_fill_( + 1, + blocked_tensor, + -float("inf"), + ) + + if args.repetition_penalty != 1.0: + used = torch.unique(model_input) + selected = logits[:, used] + logits[:, used] = torch.where( + selected < 0, + selected * args.repetition_penalty, + selected / args.repetition_penalty, + ) + + if not bool(torch.isfinite(logits).any().item()): + next_token = torch.tensor( + [[int(tokenizer.eos_token_id)]], + device="cuda", + dtype=torch.long, + ) + elif args.temperature <= 0: + next_token = logits.argmax(dim=-1, keepdim=True) + else: + logits = logits / max(args.temperature, 1e-5) + if args.top_k > 0: + threshold = torch.topk( + logits, + min(args.top_k, logits.shape[-1]), + dim=-1, + ).values[:, -1:] + logits = logits.masked_fill( + logits < threshold, + -float("inf"), + ) + probabilities = torch.softmax(logits, dim=-1) + if args.top_p < 1.0: + sorted_probabilities, sorted_indices = torch.sort( + probabilities, + descending=True, + dim=-1, + ) + cumulative = sorted_probabilities.cumsum(dim=-1) + remove = cumulative > args.top_p + remove[:, 1:] = remove[:, :-1].clone() + remove[:, 0] = False + sorted_probabilities = ( + sorted_probabilities.masked_fill(remove, 0.0) + ) + denominator = sorted_probabilities.sum( + dim=-1, + keepdim=True, + ).clamp_min(1e-12) + sorted_probabilities /= denominator + sampled = torch.multinomial( + sorted_probabilities, + 1, + ) + next_token = sorted_indices.gather(-1, sampled) + else: + next_token = torch.multinomial(probabilities, 1) + + token_id = int(next_token.item()) + if token_id in blocked_ids: + raise RuntimeError( + "A blocked structural control token escaped masking: " + f"id={token_id}, token={tokenizer.decode([token_id], skip_special_tokens=False)!r}" + ) + + generated.append(token_id) + input_ids = torch.cat((input_ids, next_token), dim=-1) + if token_id == int(tokenizer.eos_token_id): + break + + completion = tokenizer.decode( + generated, + skip_special_tokens=False, + clean_up_tokenization_spaces=False, + ) + print(completion) + return completion + + +# --------------------------------------------------------------------------- +# Combined quick cycle +# --------------------------------------------------------------------------- + + +def cycle(args: argparse.Namespace) -> None: + sync_namespace = argparse.Namespace( + data=args.data, + inbox=args.inbox, + archive=args.archive, + work_dir=args.work_dir / "sync", + seed=args.seed, + recursive=args.recursive, + skip_invalid_files=False, + backup=args.backup, + audit=args.work_dir / "last_sync.json", + ) + sync_dataset(sync_namespace) + + pack_namespace = argparse.Namespace( + data=args.data, + tokenizer=args.tokenizer, + output=args.packed, + context_length=4096, + validation_ratio=args.validation_ratio, + force=False, + ) + pack_dataset(pack_namespace) + + latest = find_latest_checkpoint(args.output.resolve()) + current_step = checkpoint_step(latest) + target_step = current_step + args.additional_steps + + train_namespace = argparse.Namespace( + tokenizer=args.tokenizer, + packed=args.packed, + output=args.output, + resume="auto", + seed=args.seed, + dtype="bf16", + batch_size=args.batch_size, + gradient_accumulation=args.gradient_accumulation, + learning_rate=args.learning_rate, + beta1=0.9, + beta2=0.95, + adam_epsilon=1e-8, + weight_decay=0.1, + max_grad_norm=1.0, + max_steps=target_step, + epochs=1, + warmup_steps=-1, + warmup_ratio=0.02, + minimum_lr_ratio=0.1, + lr_decay_steps=args.lr_decay_steps, + log_every=args.log_every, + eval_every=args.eval_every, + eval_batches=args.eval_batches, + save_every=args.save_every, + keep_checkpoints=args.keep_checkpoints, + num_workers=args.num_workers, + pin_memory=True, + persistent_workers=args.num_workers > 0, + prefetch_factor=2, + compile=args.compile, + compile_mode=args.compile_mode, + compile_fullgraph=args.compile_fullgraph, + fused_optimizer=True, + target_parameters=args.target_parameters, + hidden_size=args.hidden_size, + embedding_size=args.embedding_size, + ffn_latent_size=args.ffn_latent_size, + layers=args.layers, + heads=args.heads, + kv_heads=args.kv_heads, + attention_every=args.attention_every, + window_size=args.window_size, + conv_kernel_size=args.conv_kernel_size, + memory_size=args.memory_size, + memory_heads=args.memory_heads, + attention_residual_group_size=args.attention_residual_group_size, + mtp_loss_weight=args.mtp_loss_weight, + eval_at_start=True, + nan_action="rollback", + nan_lr_factor=0.5, + min_learning_rate=1e-7, + max_nan_recoveries=20, + finite_check_every=100, + ) + train_model(train_namespace) + + +# --------------------------------------------------------------------------- +# CLI +# --------------------------------------------------------------------------- + + +def add_architecture_arguments(parser: argparse.ArgumentParser) -> None: + parser.add_argument( + "--target-parameters", + type=int, + default=60_000_000, + ) + parser.add_argument("--hidden-size", type=int, default=512) + parser.add_argument("--embedding-size", type=int, default=256) + parser.add_argument("--ffn-latent-size", type=int, default=256) + parser.add_argument("--layers", type=int, default=24) + parser.add_argument("--heads", type=int, default=8) + parser.add_argument("--kv-heads", type=int, default=2) + parser.add_argument("--attention-every", type=int, default=4) + parser.add_argument("--window-size", type=int, default=512) + parser.add_argument("--conv-kernel-size", type=int, default=4) + parser.add_argument("--memory-size", type=int, default=128) + parser.add_argument("--memory-heads", type=int, default=4) + parser.add_argument( + "--attention-residual-group-size", + type=int, + default=4, + ) + parser.add_argument("--mtp-loss-weight", type=float, default=0.20) + + +def add_compile_arguments(parser: argparse.ArgumentParser) -> None: + parser.add_argument( + "--compile", + action=argparse.BooleanOptionalAction, + default=True, + ) + parser.add_argument( + "--compile-mode", + choices=[ + "default", + "reduce-overhead", + "max-autotune", + "max-autotune-no-cudagraphs", + ], + default="default", + ) + parser.add_argument( + "--compile-fullgraph", + action=argparse.BooleanOptionalAction, + default=False, + ) + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description=( + "Train a deeper speed-first ~60M hybrid byte language model." + ), + formatter_class=argparse.ArgumentDefaultsHelpFormatter, + ) + subcommands = parser.add_subparsers(dest="command", required=True) + + tokenizer_parser = subcommands.add_parser( + "tokenizer", + help="Build the byte + universal-special tokenizer.", + ) + tokenizer_parser.add_argument( + "--inventory", + type=Path, + default=DEFAULT_INVENTORY, + ) + tokenizer_parser.add_argument( + "--output", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + tokenizer_parser.add_argument( + "--context-length", + type=int, + default=4096, + ) + tokenizer_parser.set_defaults(function=build_tokenizer) + + sync_parser = subcommands.add_parser( + "sync-data", + help=( + "Append, exact-dedupe, deterministically shuffle, and archive " + "new batches from a directory." + ), + ) + sync_parser.add_argument("--data", type=Path, default=Path("rewrite.jsonl")) + sync_parser.add_argument("--inbox", type=Path, required=True) + sync_parser.add_argument("--archive", type=Path) + sync_parser.add_argument( + "--work-dir", + type=Path, + default=Path(".bytefalcon-work/sync"), + ) + sync_parser.add_argument("--seed", type=int, default=42) + sync_parser.add_argument("--recursive", action="store_true") + sync_parser.add_argument("--skip-invalid-files", action="store_true") + sync_parser.add_argument("--backup", action="store_true") + sync_parser.add_argument("--audit", type=Path) + sync_parser.set_defaults(function=sync_dataset) + + pack_parser = subcommands.add_parser( + "pack", + help="Pack rewrite.jsonl into train/validation uint16 streams.", + ) + pack_parser.add_argument("--data", type=Path, default=Path("rewrite.jsonl")) + pack_parser.add_argument( + "--tokenizer", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + pack_parser.add_argument( + "--output", + type=Path, + default=Path("artifacts/packed-4096"), + ) + pack_parser.add_argument("--context-length", type=int, default=4096) + pack_parser.add_argument("--validation-ratio", type=float, default=0.005) + pack_parser.add_argument("--force", action="store_true") + pack_parser.set_defaults(function=pack_dataset) + + audit_parser = subcommands.add_parser( + "audit-packed", + help="Count reserved control IDs inside packed train/validation streams.", + ) + audit_parser.add_argument( + "--tokenizer", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + audit_parser.add_argument( + "--packed", + type=Path, + default=Path("artifacts/packed-4096"), + ) + audit_parser.set_defaults(function=audit_packed_dataset) + + init_parser = subcommands.add_parser( + "init", + help="Initialize and save the deeper ~60M speed-first model.", + ) + init_parser.add_argument( + "--tokenizer", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + init_parser.add_argument( + "--output", + type=Path, + default=Path("runs/bytefast-60m/initial"), + ) + init_parser.add_argument("--context-length", type=int, default=4096) + add_architecture_arguments(init_parser) + init_parser.set_defaults(function=initialize_model) + + train_parser = subcommands.add_parser( + "train", + help="Train from scratch or resume a checkpoint.", + ) + train_parser.add_argument( + "--tokenizer", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + train_parser.add_argument( + "--packed", + type=Path, + default=Path("artifacts/packed-4096"), + ) + train_parser.add_argument( + "--output", + type=Path, + default=Path("runs/bytefast-60m"), + ) + train_parser.add_argument( + "--resume", + default="auto", + help="'auto', 'none', or a checkpoint path.", + ) + train_parser.add_argument("--seed", type=int, default=42) + train_parser.add_argument( + "--dtype", + choices=["bf16", "fp16", "fp32"], + default="bf16", + ) + train_parser.add_argument("--batch-size", type=int, default=4) + train_parser.add_argument( + "--gradient-accumulation", + type=int, + default=4, + ) + train_parser.add_argument("--learning-rate", type=float, default=2e-5) + train_parser.add_argument("--beta1", type=float, default=0.9) + train_parser.add_argument("--beta2", type=float, default=0.95) + train_parser.add_argument("--adam-epsilon", type=float, default=1e-8) + train_parser.add_argument("--weight-decay", type=float, default=0.1) + train_parser.add_argument("--max-grad-norm", type=float, default=1.0) + train_parser.add_argument("--max-steps", type=int, default=0) + train_parser.add_argument("--epochs", type=int, default=1) + train_parser.add_argument("--warmup-steps", type=int, default=-1) + train_parser.add_argument("--warmup-ratio", type=float, default=0.02) + train_parser.add_argument("--minimum-lr-ratio", type=float, default=0.1) + train_parser.add_argument("--lr-decay-steps", type=int, default=100000) + train_parser.add_argument("--log-every", type=int, default=100) + train_parser.add_argument("--eval-every", type=int, default=500) + train_parser.add_argument("--eval-batches", type=int, default=8) + train_parser.add_argument("--save-every", type=int, default=1000) + train_parser.add_argument("--keep-checkpoints", type=int, default=5) + train_parser.add_argument( + "--eval-at-start", + action=argparse.BooleanOptionalAction, + default=True, + help="Evaluate and materialize best/ before the first optimizer update.", + ) + train_parser.add_argument( + "--nan-action", + choices=["rollback", "skip", "stop"], + default="rollback", + ) + train_parser.add_argument("--nan-lr-factor", type=float, default=0.5) + train_parser.add_argument("--min-learning-rate", type=float, default=1e-7) + train_parser.add_argument("--max-nan-recoveries", type=int, default=20) + train_parser.add_argument( + "--finite-check-every", + type=int, + default=100, + help="Scan all model parameters for NaN/Inf every N optimizer steps.", + ) + train_parser.add_argument("--num-workers", type=int, default=2) + train_parser.add_argument( + "--pin-memory", + action=argparse.BooleanOptionalAction, + default=True, + ) + train_parser.add_argument( + "--persistent-workers", + action=argparse.BooleanOptionalAction, + default=True, + ) + train_parser.add_argument("--prefetch-factor", type=int, default=2) + train_parser.add_argument( + "--fused-optimizer", + action=argparse.BooleanOptionalAction, + default=True, + ) + add_compile_arguments(train_parser) + add_architecture_arguments(train_parser) + train_parser.set_defaults(function=train_model) + + cycle_parser = subcommands.add_parser( + "cycle", + help="Sync, repack if changed, and resume fast training.", + ) + cycle_parser.add_argument("--data", type=Path, default=Path("rewrite.jsonl")) + cycle_parser.add_argument("--inbox", type=Path, required=True) + cycle_parser.add_argument("--archive", type=Path) + cycle_parser.add_argument( + "--work-dir", + type=Path, + default=Path(".bytefalcon-work"), + ) + cycle_parser.add_argument( + "--tokenizer", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + cycle_parser.add_argument( + "--packed", + type=Path, + default=Path("artifacts/packed-4096"), + ) + cycle_parser.add_argument( + "--output", + type=Path, + default=Path("runs/bytefast-60m"), + ) + cycle_parser.add_argument("--additional-steps", type=int, default=500) + cycle_parser.add_argument("--seed", type=int, default=42) + cycle_parser.add_argument("--recursive", action="store_true") + cycle_parser.add_argument("--backup", action="store_true") + cycle_parser.add_argument("--validation-ratio", type=float, default=0.005) + cycle_parser.add_argument("--batch-size", type=int, default=4) + cycle_parser.add_argument("--gradient-accumulation", type=int, default=4) + cycle_parser.add_argument("--learning-rate", type=float, default=2e-5) + cycle_parser.add_argument("--lr-decay-steps", type=int, default=100000) + cycle_parser.add_argument("--log-every", type=int, default=100) + cycle_parser.add_argument("--eval-every", type=int, default=500) + cycle_parser.add_argument("--eval-batches", type=int, default=8) + cycle_parser.add_argument("--save-every", type=int, default=500) + cycle_parser.add_argument("--keep-checkpoints", type=int, default=5) + cycle_parser.add_argument("--num-workers", type=int, default=4) + add_compile_arguments(cycle_parser) + add_architecture_arguments(cycle_parser) + cycle_parser.set_defaults(function=cycle) + + doctor_parser = subcommands.add_parser( + "doctor", + help="Audit ROCm, bf16, torch.compile, and SDPA.", + ) + doctor_parser.add_argument( + "--tokenizer", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + doctor_parser.set_defaults(function=doctor) + + inspect_parser = subcommands.add_parser( + "inspect", + help="Show tokenizer and architecture details.", + ) + inspect_parser.add_argument( + "--tokenizer", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + inspect_parser.add_argument("--model", type=Path) + inspect_parser.set_defaults(function=inspect_project) + + benchmark_parser = subcommands.add_parser( + "benchmark", + help="Measure steady-state training throughput on the GPU.", + ) + benchmark_parser.add_argument("--model", type=Path) + benchmark_parser.add_argument( + "--tokenizer", + type=Path, + default=Path("artifacts/byte-tokenizer"), + ) + benchmark_parser.add_argument("--context-length", type=int, default=4096) + benchmark_parser.add_argument("--batch-size", type=int, default=2) + benchmark_parser.add_argument("--warmup", type=int, default=3) + benchmark_parser.add_argument("--steps", type=int, default=100) + benchmark_parser.add_argument( + "--dtype", + choices=["bf16", "fp16", "fp32"], + default="bf16", + ) + benchmark_parser.add_argument("--learning-rate", type=float, default=2e-5) + benchmark_parser.add_argument("--beta1", type=float, default=0.9) + benchmark_parser.add_argument("--beta2", type=float, default=0.95) + benchmark_parser.add_argument("--adam-epsilon", type=float, default=1e-8) + benchmark_parser.add_argument("--weight-decay", type=float, default=0.1) + benchmark_parser.add_argument( + "--fused-optimizer", + action=argparse.BooleanOptionalAction, + default=True, + ) + add_compile_arguments(benchmark_parser) + add_architecture_arguments(benchmark_parser) + benchmark_parser.set_defaults(function=benchmark_model) + + generate_parser = subcommands.add_parser( + "generate", + help="Generate from a trained checkpoint.", + ) + generate_parser.add_argument("--model", type=Path, required=True) + generate_parser.add_argument("--prompt", required=True) + generate_parser.add_argument("--max-new-tokens", type=int, default=128) + generate_parser.add_argument("--temperature", type=float, default=0.7) + generate_parser.add_argument("--top-p", type=float, default=0.95) + generate_parser.add_argument("--top-k", type=int, default=50) + generate_parser.add_argument("--repetition-penalty", type=float, default=1.1) + generate_parser.add_argument( + "--allow-control-tokens", + action="store_true", + help="Allow structural IDs such as ; disabled by default.", + ) + generate_parser.add_argument( + "--show-top-tokens", + type=int, + default=0, + help="Print the raw top-N logits before reserved-token masking.", + ) + generate_parser.set_defaults(function=generate_text) + + return parser + + +def validate_args(args: argparse.Namespace) -> None: + if hasattr(args, "context_length") and args.context_length != 4096: + raise ValueError("This project is fixed to context length 4096.") + if hasattr(args, "validation_ratio") and not ( + 0.0 <= args.validation_ratio < 0.5 + ): + raise ValueError("--validation-ratio must be in [0, 0.5).") + if hasattr(args, "batch_size") and args.batch_size <= 0: + raise ValueError("--batch-size must be positive.") + if ( + hasattr(args, "gradient_accumulation") + and args.gradient_accumulation <= 0 + ): + raise ValueError("--gradient-accumulation must be positive.") + if hasattr(args, "lr_decay_steps") and args.lr_decay_steps <= 0: + raise ValueError("--lr-decay-steps must be positive.") + if hasattr(args, "window_size") and 4096 % args.window_size != 0: + raise ValueError("--window-size must divide 4096 exactly.") + if ( + hasattr(args, "hidden_size") + and hasattr(args, "heads") + and args.hidden_size % args.heads != 0 + ): + raise ValueError("--hidden-size must be divisible by --heads.") + if ( + hasattr(args, "memory_size") + and hasattr(args, "memory_heads") + and args.memory_size % args.memory_heads != 0 + ): + raise ValueError( + "--memory-size must be divisible by --memory-heads." + ) + if ( + hasattr(args, "heads") + and hasattr(args, "kv_heads") + and args.heads % args.kv_heads != 0 + ): + raise ValueError("--heads must be divisible by --kv-heads.") + if hasattr(args, "attention_every") and args.attention_every <= 0: + raise ValueError("--attention-every must be positive.") + if hasattr(args, "ffn_latent_size") and ( + args.ffn_latent_size <= 0 + or args.ffn_latent_size > args.hidden_size + ): + raise ValueError( + "--ffn-latent-size must be positive and no larger than hidden size." + ) + if hasattr(args, "mtp_loss_weight") and args.mtp_loss_weight < 0: + raise ValueError("--mtp-loss-weight must be non-negative.") + if hasattr(args, "nan_lr_factor") and not ( + 0.0 < args.nan_lr_factor < 1.0 + ): + raise ValueError("--nan-lr-factor must be in (0, 1).") + if hasattr(args, "min_learning_rate") and args.min_learning_rate <= 0: + raise ValueError("--min-learning-rate must be positive.") + if hasattr(args, "max_nan_recoveries") and args.max_nan_recoveries < 0: + raise ValueError("--max-nan-recoveries must be non-negative.") + if hasattr(args, "finite_check_every") and args.finite_check_every < 0: + raise ValueError("--finite-check-every must be non-negative.") + + +def main() -> int: + parser = build_parser() + args = parser.parse_args() + validate_args(args) + args.function(args) + return 0 + + +if __name__ == "__main__": + try: + raise SystemExit(main()) + except KeyboardInterrupt: + print("\nInterrupted.", file=sys.stderr) + raise SystemExit(130)