File size: 4,984 Bytes
a974167
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
#!/usr/bin/env python
"""Convert the repo's torch-saved HMoE child models to a torch-free bundle.

This is the ONLY step that requires torch and the private repo layout. It reads
the unified model directory + family gate feature list and writes a fully
self-contained, torch-free weights bundle that the open-source `hmoe_da`
package loads with numpy alone:

    weights/
      meta.json           # tree, root_key, gate_nodes, gates, gate_repr, leaves, excluded_leaves
      features.json       # 21,604 mouse Title-case training genes (the alignment axis)
      child_models.npz    # all 34 logreg child classifiers (weights/bias/gene_idx)
      MANIFEST.json       # sha256 checksums + provenance

Run inside the m2h env (has torch):
    ~/miniforge3/envs/m2h_env/bin/python hmoe_open/scripts/convert_weights.py

The child models are `kind == "logreg"`: a per-node binary classifier storing a
200-dim weight vector, a bias, and the integer `gene_idx` into the 21,604-gene
space. No scalers are present (verified), so a converted model is exactly
`X[:, gene_idx] @ weights + bias`.
"""
from __future__ import annotations

import hashlib
import json
from pathlib import Path

import numpy as np
import torch

REPO = Path(__file__).resolve().parents[2]
UNI_DIR = REPO / "models" / "moe_v4" / "v4_improved_unified"
FAMILY_FEATURES = (
    REPO / "models" / "moe_v4" / "family_gate" / "family_gate_improved_safe" / "features.json"
)
OUT_DIR = Path(__file__).resolve().parents[1] / "hmoe_da" / "weights"


def safe_id(v: str) -> str:
    return str(v).replace("|", "__").replace(":", "_").replace("/", "_")


def child_key(gate: str, child: str) -> str:
    """Stable, filesystem-agnostic key for a (gate, child) child model."""
    return f"{safe_id(gate)}::{safe_id(child)}"


def main() -> None:
    OUT_DIR.mkdir(parents=True, exist_ok=True)
    meta = json.load(open(UNI_DIR / "meta.json"))

    # 1) meta.json: keep only the fields inference needs (drops training logs).
    meta_out = {
        "pipeline": meta["pipeline"],
        "tree": meta["tree"],
        "root_key": meta["root_key"],
        "gate_nodes": meta["gate_nodes"],
        "gates": meta["gates"],
        "gate_repr": meta["gate_repr"],
        "leaves": meta["leaves"],
        "excluded_leaves": meta.get("excluded_leaves", []),
        "config": {
            "topk": meta.get("config", {}).get("topk"),
            "feature_method": meta.get("config", {}).get("feature_method"),
        },
    }
    (OUT_DIR / "meta.json").write_text(json.dumps(meta_out, indent=2))

    # 2) features.json: the 21,604-gene alignment axis (mouse Title case).
    feats = json.load(open(FAMILY_FEATURES))
    if isinstance(feats, dict):  # some releases wrap it as {"feature_names": [...]}
        feats = feats.get("feature_names", feats)
    assert isinstance(feats, list) and len(feats) == 21604, f"unexpected features: {len(feats)}"
    (OUT_DIR / "features.json").write_text(json.dumps(feats))

    # 3) child_models.npz: every logreg child model, torch-free.
    arrays: dict[str, np.ndarray] = {}
    n_models = 0
    for gk in meta["gate_nodes"]:
        rt = meta["gate_repr"][gk]  # "raw" or "sct"
        for c in meta["gates"][gk]:
            p = UNI_DIR / "child_models" / rt / safe_id(gk) / f"{safe_id(c)}.safe.pt"
            if not p.exists():
                raise FileNotFoundError(p)
            sd = torch.load(p, map_location="cpu", weights_only=False)["state_dict"]
            kind = sd.get("kind")
            if kind != "logreg":
                raise ValueError(f"expected logreg, got {kind!r} at {p}")
            if sd.get("scaler"):
                raise ValueError(f"unexpected scaler at {p}; converter assumes none")
            key = child_key(gk, c)
            arrays[f"{key}::weights"] = np.asarray(sd["weights"], dtype=np.float32).reshape(-1)
            arrays[f"{key}::bias"] = np.asarray([float(sd["bias"])], dtype=np.float32)
            arrays[f"{key}::gene_idx"] = np.asarray(sd["gene_idx"], dtype=np.int32).reshape(-1)
            n_models += 1
    np.savez_compressed(OUT_DIR / "child_models.npz", **arrays)

    # 4) MANIFEST.json: checksums + provenance.
    def sha256(path: Path) -> str:
        return hashlib.sha256(path.read_bytes()).hexdigest()

    manifest = {
        "model": "v4_improved_unified",
        "source_dir": str(UNI_DIR.relative_to(REPO)),
        "n_child_models": n_models,
        "n_leaves": len(meta_out["leaves"]),
        "n_features": len(feats),
        "checksums": {
            f: sha256(OUT_DIR / f)
            for f in ("meta.json", "features.json", "child_models.npz")
        },
    }
    (OUT_DIR / "MANIFEST.json").write_text(json.dumps(manifest, indent=2))

    print(f"Wrote {n_models} child models + meta + {len(feats)} genes to {OUT_DIR}")
    for f in ("meta.json", "features.json", "child_models.npz", "MANIFEST.json"):
        print(f"  {f}: {(OUT_DIR / f).stat().st_size/1024:.1f} KB")


if __name__ == "__main__":
    main()