#!/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()