hmoe-da / scripts /convert_weights.py
elanschonfeld's picture
initial open-weights release: package + weights + model card
a974167 verified
Raw History Blame
4.98 kB
#!/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()