CoolFace
Apppublic

OpenTransformer/AGILLM-4.3-ZeroGPU-GUI

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
agillm41.py10221 linesDownload Raw Back to root
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"):

Showing the first 1,200 of 10221 lines. Download the file for the rest.