OpenTransformer/AGILLM-4.3-ZeroGPU-GUI
0
1#!/usr/bin/env python3
2"""AGILLM4.1 mainline single-file trainer/inference runtime.
3
4AGILLM4.1 is the promoted AGILLM4 mainline evolved from the AGILLM3.5
5prototype, and it is larger than AGILLM3/AGILLM3.5. Resumed checkpoints are
6the source of truth for the exact architecture, with AGILLM4 presets available
7for fresh starts. This file is mechanically folded from AGILLM4 plus
8compatibility patches:
9- DeepSeek-V4-Pro tokenizer/checkpoint support by default
10- DeepSeek-V3.2 legacy compatibility support through the agillm35 shim
11- AR + SAT checkpoint schema compatibility; NAT can be disabled with --agillm3_compat
12- DiffusionBlock training support and optional async side-update ingestion
13"""
14from __future__ import annotations
15
16# Single-file module alias: helper code still imports the historical module names.
17import sys as _agillm41_sys
18_agillm41_sys.modules.setdefault("nB300_agillm4", _agillm41_sys.modules[__name__])
19_agillm41_sys.modules.setdefault("agillm35", _agillm41_sys.modules[__name__])
20_agillm41_sys.modules.setdefault("agillm41", _agillm41_sys.modules[__name__])
21_agillm41_sys.modules.setdefault("dblocks_train", _agillm41_sys.modules[__name__])
22_agillm41_sys.modules.setdefault("fused_ce", _agillm41_sys.modules[__name__])
23_agillm41_sys.modules.setdefault("anchor_memory", _agillm41_sys.modules[__name__])
24
25import types as _agillm41_types
26
27# ===== BEGIN agillm_checkpoint_provenance.py (folded) =====
28_AGILLM_CHECKPOINT_PROVENANCE_SOURCE = '"""agillm_checkpoint_provenance.py — git-style lineage tracking for checkpoints.\n\nEvery full checkpoint (.pt) carries a `provenance` dict that records:\n - warmstart source & its provenance (chained like git commits)\n - training step, tokens seen, loss (total + per-head)\n - training script name + SHA256, full argv\n - creation time, hostname, PID, GPU metrics\n - inference samples (3 short generations from the model)\n - dataset provenance snapshot\n\nCLI usage:\n python3 agillm_checkpoint_provenance.py show <checkpoint.pt>\n python3 agillm_checkpoint_provenance.py lineage <checkpoint.pt>\n python3 agillm_checkpoint_provenance.py compare <ckpt_a.pt> <ckpt_b.pt>\n"""\n\nfrom __future__ import annotations\n\nimport argparse\nimport hashlib\nimport json\nimport os\nimport platform\nimport re\nimport subprocess\nimport sys\nimport time\nimport pathlib\nfrom typing import Any, Dict, List, Optional, Tuple\n\n# ---------------------------------------------------------------------------\n# Schema key\n# ---------------------------------------------------------------------------\nPROVENANCE_KEY = "agillm43_provenance"\nPROVENANCE_SCHEMA_VERSION = 1\n\n# ---------------------------------------------------------------------------\n# Provenance dict shape\n# ---------------------------------------------------------------------------\n"""\nprovenance = {\n "schema_version": 1,\n "checkpoint_type": "full" | "delta",\n\n # Identity\n "created_at_iso": "2026-06-23T03:14:00Z",\n "created_at_unix": 1750000000.0,\n "hostname": "agillm43-boxa",\n "pid": 1372905,\n "lane": "a0",\n\n # Training state\n "step": 13886,\n "seen_tok": 850000000,\n "loss": 2.345,\n "loss_ar": 2.1,\n "loss_sat": 0.15,\n "loss_nat": 0.095,\n "batch_size": 56,\n "block_size": 1536,\n\n # Source\n "train_script": "agillm41.py",\n "train_script_sha256": "abc123...",\n "train_argv": "--warmstart_from /workspace/... --preset agillm4_floor ...",\n\n # Warmstart chain (like git parent)\n "warmstart_source_path": "/workspace/agillm4_v100_master_ckpts/pretrain_step02182564.pt",\n "warmstart_source_provenance": { ... } or None,\n\n # Config snapshot\n "cfg_keys": ["dmodel", "layers", "heads", ...],\n\n # Inference samples (3 short generations)\n "inference_samples": [\n {"prompt": "The meaning of life is", "generation": " to find", "tokens": 5},\n ...\n ],\n\n # GPU state at save time\n "gpu": {\n "allocated_gb": 30.5,\n "reserved_gb": 31.2,\n "peak_allocated_gb": 32.0,\n },\n\n # Dataset provenance fragment\n "dataset_provenance": { ... },\n\n # Tokenizer info\n "tokenizer_id": "...",\n}\n"""\n\n\n# ---------------------------------------------------------------------------\n# Utilities\n# ---------------------------------------------------------------------------\n\ndef _iso_now() -> str:\n return time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())\n\n\ndef _sha256_file(path: pathlib.Path) -> str:\n h = hashlib.sha256()\n with open(path, "rb") as f:\n while True:\n chunk = f.read(1 << 20)\n if not chunk:\n break\n h.update(chunk)\n return h.hexdigest()\n\n\ndef _sha256_bytes(data: bytes) -> str:\n return hashlib.sha256(data).hexdigest()\n\n\ndef _gpu_metrics() -> dict:\n """Collect GPU memory usage if CUDA is available."""\n try:\n import torch\n if not torch.cuda.is_available():\n return {}\n return {\n "allocated_gb": round(torch.cuda.memory_allocated() / (1024**3), 2),\n "reserved_gb": round(torch.cuda.memory_reserved() / (1024**3), 2),\n "peak_allocated_gb": round(torch.cuda.max_memory_allocated() / (1024**3), 2),\n }\n except Exception:\n return {}\n\n\ndef _script_sha256() -> Tuple[str, str]:\n """SHA256 of the running training script. Returns (basename, hexdigest)."""\n try:\n main = sys.modules.get("__main__")\n if main and hasattr(main, "__file__") and main.__file__:\n p = pathlib.Path(main.__file__).resolve()\n return p.name, _sha256_file(p)\n except Exception:\n pass\n return ("", "")\n\n\ndef _script_argv() -> str:\n return " ".join(sys.argv)\n\n\ndef _read_proc_cmdline(pid: str = "self") -> str:\n try:\n raw = pathlib.Path("/proc") / str(pid) / "cmdline"\n data = raw.read_bytes()\n return " ".join(part.decode("utf-8", "replace") for part in data.split(b"\\0") if part)\n except Exception:\n return ""\n\n\ndef _safe_env_snapshot() -> dict:\n """Capture useful launch env without leaking tokens or credentials."""\n prefixes = ("AGILLM", "CUDA_", "HF_HUB_", "HF_DATASETS_", "PYTORCH_", "OMP_", "MKL_")\n allow = {\n "CUDA_VISIBLE_DEVICES",\n "HF_HUB_DISABLE_XET",\n "HF_DATASETS_TRUST_REMOTE_CODE",\n "PYTORCH_CUDA_ALLOC_CONF",\n "OMP_NUM_THREADS",\n "MKL_NUM_THREADS",\n }\n secret_fragments = ("TOKEN", "SECRET", "PASSWORD", "PASSWD", "KEY", "CREDENTIAL", "AUTH", "COOKIE")\n out = {}\n for key, value in sorted(os.environ.items()):\n if not (key in allow or key.startswith(prefixes)):\n continue\n if any(fragment in key.upper() for fragment in secret_fragments):\n out[key] = "<redacted>"\n else:\n out[key] = str(value)[:2048]\n return out\n\n\ndef _redact_text(text: str) -> str:\n secret_fragments = ("TOKEN", "SECRET", "PASSWORD", "PASSWD", "API_KEY", "AUTH", "COOKIE")\n lines = []\n for line in str(text).splitlines()[:240]:\n upper = line.upper()\n if any(fragment in upper for fragment in secret_fragments):\n lines.append("<redacted secret-bearing line>")\n else:\n lines.append(line[:4096])\n return "\\n".join(lines)\n\n\ndef _launch_metadata() -> dict:\n meta = {\n "schema": "agillm.launch.v1",\n "argv": list(sys.argv),\n "argv_string": _script_argv(),\n "cwd": "",\n "pid": os.getpid(),\n "ppid": os.getppid(),\n "proc_cmdline": _read_proc_cmdline("self"),\n "parent_proc_cmdline": _read_proc_cmdline(str(os.getppid())),\n "env": _safe_env_snapshot(),\n }\n try:\n meta["cwd"] = str(pathlib.Path.cwd())\n except Exception:\n pass\n launch_script = os.environ.get("AGILLM43_LAUNCH_SCRIPT") or os.environ.get("AGILLM_LAUNCH_SCRIPT") or ""\n launch_command = os.environ.get("AGILLM43_LAUNCH_COMMAND") or os.environ.get("AGILLM_LAUNCH_COMMAND") or ""\n if launch_command:\n meta["launch_command"] = _redact_text(launch_command)\n if launch_script:\n sp = pathlib.Path(launch_script)\n info = {"path": str(sp)}\n try:\n if sp.exists() and sp.is_file():\n info["size_bytes"] = sp.stat().st_size\n info["sha256"] = _sha256_file(sp)\n info["preview_redacted"] = _redact_text(sp.read_text(errors="replace"))\n except Exception as exc:\n info["error"] = str(exc)\n meta["launch_script"] = info\n return meta\n\n\ndef _infer_samples(core, ar_h, sat_h, tok, device: str, prompt_texts: List[str],\n max_new: int = 32, temperature: float = 0.5, top_k: int = 20) -> List[dict]:\n """Generate a few short inference samples from the model.\n\n This is called at save time with gradients off (torch.no_grad).\n If anything fails, returns an empty list — never crashes a save.\n """\n samples = []\n try:\n import torch\n core.eval()\n ar_h.eval()\n if sat_h is not None:\n sat_h.eval()\n\n for prompt in prompt_texts:\n try:\n input_ids = tok.encode(prompt, return_tensors="pt").to(device)\n if input_ids.numel() == 0:\n continue\n generated = input_ids.clone()\n for _ in range(max_new):\n with torch.no_grad():\n h = core(generated, None)\n logits = ar_h(h[:, -1:])\n probs = torch.softmax(logits[:, -1] / max(temperature, 1e-8), dim=-1)\n if top_k > 0:\n vals, idxs = torch.topk(probs, min(top_k, probs.size(-1)))\n probs = torch.zeros_like(probs).scatter_(-1, idxs, vals)\n next_id = torch.multinomial(probs, 1)\n generated = torch.cat([generated, next_id], dim=1)\n if next_id.item() == 0: # EOS\n break\n text = tok.decode(generated[0].tolist(), skip_special_tokens=True)\n new_tokens = generated.size(1) - input_ids.size(1)\n samples.append({\n "prompt": prompt,\n "generation": text[len(prompt):] if text.startswith(prompt) else text,\n "tokens": new_tokens,\n })\n except Exception:\n samples.append({"prompt": prompt, "generation": "", "tokens": 0})\n except Exception:\n pass\n return samples\n\n\n# ---------------------------------------------------------------------------\n# Core provenance construction\n# ---------------------------------------------------------------------------\n\ndef _step_from_text(text: Optional[str]) -> Optional[int]:\n m = re.search(r"step(\\d+)", str(text or ""))\n return int(m.group(1)) if m else None\n\n\ndef _origin_step_from_provenance(prov: Optional[dict]) -> int:\n if not isinstance(prov, dict):\n return 0\n for key in ("global_origin_step", "warmstart_base_step"):\n try:\n value = int(prov.get(key) or 0)\n except Exception:\n value = 0\n if value > 0:\n return value\n parent = prov.get("warmstart_source_path") or prov.get("source_path") or ""\n parent_step = _step_from_text(parent)\n if parent_step and parent_step > 0: # AGILLM-LINEAGE-FIX 20260702\n return int(parent_step)\n return 0\n\n\ndef _origin_seen_tok_from_provenance(prov: Optional[dict]) -> int:\n if not isinstance(prov, dict):\n return 0\n for key in ("global_origin_seen_tok", "warmstart_base_seen_tok"):\n try:\n value = int(prov.get(key) or 0)\n except Exception:\n value = 0\n if value > 0:\n return value\n return 0\n\n\ndef collect(args, *, step: int, seen_tok: int, loss: float,\n loss_ar: Optional[float] = None, loss_sat: Optional[float] = None,\n loss_nat: Optional[float] = None,\n batch_size: int = 0, block_size: int = 0,\n warmstart_source_path: Optional[str] = None,\n warmstart_source_provenance: Optional[dict] = None,\n dataset_provenance: Optional[dict] = None,\n lane: str = "",\n inference_samples: Optional[list] = None,\n checkpoint_type: str = "full",\n _sample_core=None, _sample_ar=None, _sample_sat=None,\n _sample_tok=None, _sample_device: str = "",\n _sample_prompts: Optional[List[str]] = None) -> dict:\n """Build a provenance dict to embed in the checkpoint."""\n\n script_name, script_sha = _script_sha256()\n\n prov: dict = {\n "schema_version": PROVENANCE_SCHEMA_VERSION,\n "checkpoint_type": checkpoint_type,\n "created_at_iso": _iso_now(),\n "created_at_unix": time.time(),\n "hostname": platform.node(),\n "pid": os.getpid(),\n "lane": lane or "",\n "step": int(step),\n "seen_tok": int(seen_tok),\n "loss": float(loss),\n "batch_size": int(batch_size),\n "block_size": int(block_size),\n "train_script": script_name,\n "train_argv": _script_argv(),\n "launch": _launch_metadata(),\n "gpu": _gpu_metrics(),\n }\n\n if script_sha:\n prov["train_script_sha256"] = script_sha\n\n if loss_ar is not None:\n prov["loss_ar"] = float(loss_ar)\n if loss_sat is not None:\n prov["loss_sat"] = float(loss_sat)\n if loss_nat is not None:\n prov["loss_nat"] = float(loss_nat)\n\n source_step = _step_from_text(warmstart_source_path)\n origin_step = _origin_step_from_provenance(warmstart_source_provenance)\n origin_seen_tok = _origin_seen_tok_from_provenance(warmstart_source_provenance)\n if not origin_step and source_step and source_step > 0: # AGILLM-LINEAGE-FIX 20260702\n origin_step = int(source_step)\n\n prov["local_step"] = int(step)\n if source_step is not None:\n prov["warmstart_source_step"] = int(source_step)\n prov["global_origin_step"] = int(origin_step or 0)\n prov["warmstart_base_step"] = int(origin_step or 0)\n prov["effective_global_step"] = int((origin_step + int(step)) if origin_step else int(step))\n prov["global_origin_seen_tok"] = int(origin_seen_tok or 0)\n prov["warmstart_base_seen_tok"] = int(origin_seen_tok or 0)\n prov["effective_seen_tok"] = int(int(origin_seen_tok or 0) + int(seen_tok))\n\n if warmstart_source_path:\n prov["warmstart_source_path"] = str(warmstart_source_path)\n if warmstart_source_provenance:\n prov["warmstart_source_provenance"] = warmstart_source_provenance\n\n if dataset_provenance:\n prov["dataset_provenance"] = dataset_provenance\n\n if inference_samples is not None:\n prov["inference_samples"] = inference_samples\n elif _sample_core is not None and _sample_ar is not None and _sample_tok is not None:\n try:\n prompts = _sample_prompts or ["The meaning of", "def hello():", "2 + 2 ="]\n prov["inference_samples"] = _infer_samples(\n _sample_core, _sample_ar, _sample_sat,\n _sample_tok, _sample_device or "cpu", prompts, max_new=12)\n except Exception:\n prov["inference_samples"] = []\n\n return prov\n\n\ndef embed(state_dict: dict, provenance: dict) -> dict:\n """Embed provenance into the checkpoint state dict (mutates + returns)."""\n state_dict[PROVENANCE_KEY] = provenance\n return state_dict\n\n\n# ---------------------------------------------------------------------------\n# Extraction (lightweight — only reads provenance from .pt wrapper)\n# ---------------------------------------------------------------------------\n\ndef extract(path: pathlib.Path) -> Optional[dict]:\n """Extract the provenance dict from a saved .pt checkpoint.\n\n This reads only the top-level wrapper, not the full model weights.\n For zstd-wrapped checkpoints, it only decompresses enough to find the\n provenance key.\n\n Returns None if no provenance is found.\n """\n try:\n import torch\n # The checkpoint may be zstd-wrapped. Load the wrapper first.\n wrapper = torch.load(str(path), map_location="cpu", weights_only=False)\n if not isinstance(wrapper, dict):\n return None\n\n # If zstd-wrapped, decompress and get inner dict\n inner = wrapper\n if wrapper.get("__agillm43_payload_codec__") == "agillm43_zstd_torch_v1":\n import zstandard as zstd\n raw = zstd.ZstdDecompressor().decompress(bytes(wrapper["payload"].tolist()))\n import io\n inner = torch.load(io.BytesIO(raw), map_location="cpu", weights_only=False)\n\n if not isinstance(inner, dict):\n return None\n\n provenance = inner.get(PROVENANCE_KEY)\n if provenance is not None:\n return provenance\n\n # Fallback: check for sidecar\n sidecar = path.with_suffix(".provenance.json")\n if sidecar.exists():\n return json.loads(sidecar.read_text())\n\n return None\n except Exception:\n return None\n\n\ndef extract_provenance_sidecar(ckpt_path: pathlib.Path) -> Optional[dict]:\n """Read the .provenance.json sidecar without touching the .pt at all."""\n sidecar = ckpt_path.with_suffix(".provenance.json")\n if sidecar.exists():\n try:\n return json.loads(sidecar.read_text())\n except Exception:\n pass\n return None\n\n\ndef write_sidecar(ckpt_path: pathlib.Path, provenance: dict) -> None:\n """Write .provenance.json sidecar beside the checkpoint."""\n sidecar = ckpt_path.with_suffix(".provenance.json")\n tmp = sidecar.with_suffix(".provenance.json.tmp")\n try:\n tmp.write_text(json.dumps(provenance, indent=2, sort_keys=True) + "\\n")\n tmp.replace(sidecar)\n except Exception as exc:\n print(f"[provenance] WARNING: failed to write sidecar {sidecar}: {exc}")\n\n\n# ---------------------------------------------------------------------------\n# Display / CLI\n# ---------------------------------------------------------------------------\n\ndef format_provenance(prov: dict, indent: int = 0) -> str:\n """Format a provenance dict as a readable block."""\n pad = " " * indent\n lines = [f"{pad}┌── Checkpoint Provenance ──"]\n if not prov:\n return f"{pad}└── (no provenance)"\n\n def kv(k, v, default="—"):\n val = v if v is not None else default\n return f"{pad} {k}: {val}"\n\n lines.append(kv("Schema version", prov.get("schema_version")))\n lines.append(kv("Type", prov.get("checkpoint_type")))\n lines.append(kv("Step", prov.get("step")))\n lines.append(kv("Tokens seen", f"{prov.get(\'seen_tok\', 0):,}"))\n lines.append(kv("Loss", prov.get("loss")))\n if prov.get("loss_ar") is not None:\n lines.append(kv(" ├ AR loss", prov["loss_ar"]))\n if prov.get("loss_sat") is not None:\n lines.append(kv(" ├ SAT loss", prov["loss_sat"]))\n if prov.get("loss_nat") is not None:\n lines.append(kv(" └ NAT loss", prov["loss_nat"]))\n lines.append(kv("Batch / Block", f"{prov.get(\'batch_size\')} / {prov.get(\'block_size\')}"))\n lines.append(kv("Created (ISO)", prov.get("created_at_iso")))\n lines.append(kv("Hostname", prov.get("hostname")))\n lines.append(kv("PID", prov.get("pid")))\n lines.append(kv("Lane", prov.get("lane", "—")))\n lines.append(kv("Train script", prov.get("train_script")))\n if prov.get("train_script_sha256"):\n lines.append(kv(" └ SHA256", prov["train_script_sha256"][:16] + "..."))\n gpu = prov.get("gpu", {})\n if gpu:\n lines.append(kv("GPU alloc/resrv/peak",\n f"{gpu.get(\'allocated_gb\', \'?\')}G / {gpu.get(\'reserved_gb\', \'?\')}G / {gpu.get(\'peak_allocated_gb\', \'?\')}G"))\n\n ws = prov.get("warmstart_source_path")\n if ws:\n lines.append(kv("Warmstart source", ws))\n wprov = prov.get("warmstart_source_provenance")\n if wprov:\n lines.append(f"{pad} └ step={wprov.get(\'step\', \'?\')} loss={wprov.get(\'loss\', \'?\')}")\n\n samples = prov.get("inference_samples", [])\n if samples:\n lines.append(f"{pad}Inference samples ({len(samples)}):")\n for i, s in enumerate(samples):\n gen = s.get("generation", "")\n if len(gen) > 60:\n gen = gen[:60] + "..."\n lines.append(f"{pad} [{i}] prompt={s.get(\'prompt\',\'\')!r}")\n lines.append(f"{pad} → {gen!r} ({s.get(\'tokens\', 0)} tokens)")\n\n lines.append(f"{pad}└──")\n return "\\n".join(lines)\n\n\ndef show_lineage(path: pathlib.Path, max_depth: int = 32) -> List[dict]:\n """Walk the provenance chain (like git log) and return ordered list [oldest..newest]."""\n chain: List[dict] = []\n seen = set()\n current = path.resolve() if path.exists() else path\n\n for _ in range(max_depth):\n prov = extract(current)\n if prov is None:\n break\n\n key = str(current)\n if key in seen:\n break\n seen.add(key)\n\n entry = prov.copy()\n entry["_checkpoint_path"] = str(current)\n chain.append(entry)\n\n # Walk to warmstart parent\n ws = prov.get("warmstart_source_path")\n if not ws:\n break\n wprov = prov.get("warmstart_source_provenance")\n if not wprov:\n break\n current = pathlib.Path(ws)\n # Avoid infinite loop if parent points to itself\n if str(current) == key:\n break\n else:\n chain.append({"_checkpoint_path": f"(truncated at {max_depth} hops)"})\n\n chain.reverse() # oldest first\n return chain\n\n\ndef format_lineage(chain: List[dict]) -> str:\n """Format a lineage chain as a readable tree."""\n lines = ["Checkpoint Lineage (oldest → newest):", ""]\n for i, entry in enumerate(chain):\n path = entry.get("_checkpoint_path", "?")\n step = entry.get("step", "?")\n loss = entry.get("loss", "?")\n iso = entry.get("created_at_iso", "?")\n ws = entry.get("warmstart_source_path", "")\n marker = "●" if i == len(chain) - 1 else "│" if i < len(chain) - 1 else "○"\n lines.append(f" {marker} step={step} loss={loss} {iso}")\n lines.append(f" │ {path}")\n if ws and i < len(chain) - 1:\n lines.append(f" │ warmstart ← {pathlib.Path(ws).name}")\n lines.append("")\n return "\\n".join(lines)\n\n\n# ---------------------------------------------------------------------------\n# CLI\n# ---------------------------------------------------------------------------\n\ndef _cmd_show(args_cli):\n path = pathlib.Path(args_cli.checkpoint)\n if not path.exists():\n print(f"ERROR: {path} not found")\n sys.exit(1)\n prov = extract(path)\n if prov is None:\n prov = extract_provenance_sidecar(path)\n if prov is None:\n print(f"No provenance found in {path}")\n sys.exit(1)\n print(format_provenance(prov))\n if args_cli.verbose:\n print("\\nFull provenance JSON:")\n print(json.dumps(prov, indent=2, sort_keys=True))\n\n\ndef _cmd_lineage(args_cli):\n path = pathlib.Path(args_cli.checkpoint)\n if not path.exists():\n print(f"ERROR: {path} not found")\n sys.exit(1)\n chain = show_lineage(path, max_depth=args_cli.max_depth)\n print(format_lineage(chain))\n\n\ndef _cmd_compare(args_cli):\n a = pathlib.Path(args_cli.checkpoint_a)\n b = pathlib.Path(args_cli.checkpoint_b)\n for p, label in [(a, "A"), (b, "B")]:\n if not p.exists():\n print(f"ERROR: {label}={p} not found")\n sys.exit(1)\n\n pa = extract(a) or {}\n pb = extract(b) or {}\n\n def safe(key, d, default="—"):\n return d.get(key, default)\n\n print(f"Compare: {a.name} vs {b.name}")\n print()\n keys = ["step", "seen_tok", "loss", "loss_ar", "loss_sat", "loss_nat",\n "batch_size", "block_size", "created_at_iso", "hostname", "lane"]\n for k in keys:\n va = safe(k, pa)\n vb = safe(k, pb)\n changed = " ←" if str(va) != str(vb) else ""\n print(f" {k:20s} {str(va):>20s} {str(vb):>20s}{changed}")\n\n sa = pa.get("inference_samples", [])\n sb = pb.get("inference_samples", [])\n if sa or sb:\n print()\n print(f" Inference samples: A={len(sa)} B={len(sb)}")\n\n\ndef main():\n parser = argparse.ArgumentParser(\n description="agillm checkpoint provenance — git for checkpoints",\n formatter_class=argparse.RawDescriptionHelpFormatter,\n epilog=__doc__,\n )\n sub = parser.add_subparsers(dest="command")\n\n p_show = sub.add_parser("show", help="Show provenance for a checkpoint")\n p_show.add_argument("checkpoint", type=str, help="Path to .pt checkpoint")\n p_show.add_argument("-v", "--verbose", action="store_true", help="Also dump full JSON")\n\n p_lineage = sub.add_parser("lineage", help="Show full warmstart chain (git log)")\n p_lineage.add_argument("checkpoint", type=str, help="Path to .pt checkpoint")\n p_lineage.add_argument("--max-depth", type=int, default=32, help="Max hops to follow")\n\n p_cmp = sub.add_parser("compare", help="Compare two checkpoints")\n p_cmp.add_argument("checkpoint_a", type=str)\n p_cmp.add_argument("checkpoint_b", type=str)\n\n args_cli = parser.parse_args()\n if args_cli.command == "show":\n _cmd_show(args_cli)\n elif args_cli.command == "lineage":\n _cmd_lineage(args_cli)\n elif args_cli.command == "compare":\n _cmd_compare(args_cli)\n else:\n parser.print_help()\n sys.exit(1)\n\n\nif __name__ == "__main__":\n main()\n'
29_agillm_provenance = _agillm41_types.ModuleType("agillm_checkpoint_provenance")
30_agillm_provenance.__file__ = __file__ + "#agillm_checkpoint_provenance"
31exec(compile(_AGILLM_CHECKPOINT_PROVENANCE_SOURCE, _agillm_provenance.__file__, "exec"), _agillm_provenance.__dict__)
32_agillm41_sys.modules.setdefault("agillm_checkpoint_provenance", _agillm_provenance)
33# ===== END agillm_checkpoint_provenance.py (folded) =====
34
35
36# ===== BEGIN anchor_memory.py =====
37#!/usr/bin/env python3
38
39from dataclasses import dataclass
40
41import torch
42import torch.nn as nn
43import torch.nn.functional as F
44
45
46@dataclass
47class AnchorMemoryConfig:
48 d_model: int
49 heads: int
50 anchor_stride: int = 256
51 max_anchors: int = 2048
52 dropout: float = 0.0
53
54
55class AnchorCompressor(nn.Module):
56 """Compress local token spans into trainable anchor vectors."""
57
58 def __init__(self, d_model: int, anchor_stride: int):
59 super().__init__()
60 self.anchor_stride = anchor_stride
61 self.score = nn.Linear(d_model, 1)
62 self.mix = nn.Sequential(
63 nn.LayerNorm(d_model),
64 nn.Linear(d_model, 4 * d_model),
65 nn.GELU(),
66 nn.Linear(4 * d_model, d_model),
67 )
68
69 def forward(self, x: torch.Tensor) -> torch.Tensor:
70 bsz, seq, dim = x.shape
71 pad = (-seq) % self.anchor_stride
72 if pad:
73 x = F.pad(x, (0, 0, 0, pad))
74 chunks = x.view(bsz, -1, self.anchor_stride, dim)
75 weights = self.score(chunks).softmax(dim=2)
76 pooled = (chunks * weights).sum(dim=2)
77 return pooled + self.mix(pooled)
78
79
80class AnchorMemoryLayer(nn.Module):
81 """Local-token stream reads from a bounded bank of learned anchors."""
82
83 def __init__(self, cfg: AnchorMemoryConfig):
84 super().__init__()
85 self.cfg = cfg
86 self.compress = AnchorCompressor(cfg.d_model, cfg.anchor_stride)
87 self.q_ln = nn.LayerNorm(cfg.d_model)
88 self.mem_ln = nn.LayerNorm(cfg.d_model)
89 self.read = nn.MultiheadAttention(
90 cfg.d_model,
91 cfg.heads,
92 dropout=cfg.dropout,
93 batch_first=True,
94 )
95 self.gate = nn.Sequential(nn.Linear(2 * cfg.d_model, cfg.d_model), nn.Sigmoid())
96 self.out_ln = nn.LayerNorm(cfg.d_model)
97
98 def forward(
99 self,
100 x: torch.Tensor,
101 memory: torch.Tensor | None = None,
102 *,
103 detach_memory: bool = False,
104 ) -> tuple[torch.Tensor, torch.Tensor]:
105 new_anchors = self.compress(x)
106 if detach_memory:
107 new_anchors = new_anchors.detach()
108 if memory is None:
109 bank = new_anchors
110 else:
111 bank = torch.cat([memory, new_anchors], dim=1)
112 if bank.size(1) > self.cfg.max_anchors:
113 bank = bank[:, -self.cfg.max_anchors :]
114
115 recalled, _ = self.read(self.q_ln(x), self.mem_ln(bank), self.mem_ln(bank), need_weights=False)
116 gate = self.gate(torch.cat([x, recalled], dim=-1))
117 mixed = x + gate * recalled
118 return self.out_ln(mixed), bank
119
120
121def smoke_test() -> None:
122 cfg = AnchorMemoryConfig(d_model=128, heads=8, anchor_stride=32, max_anchors=64)
123 layer = AnchorMemoryLayer(cfg)
124 x = torch.randn(2, 256, 128)
125 y, memory = layer(x)
126 assert y.shape == x.shape
127 assert memory.shape == (2, 8, 128)
128 y2, memory2 = layer(x, memory)
129 assert y2.shape == x.shape
130 assert memory2.shape == (2, 16, 128)
131 print("anchor_memory smoke OK", y.shape, memory2.shape)
132
133
134
135# ===== END anchor_memory.py =====
136
137
138# ===== BEGIN fused_ce.py =====
139"""Fused cross-entropy: streams over the VOCAB dimension (online-softmax) so the
140[N x V] logit matrix is NEVER materialized -- only [N x vchunk]. Custom backward
141recomputes softmax per vocab-chunk (grad = softmax - onehot). This is the
142DiffusionBlocks 'process in chunks, don't hold the whole thing' idea applied to
143the output head instead of network depth."""
144import torch
145
146class FusedCE(torch.autograd.Function):
147 @staticmethod
148 def forward(ctx, h, W, tgt, vchunk=16384):
149 with torch.cuda.amp.autocast(enabled=True):
150 hf = h.float()
151 Wf = W.float()
152 N, d = h.shape
153 V = W.shape[0]
154 m = torch.full((N,), -1e30, device=h.device, dtype=torch.float32)
155 s = torch.zeros(N, device=h.device, dtype=torch.float32)
156 zt = torch.zeros(N, device=h.device, dtype=torch.float32)
157 for c in range(0, V, vchunk):
158 lg = hf @ Wf[c:c+vchunk].T # [N,vchunk] transient only
159 cm = lg.max(1).values
160 nm = torch.maximum(m, cm)
161 s = s * torch.exp(m - nm) + torch.exp(lg - nm[:, None]).sum(1)
162 m = nm
163 ic = (tgt >= c) & (tgt < c+vchunk)
164 if ic.any():
165 zt[ic] = lg[ic, tgt[ic] - c].float()
166 lse = m + torch.log(s)
167 ctx.save_for_backward(h, W, tgt, lse)
168 ctx.vchunk = vchunk
169 return (lse - zt).mean()
170
171 @staticmethod
172 def backward(ctx, go):
173 h, W, tgt, lse = ctx.saved_tensors
174 vc = ctx.vchunk
175 N, d = h.shape
176 V = W.shape[0]
177 with torch.cuda.amp.autocast(enabled=True):
178 hf = h.float()
179 Wc_all = W.float()
180 gh = torch.zeros_like(hf)
181 gW = torch.zeros(W.shape, device=W.device, dtype=torch.float32)
182 sc = float(go) / N
183 for c in range(0, V, vc):
184 Wc = Wc_all[c:c+vc]
185 p = torch.exp(hf @ Wc.T - lse[:, None]) # softmax chunk [N,vchunk]
186 ic = (tgt >= c) & (tgt < c+vc)
187 if ic.any():
188 p[ic, tgt[ic] - c] -= 1.0
189 p *= sc
190 gh += p @ Wc
191 gW[c:c+vc] += p.T @ hf
192 return gh.to(h.dtype), gW.to(W.dtype), None, None
193
194def fused_ce(h, W, tgt, vchunk=16384):
195 return FusedCE.apply(h.reshape(-1, h.size(-1)), W, tgt.reshape(-1), vchunk)
196
197# ===== END fused_ce.py =====
198
199
200# ===== BEGIN dblocks_train.py =====
201"""DiffusionBlocks training mode folded into AGILLM-4 (gated by --dblock).
202
203Block-wise EDM denoising on the real Encoder blocks, supervising AR + SAT(fixed+var)
204+ NAT each step on ONE block, with grad-checkpointed layers and fused vocab-streaming
205CE. Reuses the live data stream / optimizer / checkpointing of nB300_agillm4.
206Lazy-imports nB300 inside functions to avoid a circular import.
207"""
208import math
209import random
210import time
211from collections import defaultdict
212import numpy as np
213import torch
214import torch.nn as nn
215import torch.nn.functional as F
216import torch.utils.checkpoint as _ck
217
218# Optional CuPy hook for future AGILLM agents.
219# Keep the main trainer on PyTorch CUDA: autograd, AMP, SDPA, MoE, and DBlock
220# losses are already torch-native. This helper is deliberately lazy and disabled
221# by default so importing the trainer never depends on CuPy or CUDA toolkit
222# headers. Use it only for side/offline NumPy-heavy, non-autograd helpers such as
223# checkpoint/delta diagnostics, custom array probes, or preprocessing experiments.
224_CUPY_DISABLED = object()
225_OPTIONAL_CUPY = _CUPY_DISABLED
226
227
228def _optional_cupy_backend(reason=""):
229 """Return cupy when AGILLM_ENABLE_CUPY=1, otherwise None.
230
231 CuPy is useful for large NumPy-style array work on CUDA/ROCm hosts, but it is
232 not a replacement for torch in the AGILLM4.3 training hot path. Callers must
233 keep data on the GPU and avoid CPU<->GPU ping-pong. On Vast CUDA images, CuPy
234 may need CUDA_PATH=/usr/local/cuda so elementwise kernels can find headers.
235 """
236 global _OPTIONAL_CUPY
237 import os as _os
238
239 if _os.environ.get("AGILLM_ENABLE_CUPY", "0") != "1":
240 return None
241 if _OPTIONAL_CUPY is _CUPY_DISABLED:
242 if not _os.environ.get("CUDA_PATH") and _os.path.exists("/usr/local/cuda"):
243 _os.environ["CUDA_PATH"] = "/usr/local/cuda"
244 try:
245 import cupy as _cp # type: ignore
246 _OPTIONAL_CUPY = _cp
247 label = f" for {reason}" if reason else ""
248 print(f"[cupy] optional backend enabled{label}: cupy={_cp.__version__}", flush=True)
249 except Exception as exc:
250 _OPTIONAL_CUPY = None
251 print(f"[cupy] optional backend unavailable: {type(exc).__name__}: {exc}", flush=True)
252 return _OPTIONAL_CUPY
253
254SD = 0.5
255
256
257
258
259def _profile_active(state, args):
260 limit = int(getattr(args, "profile_steps", 0) or 0)
261 return limit > 0 and int(state.get("profile_n", 0)) < limit
262
263
264def _profile_add(state, name, seconds):
265 if seconds is None:
266 return
267 prof = state.setdefault("profile_times", defaultdict(float))
268 prof[name] += float(seconds)
269
270
271def _profile_tic(enabled):
272 if not enabled:
273 return None
274 if torch.cuda.is_available():
275 torch.cuda.synchronize()
276 return time.perf_counter()
277
278
279def _profile_toc(state, name, start):
280 if start is None:
281 return
282 if torch.cuda.is_available():
283 torch.cuda.synchronize()
284 _profile_add(state, name, time.perf_counter() - start)
285
286
287def _profile_step_done(state, args):
288 limit = int(getattr(args, "profile_steps", 0) or 0)
289 if limit <= 0:
290 return
291 n_prev = int(state.get("profile_n", 0))
292 if n_prev >= limit:
293 return
294 state["profile_n"] = n_prev + 1
295 n = int(state["profile_n"])
296 log_every = max(1, int(getattr(args, "profile_log_every", 25) or 25))
297 if n % log_every != 0 and n != limit:
298 return
299 times = state.get("profile_times", {})
300 keys = [
301 "data_stream", "tensor", "setup",
302 "ar_forward", "ar_ce", "ar_backward",
303 "sat_forward", "sat_ce", "sat_backward",
304 "nat_forward", "nat_ce", "nat_backward",
305 "opt_step", "step_total",
306 ]
307 parts = []
308 for key in keys:
309 val = float(times.get(key, 0.0)) * 1000.0 / max(1, n)
310 if val > 0.01:
311 parts.append(f"{key}={val:.2f}ms")
312 print(f"[profile] n={n}/{limit} avg " + " ".join(parts), flush=True)
313
314def _cdf(x):
315 return 0.5 * (1 + math.erf(x / math.sqrt(2)))
316
317
318def _ppf(p):
319 return float(torch.erfinv(torch.tensor(2 * p - 1.0)) * math.sqrt(2))
320
321
322def _dblock_sigma_config(args=None):
323 smin = float(getattr(args, "dblock_sigma_min", 0.002) if args is not None else 0.002)
324 smax = float(getattr(args, "dblock_sigma_max", 80.0) if args is not None else 80.0)
325 pm = float(getattr(args, "dblock_sigma_pmean", -1.2) if args is not None else -1.2)
326 ps = float(getattr(args, "dblock_sigma_pstd", 1.2) if args is not None else 1.2)
327 smin = max(smin, 1e-6)
328 smax = max(smax, smin * 1.0001)
329 ps = max(ps, 1e-6)
330 return smin, smax, pm, ps
331
332
333def _block_sigmas(B, smin=0.002, smax=80.0, pm=-1.2, ps=1.2):
334 smin = max(float(smin), 1e-6)
335 smax = max(float(smax), smin * 1.0001)
336 ps = max(float(ps), 1e-6)
337 a, b = _cdf((math.log(smin) - pm) / ps), _cdf((math.log(smax) - pm) / ps)
338 return [float(np.exp(pm + ps * _ppf(a + (b - a) * (i / B)))) for i in range(B + 1)]
339
340
341def _edm_pre(s):
342 s = s[:, None, None]
343 return SD**2 / (s**2 + SD**2), s * SD / (s**2 + SD**2) ** 0.5, 1 / (s**2 + SD**2) ** 0.5
344
345
346def _edm_w(s, wmax=5.0):
347 return float(((s**2 + SD**2) / (s * SD) ** 2).clamp(max=wmax).mean())
348
349
350_DBLOCK_ROUTER_EVENT_FEATURES = 10
351_DBLOCK_ROUTER_HISTORY = 32
352
353
354class _DblockLearnedRouter(nn.Module):
355 # Transformer DBlock router conditioned on the network's running representation
356 # plus a bounded route/outcome memory. Sequence = [CTX] + B block tokens + H
357 # recent outcome tokens, so routing can learn from what the model is seeing now
358 # and what the previous routing choices actually did to loss.
359 def __init__(self, ctx_dim, d_model=64, heads=4, layers=2, feat_dim=6, n_blocks_max=64, history=_DBLOCK_ROUTER_HISTORY, event_dim=_DBLOCK_ROUTER_EVENT_FEATURES):
360 super().__init__()
361 d_model = max(16, int(d_model))
362 heads = max(1, int(heads))
363 if d_model % heads != 0:
364 heads = 1
365 self.ctx_dim = int(ctx_dim)
366 self.feat_dim = int(feat_dim)
367 self.history = max(0, int(history))
368 self.event_dim = int(event_dim)
369 self.block_emb = nn.Embedding(int(n_blocks_max), d_model)
370 self.feat_proj = nn.Linear(int(feat_dim), d_model)
371 self.ctx_proj = nn.Linear(int(ctx_dim), d_model)
372 self.event_proj = nn.Linear(self.event_dim, d_model)
373 self.kind_emb = nn.Embedding(3, d_model)
374 self.event_pos = nn.Embedding(max(1, self.history), d_model)
375 self.cls = nn.Parameter(torch.zeros(1, 1, d_model))
376 enc = nn.TransformerEncoderLayer(
377 d_model=d_model, nhead=heads, dim_feedforward=max(32, d_model * 4),
378 dropout=0.0, activation="gelu", batch_first=True, norm_first=True,
379 )
380 self.encoder = nn.TransformerEncoder(enc, num_layers=max(1, int(layers)))
381 self.ln = nn.LayerNorm(d_model)
382 self.value = nn.Sequential(
383 nn.LayerNorm(d_model * 2),
384 nn.Linear(d_model * 2, d_model),
385 nn.GELU(),
386 nn.Linear(d_model, 1),
387 )
388 nn.init.normal_(self.cls, std=0.02)
389
390 @staticmethod
391 def _fit_last_dim(x, dim):
392 if x.size(-1) == dim:
393 return x
394 if x.size(-1) > dim:
395 return x[..., :dim]
396 return F.pad(x, (0, dim - x.size(-1)))
397
398 def forward(self, block_ids, feats, ctx, history=None):
399 feats = self._fit_last_dim(feats.float(), self.feat_dim)
400 ctx = self._fit_last_dim(ctx.float(), self.ctx_dim)
401 B = feats.size(1)
402 bt = self.block_emb(block_ids.clamp(min=0, max=self.block_emb.num_embeddings - 1)) + self.feat_proj(feats)
403 bt = bt + self.kind_emb(torch.ones(B, dtype=torch.long, device=feats.device)).unsqueeze(0)
404 ctx_tok = self.cls + self.ctx_proj(ctx).unsqueeze(1)
405 ctx_tok = ctx_tok + self.kind_emb(torch.zeros(1, dtype=torch.long, device=feats.device)).view(1, 1, -1)
406 tokens = [ctx_tok, bt]
407 if history is not None and self.history > 0:
408 if not torch.is_tensor(history):
409 history = torch.tensor(history, dtype=feats.dtype, device=feats.device)
410 else:
411 history = history.to(device=feats.device, dtype=feats.dtype)
412 if history.dim() == 2:
413 history = history.unsqueeze(0)
414 if history.dim() == 3 and history.numel() > 0:
415 if history.size(0) == 1 and feats.size(0) > 1:
416 history = history.expand(feats.size(0), -1, -1)
417 elif history.size(0) != feats.size(0):
418 history = history[:1].expand(feats.size(0), -1, -1)
419 if history.size(1) > self.history:
420 history = history[:, -self.history :, :]
421 history = self._fit_last_dim(history, self.event_dim)
422 H = history.size(1)
423 if H > 0:
424 pos = torch.arange(H, dtype=torch.long, device=feats.device).clamp(max=max(0, self.history - 1))
425 kind = torch.full((H,), 2, dtype=torch.long, device=feats.device)
426 ht = self.event_proj(history) + self.event_pos(pos).unsqueeze(0) + self.kind_emb(kind).unsqueeze(0)
427 tokens.append(ht)
428 h = self.ln(self.encoder(torch.cat(tokens, dim=1)))
429 ctx_h = h[:, 0:1, :].expand(-1, B, -1)
430 block_h = h[:, 1 : 1 + B, :]
431 return self.value(torch.cat([block_h, ctx_h], dim=-1)).squeeze(-1)
432
433
434def _dblock_router_mode(args):
435 return str(getattr(args, "dblock_router", "heuristic") or "heuristic").lower()
436
437
438def _dblock_router_enabled(args):
439 return _dblock_router_mode(args) in {"transformer", "learned", "neural"}
440
441
442def _dblock_router_boot(state, args, ctx_dim=None):
443 if not _dblock_router_enabled(args):
444 return
445 hidden = int(getattr(args, "dblock_router_hidden", 64) or 64)
446 heads = int(getattr(args, "dblock_router_heads", 4) or 4)
447 layers = int(getattr(args, "dblock_router_layers", 2) or 2)
448 lr = float(getattr(args, "dblock_router_lr", 0.002) or 0.002)
449 history = max(8, min(128, int(getattr(args, "dblock_router_history", _DBLOCK_ROUTER_HISTORY) or _DBLOCK_ROUTER_HISTORY)))
450 cdim = int(ctx_dim or state.get("router_ctx_dim", 0) or 64)
451 state["router_ctx_dim"] = cdim
452 router = _DblockLearnedRouter(ctx_dim=cdim, d_model=hidden, heads=heads, layers=layers, history=history).to("cpu")
453 state["router"] = router
454 state["router_opt"] = torch.optim.AdamW(router.parameters(), lr=lr, weight_decay=1e-3)
455 state["router_target_ema"] = None
456 state["router_target_abs_ema"] = None
457 state["router_train_loss"] = None
458 state["router_last"] = None
459 state["router_history"] = []
460 state["router_history_limit"] = history
461 print(
462 f"[dblock] learned_router=ctx_seq_transformer hidden={hidden} heads={heads} layers={layers} ctx_dim={cdim} history={history} lr={lr:g} "
463 f"blend={float(getattr(args, 'dblock_router_blend', 0.35)):.2f} "
464 f"ramp_steps={int(getattr(args, 'dblock_router_ramp_steps', 256) or 0)}",
465 flush=True,
466 )
467
468
469def _dblock_router_features(state, args):
470 B = int(state["B"])
471 step = int(state.get("step", 0))
472 counts = list(state.get("counts", [0 for _ in range(B)]))
473 if len(counts) != B:
474 counts = [0 for _ in range(B)]
475 emas = list(state.get("loss_ema", [None for _ in range(B)]))
476 if len(emas) != B:
477 emas = [None for _ in range(B)]
478 last_seen = list(state.get("last_seen", [-1 for _ in range(B)]))
479 if len(last_seen) != B:
480 last_seen = [-1 for _ in range(B)]
481 bsig = list(state.get("bsig", _block_sigmas(B, *_dblock_sigma_config(args))))
482 max_count = max(1, max(counts) if counts else 1)
483 known = [float(x) for x in emas if x is not None and math.isfinite(float(x))]
484 center = sum(known) / len(known) if known else 0.0
485 scale = (sum((x - center) ** 2 for x in known) / len(known)) ** 0.5 if len(known) > 1 else max(1.0, abs(center) * 0.05)
486 scale = max(1e-3, scale)
487 stale = [step - last_seen[i] if last_seen[i] >= 0 else step + 1 for i in range(B)]
488 max_stale = int(getattr(args, "dblock_max_stale_steps", 64) or 0)
489 stale_denom = float(max(1, max_stale if max_stale > 0 else max(stale) if stale else 1))
490 logs = [math.log(max(1e-9, float(x))) for x in bsig]
491 log_min = min(logs) if logs else 0.0
492 log_span = max(1e-6, (max(logs) - log_min) if logs else 1.0)
493 feats = []
494 for i in range(B):
495 ema = emas[i]
496 known_flag = 1.0 if ema is not None and math.isfinite(float(ema)) else 0.0
497 loss_z = 0.0 if not known_flag else max(-5.0, min(5.0, (float(ema) - center) / scale))
498 lo = logs[min(i, len(logs) - 1)] if logs else 0.0
499 hi = logs[min(i + 1, len(logs) - 1)] if logs else lo
500 sig_mid = ((0.5 * (lo + hi)) - log_min) / log_span
501 feats.append([
502 loss_z, known_flag, float(counts[i]) / float(max_count),
503 max(0.0, float(max_count - counts[i]) / float(max_count)),
504 min(1.0, max(0.0, float(stale[i]) / stale_denom)), float(sig_mid),
505 ])
506 block_ids = torch.arange(B, dtype=torch.long).unsqueeze(0)
507 ft = torch.tensor([feats], dtype=torch.float32)
508 cdim = int(state.get("router_ctx_dim", 0) or 0)
509 ctx = state.get("router_ctx")
510 if torch.is_tensor(ctx) and cdim > 0 and ctx.numel() == cdim:
511 cv = ctx.detach().reshape(1, cdim).float()
512 else:
513 cv = torch.zeros(1, max(1, cdim))
514 return block_ids, ft, cv
515
516
517def _dblock_router_clip(x, lo=-5.0, hi=5.0):
518 try:
519 x = float(x)
520 except Exception:
521 return 0.0
522 if not math.isfinite(x):
523 return 0.0
524 return max(lo, min(hi, x))
525
526
527def _dblock_router_history_features(state, args):
528 limit = int(state.get("router_history_limit", getattr(args, "dblock_router_history", _DBLOCK_ROUTER_HISTORY)) or 0)
529 limit = max(0, min(128, limit))
530 if limit <= 0:
531 return torch.zeros((1, 0, _DBLOCK_ROUTER_EVENT_FEATURES), dtype=torch.float32)
532 hist = list(state.get("router_history", []))[-limit:]
533 if not hist:
534 return torch.zeros((1, 0, _DBLOCK_ROUTER_EVENT_FEATURES), dtype=torch.float32)
535 B = int(state["B"])
536 step = int(state.get("step", 0))
537 losses = []
538 for rec in hist:
539 try:
540 loss = float(rec.get("loss", 0.0))
541 except Exception:
542 loss = 0.0
543 if math.isfinite(loss):
544 losses.append(loss)
545 center = sum(losses) / len(losses) if losses else 0.0
546 scale = (sum((x - center) ** 2 for x in losses) / len(losses)) ** 0.5 if len(losses) > 1 else max(1.0, abs(center) * 0.05)
547 scale = max(1e-3, scale)
548 rows = []
549 for rec in hist:
550 rec_step = int(rec.get("step", -1))
551 block = max(0, min(B - 1, int(rec.get("block", 0))))
552 age = max(0, step - rec_step)
553 try:
554 rec_loss = float(rec.get("loss", center))
555 except Exception:
556 rec_loss = center
557 loss = _dblock_router_clip((rec_loss - center) / scale)
558 rows.append([
559 float(block) / float(max(1, B - 1)),
560 _dblock_router_clip(rec.get("target", 0.0)),
561 loss,
562 max(0.0, min(1.0, float(rec.get("count_norm", 0.0)))),
563 max(0.0, min(1.0, float(rec.get("stale_norm", 0.0)))),
564 min(1.0, math.log1p(age) / math.log1p(max(2, limit))),
565 min(1.0, math.log1p(max(0, rec_step)) / math.log1p(10000.0)),
566 1.0 if float(rec.get("router_choice", 0.0)) > 0.0 else 0.0,
567 max(0.0, min(1.0, float(rec.get("blend", 0.0)))),
568 1.0,
569 ])
570 return torch.tensor([rows], dtype=torch.float32)
571
572
573def _dblock_router_append_history(state, args, bi, loss_float, target_val):
574 limit = int(state.get("router_history_limit", getattr(args, "dblock_router_history", _DBLOCK_ROUTER_HISTORY)) or _DBLOCK_ROUTER_HISTORY)
575 limit = max(0, min(128, limit))
576 if limit <= 0:
577 return
578 B = int(state["B"])
579 step = int(state.get("step", 0))
580 counts = list(state.get("counts", [0 for _ in range(B)]))
581 if len(counts) != B:
582 counts = [0 for _ in range(B)]
583 last_seen = list(state.get("last_seen", [-1 for _ in range(B)]))
584 if len(last_seen) != B:
585 last_seen = [-1 for _ in range(B)]
586 max_count = max(1, max(counts) if counts else 1)
587 stale = step - last_seen[int(bi)] if 0 <= int(bi) < len(last_seen) and last_seen[int(bi)] >= 0 else step + 1
588 max_stale = int(getattr(args, "dblock_max_stale_steps", 64) or 0)
589 stale_denom = float(max(1, max_stale if max_stale > 0 else stale))
590 route = state.get("router_last")
591 router_choice = 0.0
592 blend = 0.0
593 if isinstance(route, dict):
594 router_choice = 1.0 if int(route.get("choice", -1)) == int(bi) else 0.0
595 blend = float(route.get("blend", 0.0))
596 hist = state.setdefault("router_history", [])
597 hist.append({
598 "step": int(step),
599 "block": int(bi),
600 "loss": float(loss_float),
601 "target": float(target_val),
602 "count_norm": float(counts[int(bi)]) / float(max_count) if 0 <= int(bi) < len(counts) else 0.0,
603 "stale_norm": min(1.0, max(0.0, float(stale) / stale_denom)),
604 "router_choice": router_choice,
605 "blend": blend,
606 })
607 if len(hist) > limit:
608 del hist[:-limit]
609
610
611def _dblock_router_norm(xs):
612 vals = [0.0 if not math.isfinite(float(x)) else float(x) for x in xs]
613 if not vals:
614 return vals
615 mean = sum(vals) / len(vals)
616 scale = max(1e-6, (sum((x - mean) ** 2 for x in vals) / len(vals)) ** 0.5)
617 return [(x - mean) / scale for x in vals]
618
619
620def _dblock_fleet_lane_keys(args):
621 keys = []
622 for env_key in ("AGILLM_FLEET_LANE", "AGILLM_WORKER_ID", "AGILLM_LANE_ID"):
623 val = os.environ.get(env_key, "")
624 if val:
625 keys.append(str(val))
626 save_dir = str(getattr(args, "save_dir", "") or "")
627 if save_dir:
628 keys.append(os.path.basename(save_dir.rstrip("/")))
629 keys.append(save_dir)
630 return [k for i, k in enumerate(keys) if k and k not in keys[:i]]
631
632
633def _dblock_fleet_router_scores(state, args, base_scores):
634 state["fleet_router_last"] = None
635 if not base_scores:
636 return None
637 try:
638 cfg = get_hot_config()
639 except Exception:
640 return None
641 spec = cfg.get("dblock_fleet_router") or cfg.get("dblock_fleet_route")
642 if not isinstance(spec, dict):
643 return None
644 if str(spec.get("enabled", True)).lower() in {"0", "false", "off", "no"}:
645 return None
646 lanes = spec.get("lanes") if isinstance(spec.get("lanes"), dict) else {}
647 lane_key = None
648 lane = None
649 for key in _dblock_fleet_lane_keys(args):
650 cand = lanes.get(key)
651 if isinstance(cand, dict):
652 lane_key, lane = key, cand
653 break
654 if lane is None:
655 return None
656 bias = lane.get("bias", lane.get("block_bias", lane.get("biases")))
657 if not isinstance(bias, (list, tuple)):
658 return None
659 B = int(state.get("B", len(base_scores)) or len(base_scores))
660 if len(bias) != B or len(base_scores) != B:
661 return None
662 vals = []
663 for x in bias:
664 try:
665 fx = float(x)
666 except Exception:
667 fx = 0.0
668 vals.append(0.0 if not math.isfinite(fx) else max(-3.0, min(3.0, fx)))
669 if not any(abs(x) > 1e-9 for x in vals):
670 return None
671 strength = float(lane.get("strength", spec.get("strength", 0.20)) or 0.0)
672 strength = max(0.0, min(1.0, strength))
673 if strength <= 1e-9:
674 return None
675 base = [float(x) if math.isfinite(float(x)) else 0.0 for x in base_scores]
676 mean = sum(base) / len(base)
677 scale = max(1e-3, (sum((x - mean) ** 2 for x in base) / len(base)) ** 0.5)
678 adjusted = [base[i] + strength * scale * vals[i] for i in range(B)]
679 state["fleet_router_last"] = {
680 "lane": str(lane_key),
681 "role": str(lane.get("role", "")),
682 "strength": float(strength),
683 "bias": [float(x) for x in vals],
684 "updated_at": spec.get("updated_at", ""),
685 }
686 return adjusted
687
688
689def _dblock_router_choose(state, args, heuristic_scores):
690 state["router_last"] = None
691 if not _dblock_router_enabled(args):
692 return None
693 router = state.get("router")
694 if router is None:
695 return None
696 B = int(state["B"])
697 step = int(state.get("step", 0))
698 warmup = int(getattr(args, "dblock_warmup_steps", max(8, B * 2)))
699 ramp_steps = int(getattr(args, "dblock_router_ramp_steps", 256) or 0)
700 blend_base = max(0.0, min(1.0, float(getattr(args, "dblock_router_blend", 0.35) or 0.0)))
701 if step < warmup or blend_base <= 0.0:
702 return None
703 ramp = 1.0 if ramp_steps <= 0 else min(1.0, max(0.0, float(step - warmup) / float(ramp_steps)))
704 blend = blend_base * ramp
705 if blend <= 1e-6:
706 return None
707 history_features = _dblock_router_history_features(state, args)
708 with torch.no_grad():
709 router.eval()
710 pred = router(*_dblock_router_features(state, args), history=history_features)[0].detach().cpu().tolist()
711 h = _dblock_router_norm(heuristic_scores)
712 q = _dblock_router_norm(pred)
713 if len(h) != B or len(q) != B:
714 return None
715 counts = state.get("counts", [0 for _ in range(B)])
716 combined = [(1.0 - blend) * h[i] + blend * q[i] for i in range(B)]
717 choice = max(range(B), key=lambda i: (combined[i], -counts[i], -i))
718 state["router_last"] = {
719 "mode": "ctx_seq_transformer",
720 "choice": int(choice),
721 "blend": float(blend),
722 "history": int(history_features.size(1)),
723 "pred": [float(x) for x in pred],
724 }
725 return choice
726
727
728def _dblock_router_update(state, args, bi, loss_value):
729 if not _dblock_router_enabled(args):
730 return
731 router, opt = state.get("router"), state.get("router_opt")
732 if router is None or opt is None:
733 return
734 try:
735 loss_float = float(loss_value)
736 except Exception:
737 return
738 if not math.isfinite(loss_float):
739 return
740 baseline = state.get("router_target_ema")
741 scale = state.get("router_target_abs_ema")
742 if baseline is None or not math.isfinite(float(baseline)):
743 baseline = loss_float
744 if scale is None or not math.isfinite(float(scale)) or float(scale) < 1e-3:
745 scale = max(1.0, abs(loss_float) * 0.05)
746 target_val = max(-5.0, min(5.0, (loss_float - float(baseline)) / max(1e-3, float(scale))))
747 router.train()
748 pred = router(*_dblock_router_features(state, args), history=_dblock_router_history_features(state, args))[0, int(bi)]
749 fit_loss = F.smooth_l1_loss(pred, pred.detach().new_tensor(target_val))
750 opt.zero_grad(set_to_none=True)
751 fit_loss.backward()
752 nn.utils.clip_grad_norm_(router.parameters(), 1.0)
753 opt.step()
754 diff = abs(loss_float - float(baseline))
755 state["router_target_ema"] = 0.98 * float(baseline) + 0.02 * loss_float
756 state["router_target_abs_ema"] = 0.98 * float(scale) + 0.02 * max(1e-3, diff)
757 state["router_train_loss"] = float(fit_loss.detach().cpu())
758 _dblock_router_append_history(state, args, bi, loss_float, target_val)
759
760
761def _dblock_get_candidates(L):
762 c = []
763 # 1. Uniform candidates for b in [2, 3, 4, 6]
764 for b in [2, 3, 4, 6]:
765 per = max(1, L // b)
766 asg = [list(range(i * per, (i + 1) * per)) for i in range(b)]
767 asg[-1] = list(range((b - 1) * per, L))
768 c.append((b, asg, f"Uniform-{b}"))
769
770 # 2. Non-uniform candidates for B=3
771 # Middle-heavy (e.g. 25%, 50%, 25%)
772 m_h = [max(1, L // 4), max(1, L // 2)]
773 m_h.append(L - sum(m_h))
774 asg = []
775 curr = 0
776 for size in m_h:
777 asg.append(list(range(curr, curr + size)))
778 curr += size
779 c.append((3, asg, "Middle-Heavy-3"))
780
781 # End-heavy (e.g. 20%, 35%, 45%)
782 e_h = [max(1, int(L * 0.20)), max(1, int(L * 0.35))]
783 e_h.append(L - sum(e_h))
784 asg = []
785 curr = 0
786 for size in e_h:
787 asg.append(list(range(curr, curr + size)))
788 curr += size
789 c.append((3, asg, "End-Heavy-3"))
790
791 # Start-heavy (e.g. 45%, 35%, 20%)
792 s_h = [max(1, int(L * 0.45)), max(1, int(L * 0.35))]
793 s_h.append(L - sum(s_h))
794 asg = []
795 curr = 0
796 for size in s_h:
797 asg.append(list(range(curr, curr + size)))
798 curr += size
799 c.append((3, asg, "Start-Heavy-3"))
800
801 # 3. Non-uniform candidates for B=4
802 # Middle-heavy (e.g. 20%, 30%, 30%, 20%)
803 m_h4 = [max(1, int(L * 0.20)), max(1, int(L * 0.30)), max(1, int(L * 0.30))]
804 m_h4.append(L - sum(m_h4))
805 asg = []
806 curr = 0
807 for size in m_h4:
808 asg.append(list(range(curr, curr + size)))
809 curr += size
810 c.append((4, asg, "Middle-Heavy-4"))
811
812 # End-heavy (e.g. 15%, 25%, 30%, 30%)
813 e_h4 = [max(1, int(L * 0.15)), max(1, int(L * 0.25)), max(1, int(L * 0.30))]
814 e_h4.append(L - sum(e_h4))
815 asg = []
816 curr = 0
817 for size in e_h4:
818 asg.append(list(range(curr, curr + size)))
819 curr += size
820 c.append((4, asg, "End-Heavy-4"))
821
822 return c
823
824def _dblock_init(core, args):
825 L = len(core.blocks)
826 auto_search = getattr(args, "auto_dblock_search", False)
827
828 if auto_search:
829 candidates = _dblock_get_candidates(L)
830 print(f"[dblock] Auto Search enabled with {len(candidates)} candidates.")
831 B, asg, name = candidates[0]
832 state = {
833 "auto_search": True,
834 "candidates": candidates,
835 "candidate_idx": 0,
836 "search_step": 0,
837 "search_interval": 20,
838 "scores": [],
839 }
840 else:
841 B = int(getattr(args, "dblock_blocks", 4))
842 sp = max(1, L // B)
843 asg = [list(range(i * sp, (i + 1) * sp)) for i in range(B)]
844 asg[-1] = list(range((B - 1) * sp, L))
845 state = {"auto_search": False}
846
847 bsig = _block_sigmas(B, *_dblock_sigma_config(args))
848 schedule = getattr(args, "dblock_schedule", "loss_balanced")
849 print(f"[dblock] DiffusionBlocks mode: {L} layers -> {B} blocks {asg}")
850 print(f"[dblock] schedule={schedule} sigma boundaries: {[round(x, 3) for x in bsig]}")
851
852 state.update({
853 "B": B,
854 "assign": asg,
855 "bsig": bsig,
856 "step": 0,
857 "counts": [0 for _ in range(B)],
858 "loss_ema": [None for _ in range(B)],
859 "last_seen": [-1 for _ in range(B)],
860 })
861 if bool(getattr(args, "dblock_looped", False)):
862 loop_layers = int(getattr(args, "dblock_loop_layers", 0) or 0)
863 if loop_layers <= 0:
864 loop_layers = max(1, L // max(1, B))
865 loop_layers = max(1, min(loop_layers, L))
866 loop_start = max(0, min(int(getattr(args, "dblock_loop_start", 0) or 0), L - loop_layers))
867 loop_group = list(range(loop_start, loop_start + loop_layers))
868 if not hasattr(core, "dblock_loop_embed"):
869 d = int(getattr(core.emb, "embedding_dim", 0))
870 core.dblock_loop_embed = nn.Embedding(B, d).to(core.emb.weight.device)
871 nn.init.normal_(core.dblock_loop_embed.weight, mean=0.0, std=0.02)
872 state.update({
873 "looped": True,
874 "loop_group": loop_group,
875 "loop_layers": loop_layers,
876 "loop_start": loop_start,
877 })
878 print(
879 f"[dblock-looped] enabled: shared_layers={loop_group} bands={B} "
880 f"unrolled_depth={loop_layers * B} one-band-per-step no_bptt=True",
881 flush=True,
882 )
883 _dblock_router_boot(state, args, ctx_dim=int(getattr(core.emb, "embedding_dim", 0)) or None)
884 return state
885
886
887def _choose_block(state, args):
888 if not state.get("auto_search", False) and state.get("step", 0) % 100 == 0:
889 try:
890 cfg = get_hot_config()
891 if "dblock_blocks" in cfg:
892 new_B = int(cfg["dblock_blocks"])
893 if new_B != state.get("B"):
894 L = sum(len(x) for x in state["assign"]) if "assign" in state else 28
895 new_sp = max(1, L // new_B)
896 new_asg = [list(range(i * new_sp, (i + 1) * new_sp)) for i in range(new_B)]
897 new_asg[-1] = list(range((new_B - 1) * new_sp, L))
898
899 print(f"[dblock] Dynamically adjusting block configuration from hot_config: B={state['B']} -> {new_B}, assign={new_asg}", flush=True)
900 state["B"] = new_B
901 state["assign"] = new_asg
902 state["bsig"] = _block_sigmas(new_B, *_dblock_sigma_config(args))
903 state["counts"] = [0] * new_B
904 state["loss_ema"] = [None] * new_B
905 state["last_seen"] = [-1] * new_B
906 except Exception as e:
907 print(f"[dblock] Error reloading hot_config in _choose_block: {e}", flush=True)
908
909 if state.get("auto_search", False) and state["candidate_idx"] < len(state["candidates"]):
910 state["search_step"] += 1
911 if "search_start_time" not in state:
912 state["search_start_time"] = time.perf_counter()
913 state["search_tokens"] = 0
914
915 if state["search_step"] >= state["search_interval"]:
916 valid_emas = [e for e in state["loss_ema"] if e is not None]
917 avg_loss = sum(valid_emas) / max(1, len(valid_emas)) if valid_emas else float('inf')
918
919 elapsed = time.perf_counter() - state["search_start_time"]
920 tokens = state.get("search_tokens", 0)
921 tokps = tokens / max(1e-9, elapsed)
922
923 cand = state["candidates"][state["candidate_idx"]]
924 cand_name = cand[2] if len(cand) > 2 else f"Candidate-{state['candidate_idx']}"
925
926 state["scores"].append({
927 "idx": state["candidate_idx"],
928 "B": state["B"],
929 "assign": state["assign"],
930 "name": cand_name,
931 "loss": avg_loss,
932 "tokps": tokps
933 })
934 print(f"[dblock] Candidate {state['candidate_idx']} ({cand_name}) complete: loss={avg_loss:.4f} speed={tokps:.1f} tok/s", flush=True)
935
936 state["candidate_idx"] += 1
937 state["search_step"] = 0
938 if "search_start_time" in state:
939 del state["search_start_time"]
940 state["search_tokens"] = 0
941
942 if state["candidate_idx"] < len(state["candidates"]):
943 B, asg, cand_name = state["candidates"][state["candidate_idx"]]
944 state["B"] = B
945 state["assign"] = asg
946 state["bsig"] = _block_sigmas(B, *_dblock_sigma_config(args))
947 state["counts"] = [0] * B
948 state["loss_ema"] = [None] * B
949 state["last_seen"] = [-1] * B
950 print(f"[dblock] Switched to candidate {state['candidate_idx']} ({cand_name}): {B} blocks {asg}", flush=True)
951 else:
952 # Select the candidate with highest speed/loss utility
953 best_cand = None
954 best_utility = -1.0
955 for score_entry in state["scores"]:
956 loss = score_entry["loss"]
957 tokps = score_entry["tokps"]
958 utility = tokps / max(1e-3, loss)
959 score_entry["utility"] = utility
960 if utility > best_utility:
961 best_utility = utility
962 best_cand = score_entry
963
964 B = best_cand["B"]
965 asg = best_cand["assign"]
966 state["B"] = B
967 state["assign"] = asg
968 state["bsig"] = _block_sigmas(B, *_dblock_sigma_config(args))
969 state["auto_search"] = False
970 print(f"[dblock] Search complete. Locked in best candidate {best_cand['name']} (Utility={best_utility:.2f}, Loss={best_cand['loss']:.4f}, Speed={best_cand['tokps']:.1f} tok/s): {B} blocks {asg}", flush=True)
971 B = state["B"]
972 schedule = str(getattr(args, "dblock_schedule", "loss_balanced") or "loss_balanced").lower()
973 step = int(state.get("step", 0))
974 counts = state.setdefault("counts", [0 for _ in range(B)])
975 if len(counts) != B:
976 counts[:] = [0 for _ in range(B)]
977 emas = state.setdefault("loss_ema", [None for _ in range(B)])
978 if len(emas) != B:
979 emas[:] = [None for _ in range(B)]
980 last_seen = state.setdefault("last_seen", [-1 for _ in range(B)])
981 if len(last_seen) != B:
982 last_seen[:] = [-1 for _ in range(B)]
983 state["router_last"] = None
984 state["fleet_router_last"] = None
985 if schedule == "random":
986 return random.randrange(B)
987 if schedule == "roundrobin":
988 return step % B
989
990 explore = max(0.0, min(1.0, float(getattr(args, "dblock_explore", 0.05))))
991 warmup = int(getattr(args, "dblock_warmup_steps", max(8, B * 2)))
992
993 def least_trained():
994 return min(range(B), key=lambda i: (counts[i], last_seen[i], i))
995
996 if step < warmup or any(c == 0 for c in counts):
997 return least_trained()
998
999 max_stale = int(getattr(args, "dblock_max_stale_steps", 64) or 0)
1000 stale = [step - last_seen[i] if last_seen[i] >= 0 else step + 1 for i in range(B)]
1001 if max_stale > 0 and max(stale) >= max_stale:
1002 return max(range(B), key=lambda i: (stale[i], -counts[i], -i))
1003
1004 max_count = max(counts) if counts else 0
1005 min_count = min(counts) if counts else 0
1006 max_skew = float(getattr(args, "dblock_max_count_skew", 1.35) or 0.0)
1007 if max_skew > 1.0 and min_count > 0 and (max_count / max(1, min_count)) > max_skew:
1008 return least_trained()
1009
1010 if explore > 0.0 and random.random() < explore:
1011 return least_trained()
1012
1013 stale_bonus = float(getattr(args, "dblock_stale_bonus", 0.35) or 0.0)
1014 undertrain_bonus = float(getattr(args, "dblock_undertrain_bonus", 0.25) or 0.0)
1015 stale_denom = float(max(1, max_stale if max_stale > 0 else max(stale) if stale else 1))
1016 count_denom = float(max(1, max_count))
1017
1018 def score(i):
1019 loss_score = -1.0 if emas[i] is None else float(emas[i])
1020 stale_score = stale_bonus * min(1.0, max(0.0, stale[i] / stale_denom))
1021 undertrain_score = undertrain_bonus * max(0.0, (max_count - counts[i]) / count_denom)
1022 return (loss_score + stale_score + undertrain_score, -counts[i], stale[i], -i)
1023
1024 base_scores = [float(score(i)[0]) for i in range(B)]
1025 route_scores = _dblock_fleet_router_scores(state, args, base_scores) or base_scores
1026 if route_scores is base_scores:
1027 heuristic_choice = max(range(B), key=score)
1028 else:
1029 heuristic_choice = max(range(B), key=lambda i: (route_scores[i], -counts[i], stale[i], -i))
1030 learned_choice = _dblock_router_choose(state, args, route_scores)
1031 return heuristic_choice if learned_choice is None else learned_choice
1032
1033
1034def _sample_sigma(ids, lo, hi, args, state):
1035 cur_step = int(state.get("step", 0))
1036 curriculum = int(getattr(args, "dblock_sigma_curriculum_steps", 0))
1037 if curriculum > 0:
1038 frac = min(1.0, max(0.05, (cur_step + 1) / float(curriculum)))
1039 hi = lo * ((hi / max(lo, 1e-8)) ** frac)
1040 mode = str(getattr(args, "dblock_sigma_sampling", "lognormal") or "lognormal").lower()
1041 if mode in {"lognormal", "truncated_lognormal", "edm"}:
1042 _, _, pm, ps = _dblock_sigma_config(args)
1043 qa = _cdf((math.log(max(lo, 1e-6)) - pm) / ps)
1044 qb = _cdf((math.log(max(hi, lo * 1.0001)) - pm) / ps)
1045 qa = min(max(qa, 1e-7), 1.0 - 1e-7)
1046 qb = min(max(qb, qa + 1e-7), 1.0 - 1e-7)
1047 n = int(ids.size(0))
1048 if bool(getattr(args, "dblock_sigma_stratified", True)) and n > 1:
1049 # Beyond the DBT paper: randomized quantile strata reduce Monte Carlo
1050 # variance of the conditional p_noise integral for each block.
1051 u = (torch.arange(n, device=ids.device, dtype=torch.float32) + torch.rand((), device=ids.device)) / float(n)
1052 u = u.index_select(0, torch.randperm(n, device=ids.device))
1053 else:
1054 u = torch.rand(n, device=ids.device, dtype=torch.float32)
1055 q = qa + (qb - qa) * u
1056 q = q.clamp(1e-7, 1.0 - 1e-7)
1057 z = torch.erfinv(2.0 * q - 1.0) * math.sqrt(2.0)
1058 return torch.exp(torch.tensor(pm, device=ids.device, dtype=torch.float32) + float(ps) * z)
1059 sig_np = np.exp(
1060 np.random.uniform(
1061 math.log(max(lo, 1e-4)),
1062 math.log(max(hi, lo + 1e-4)),
1063 ids.size(0),
1064 ).astype("float32")
1065 )
1066 return torch.from_numpy(sig_np).to(ids.device)
1067
1068
1069def _maybe_log(
1070 state,
1071 args,
1072 bi,
1073 layers,
1074 ar_val,
1075 sat_val,
1076 nat_val,
1077 total_val,
1078 peak_alloc,
1079 peak_reserved,
1080 objective=None,
1081 raw_avg=None,
1082 raw_total=None,
1083 edm_weight=None,
1084):
1085 log_every = int(getattr(args, "dblock_log_every", 50))
1086 step = int(state.get("step", 0))
1087 if log_every <= 0 or step % log_every != 0:
1088 return
1089 counts_list = state.get("counts", [])
1090 last_seen = state.get("last_seen", [-1 for _ in counts_list])
1091 counts = ",".join(str(x) for x in counts_list)
1092 emas = ",".join("nan" if x is None else f"{x:.2f}" for x in state.get("loss_ema", []))
1093 stale = ",".join(str(max(0, step - int(last_seen[i]))) for i in range(min(len(counts_list), len(last_seen))))
1094 mem = ""
1095 if peak_alloc is not None:
1096 mem = f" peak_alloc={peak_alloc:.2f}GB peak_reserved={peak_reserved:.2f}GB"
1097 display = float(raw_avg) if raw_avg is not None and math.isfinite(float(raw_avg)) else float(total_val)
1098 raw_part = ""
1099 if raw_total is not None:
1100 raw_part += f" raw_sum={float(raw_total):.3f}"
1101 if edm_weight is not None:
1102 raw_part += f" edm_w={float(edm_weight):.3f}"
1103 route = state.get("router_last")
1104 if isinstance(route, dict):
1105 pred = ",".join(f"{float(x):.2f}" for x in route.get("pred", []))
1106 hist = route.get("history")
1107 hist_part = "" if hist is None else f" hist={int(hist)}"
1108 raw_part += f" router={route.get('mode', 'none')} blend={float(route.get('blend', 0.0)):.2f}{hist_part} pred=[{pred}]"
1109 rloss = state.get("router_train_loss")
1110 if rloss is not None:
1111 raw_part += f" router_fit={float(rloss):.3f}"
1112 fleet = state.get("fleet_router_last")
1113 if isinstance(fleet, dict):
1114 fbias = fleet.get("bias", [])
1115 top = []
1116 try:
1117 top = sorted(range(len(fbias)), key=lambda j: abs(float(fbias[j])), reverse=True)[:3]
1118 except Exception:
1119 top = []
1120 top_part = ",".join(f"{j}:{float(fbias[j]):+.2f}" for j in top)
1121 raw_part += f" fleet={fleet.get('lane', '')} role={fleet.get('role', '')} strength={float(fleet.get('strength', 0.0)):.2f}"
1122 if top_part:
1123 raw_part += f" fleet_bias=[{top_part}]"
1124 print(
1125 f"[dblock] step={step} block={bi} obj={objective or 'mixed'} layers={layers} "
1126 f"loss={display:.3f} weighted={total_val:.3f} ar={ar_val:.3f} sat={sat_val:.3f} nat={nat_val:.3f}"
1127 f"{raw_part} counts=[{counts}] ema=[{emas}] stale=[{stale}]{mem}",
1128 flush=True,
1129 )
1130
1131
1132def _update_stats(state, bi, loss_value, args=None):
1133 if args is not None:
1134 _dblock_router_update(state, args, bi, loss_value)
1135 B = state["B"]
1136 counts = state.setdefault("counts", [0 for _ in range(B)])
1137 emas = state.setdefault("loss_ema", [None for _ in range(B)])
1138 last_seen = state.setdefault("last_seen", [-1 for _ in range(B)])
1139 if len(last_seen) != B:
1140 last_seen[:] = [-1 for _ in range(B)]
1141 counts[bi] += 1
1142 last_seen[bi] = int(state.get("step", 0))
1143 prev = emas[bi]
1144 beta = 0.96
1145 emas[bi] = float(loss_value) if prev is None else beta * float(prev) + (1.0 - beta) * float(loss_value)
1146 state["step"] = int(state.get("step", 0)) + 1
1147
1148
1149def _activation_offload_enabled(args):
1150 return bool(getattr(args, "dblock_activation_offload", False)) and torch.cuda.is_available()
1151
1152
1153def _activation_offload_hooks(args):
1154 min_bytes = int(float(getattr(args, "dblock_activation_offload_min_mb", 1.0) or 1.0) * 1024 * 1024)
1155
1156 def pack(t):
1157 if not torch.is_tensor(t) or not t.is_cuda or not t.is_floating_point() or t.numel() * t.element_size() < min_bytes:
1158 return t
1159 return ("cpu_offload", t.device, t.detach().to("cpu", non_blocking=True))
1160
1161 def unpack(x):
1162 if isinstance(x, tuple) and len(x) == 3 and x[0] == "cpu_offload":
1163 _, dev, cpu_t = x
1164 return cpu_t.to(dev, non_blocking=True)
1165 return x
1166
1167 return torch.autograd.graph.saved_tensors_hooks(pack, unpack)
1168
1169
1170def _dblock_sublayer_base_mode(args):
1171 mode = str(getattr(args, "dblock_sublayer_mode", "off") or "off").strip().lower().replace("-", "_")
1172 if mode in {"none", "disabled"}:
1173 return "off"
1174 return mode
1175
1176
1177def _dblock_sublayer_mode_for_layer(args, state, block_idx, layer_pos):
1178 mode = _dblock_sublayer_base_mode(args)
1179 if mode == "split_alt":
1180 step = int((state or {}).get("step", 0))
1181 return "attn_only" if ((step + int(block_idx) + int(layer_pos)) % 2 == 0) else "ffn_only"
1182 if mode == "cycle":
1183 step = int((state or {}).get("step", 0))
1184 return ("full", "ffn_only", "attn_only")[(step + int(block_idx) + int(layer_pos)) % 3]
1185 return mode
1186
1187
1188def _run_block_forward(block, x, mask, sublayer_mode="off"):
1189 mode = str(sublayer_mode or "off").strip().lower().replace("-", "_")
1190 if mode in {"off", "full"}:
1191 return block(x, mask)
1192 if mode == "attn_only":
1193 n = x.size(1)
1194 return x + block.mha(block.ln1(x), mask, rel_bias_tokens=n)
1195 if mode == "ffn_only":
1196 return x + block.ff(block.ln2(x))
1197 raise ValueError(f"unknown DBlock sublayer mode: {sublayer_mode}")
1198
1199
1200def _run_block(block, x, mask, use_checkpoint, args=None, sublayer_mode="off"):
