Download export_unet.py from coreai-community/Moebius-CoreAI: direct link, hf CLI and curl.
- Browser
- Download file 17 kB
-
https://huggingface.co/coreai-community/Moebius-CoreAI/resolve/main/export_unet.py
- Command line
-
hf download hf://coreai-community/Moebius-CoreAI/export_unet.py
-
curl -L -o export_unet.py https://huggingface.co/coreai-community/Moebius-CoreAI/resolve/main/export_unet.py
17 kB
| # /// script | |
| # requires-python = ">=3.11" | |
| # dependencies = [ | |
| # "coreai-core==1.0.0b2", | |
| # "coreai-torch==0.4.1", | |
| # "diffusers", | |
| # "timm", | |
| # "einops", | |
| # "pyyaml", | |
| # "numpy", | |
| # ] | |
| # | |
| # [tool.uv] | |
| # index-url = "https://pypi.org/simple" | |
| # prerelease = "allow" | |
| # index-strategy = "unsafe-best-match" | |
| # /// | |
| """Export the Moebius UNet to a CoreAI .aimodel. | |
| WHY UNET-ONLY: the UNet is 38 of the 40 forwards per image and it IS the hypothesis under test | |
| (depthwise-separable + MBConv + linear attention on ANE vs Metal — memory `mlx-no-grouped-conv3d`). | |
| The VAE is 2 calls and does not move the measurement; it can follow using coreai-models' existing | |
| VAEEncoder/VAEDecoder wrappers if the answer is favourable. | |
| STATIC SHAPES throughout — required for ANE residency, and free here: Moebius is structurally | |
| locked to 512² (spatially-baked `rel_pos_emb` + a √n reshape in the attention wrapper), so the | |
| usual static-shape constraint costs nothing. | |
| Run: uv run coreai/export_unet.py --dtype fp16 | |
| """ | |
| import argparse | |
| import importlib | |
| import shutil | |
| import sys | |
| import time | |
| import types | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| ROOT = Path(__file__).resolve().parent.parent | |
| REF = ROOT / "reference" | |
| sys.path.insert(0, str(REF)) | |
| CKPT = ROOT / "weights/Moebius/ft_places2/diffusion_pytorch_model.bin" | |
| CFG = REF / "config/model_cfg/moebius.yaml" | |
| NUM_EMBEDDINGS = 20 | |
| def load_unet(): | |
| """The reference UNet, without executing `model_lib/__init__.py` (it eagerly imports a GLA | |
| variant needing flash-linear-attention — CUDA-first and unused by Moebius).""" | |
| for name, path in [ | |
| ("model_lib", REF / "model_lib"), | |
| ("model_lib.nets", REF / "model_lib/nets"), | |
| ("model_lib.nets.layers", REF / "model_lib/nets/layers"), | |
| ]: | |
| m = types.ModuleType(name) | |
| m.__path__ = [str(path)] | |
| sys.modules[name] = m | |
| mod = importlib.import_module("model_lib.nets.unet_lambda_prune_lite") | |
| cfg = yaml.safe_load(CFG.read_text()) | |
| model_cfg = dict(cfg["model"]) | |
| model_type = model_cfg.pop("model_type") | |
| model_cfg["sample_size"] = cfg["data"]["image_size"] // cfg["vae"]["downsample_ratio"] | |
| model_cfg["num_embeddings"] = NUM_EMBEDDINGS | |
| net = getattr(mod, model_type)(**model_cfg) | |
| sd = torch.load(CKPT, map_location="cpu", weights_only=True) | |
| # The checkpoint is the RemovalModel state dict: `diff_model.*` + `embedding_layer.weight`. | |
| unet_sd = {k[len("diff_model."):]: v for k, v in sd.items() if k.startswith("diff_model.")} | |
| missing, unexpected = net.load_state_dict(unet_sd, strict=True) | |
| print(f"[export] unet load: missing={len(missing)} unexpected={len(unexpected)}") | |
| net.eval() # the 124 BatchNorms must use running statistics | |
| embedding = sd["embedding_layer.weight"] # [20, 3072] | |
| return net, embedding | |
| def patch_nearest_upsample(module: torch.nn.Module) -> int: | |
| """Replace nearest-neighbour interpolate with repeat_interleave in `Upsample2D`. | |
| LIFTED FROM coreai-models (`diffusion/components.py::_patch_nearest_upsample`) — and it is | |
| load-bearing, not cosmetic: MPSGraph's segmenter REJECTS `coreai.interpolate` with | |
| nearest_neighbor mode and routes those ops to the BNNS (CPU) backend. That both breaks | |
| single-backend execution and inserts GPU→CPU→GPU copies at every upsample boundary. Exporting | |
| without this yields a graph that quietly falls off the accelerator — and then a benchmark that | |
| measures the wrong thing. | |
| `repeat_interleave` is mathematically identical for integer scale factors. | |
| """ | |
| from diffusers.models.upsampling import Upsample2D | |
| patched = 0 | |
| for mod in module.modules(): | |
| if isinstance(mod, Upsample2D) and not mod.use_conv_transpose: | |
| def _forward(hidden_states, output_size=None, _mod=mod): | |
| h = hidden_states.repeat_interleave(2, dim=-2).repeat_interleave(2, dim=-1) | |
| return _mod.conv(h) | |
| mod.forward = _forward | |
| patched += 1 | |
| return patched | |
| def patch_lambda_einsums() -> None: | |
| """Rewrite the two λ positional einsums to rank-≤4 matmul form, for ANE eligibility. | |
| WHY (measured 2026-08-01): requesting `neuralEngine` on the unpatched export fails to compile — | |
| 17× `MPS-ANEC conversion failure: mps.reshape input/output rank 6 exceeds the max rank 5`, all | |
| from `vanillaλ.py:146-147`, then `_ANECompiler: ANECCompile() FAILED`. torch.export decomposes | |
| `einsum('n m k u, b u v m -> b n k v')` (six distinct indices) through rank-6 reshapes, and | |
| **ANE's maximum tensor rank is 5**. The GPU delegate doesn't care; the ANE hard-rejects it. | |
| Both equations fold to plain (batched) matmuls with NO change in value — same trick the MLX | |
| port's `applyPositionalLambda` uses for memory reasons. One structural quirk, two backends, | |
| two different symptoms. | |
| SEAM: `_einsum` is a module-level lambda in `layers/utils.py`, but `vanillaλ.py` binds the NAME | |
| at import (`from ..utils import _einsum`), so patching utils after the fact would be a no-op. | |
| Rebinding the vanillaλ module global covers all four call sites (self- and cross-lambda) in one | |
| move. Dispatch on the equation string; everything else falls through to the original — the | |
| remaining λ einsums are rank ≤ 4 already and drew no validation warnings. | |
| The export flow numerically gates this patch (fp32 eager, pre- vs post-patch) before casting. | |
| """ | |
| vλ = importlib.import_module("model_lib.nets.layers.λ.vanillaλ") | |
| original = vλ._einsum | |
| def _patched(eq, *ops): | |
| if eq == 'n m k u, b u v m -> b n k v': | |
| # Broadcast-matmul form: [1,N,K,MU] @ [B,1,MU,Vd] → [B,N,K,Vd]. The earlier | |
| # [NK, MU]-flattened form put N·K on one axis — 65536 at the 64² level, past the | |
| # ANE's per-axis limit; this keeps every axis ≤ max(N, MU, K, Vd). | |
| rel, V = ops # [N,M,K,U], [B,U,Vd,M] | |
| N, M, K, U = rel.shape | |
| B, _, Vd, _ = V.shape | |
| A = rel.permute(0, 2, 1, 3).reshape(1, N, K, M * U) | |
| Bm = V.permute(0, 3, 1, 2).reshape(B, 1, M * U, Vd) | |
| return (A @ Bm).contiguous() # [B,N,K,Vd] | |
| if eq == 'b h k n, b n k v -> b h v n': | |
| Q, lam = ops # [B,H,K,N], [B,N,K,Vd] | |
| Qbn = Q.permute(0, 3, 1, 2) # [B,N,H,K] | |
| Y = Qbn @ lam # [B,N,H,Vd] — batched, rank 4 | |
| return Y.permute(0, 2, 3, 1).contiguous() # [B,H,Vd,N] | |
| return original(eq, *ops) | |
| vλ._einsum = _patched | |
| def patch_self_lambda_forward() -> None: | |
| """Replace MultiQuerySelfLambda.forward with a rank-5-free, ANE-eligible formulation. | |
| WHY (stage-bisected, probe_ane_selflambda.py): the self-λ takes the LOCAL positional branch — | |
| `pos_conv = Conv3d(u, k, (1, r, r))` over V as [b,u,v,hh,ww]. The Conv3d itself compiles for | |
| ANE (s4a: OK) — but any reshape/flatten CONSUMING its rank-5 output does not (s4e/s4f: FAIL; | |
| s4d, the same matmul fed rank-4 tensors: OK). The fix never materialises rank 5: with u=1 and | |
| depth-kernel 1, the Conv3d IS a Conv2d over each v-slice, so fold v into the conv batch and | |
| land the output directly in matmul layout. The positional application then runs as a batched | |
| matmul over n (the same rewrite as the MLX port's `applyPositionalLambda` — third appearance | |
| of this contraction, third backend-specific formulation). | |
| Numerically gated by the export's fp32 pre/post-patch eager comparison, same as the einsums. | |
| """ | |
| import torch.nn.functional as F | |
| vλ = importlib.import_module("model_lib.nets.layers.λ.vanillaλ") | |
| def forward(self, x): # x: [b, hh, ww, c] | |
| b, hh, ww, _ = x.shape | |
| n = hh * ww | |
| xc = x.permute(0, 3, 1, 2) # 'b h w c -> b c h w' | |
| q = self.to_q(xc) | |
| k = self.to_k(xc) | |
| v = self.to_v(xc) | |
| Q = self.norm_q(q) | |
| V = self.norm_v(v) | |
| h, u = self.heads, self.u | |
| dk = q.shape[1] // h | |
| dv = V.shape[1] // u | |
| Q = Q.reshape(b, h, dk, n) | |
| k = k.reshape(b, u, dk, n).softmax(dim=-1) | |
| V = V.reshape(b, u, dv, n) | |
| lam_c = torch.einsum('b u k m, b u v m -> b k v', k, V) | |
| Yc = torch.einsum('b h k n, b k v -> b h v n', Q, lam_c) | |
| assert self.local_contexts and u == 1 and self.pos_conv.weight.shape[2] == 1, \ | |
| "rank-5-free fold assumes the local branch with u=1 and depth-kernel 1" | |
| w2d = self.pos_conv.weight.squeeze(2) # [k, u, r, r] | |
| Vb = V.reshape(b * dv, u, hh, ww) # u=1: ONE rank-4 reshape, no rank-5 | |
| lam = F.conv2d(Vb, w2d, self.pos_conv.bias, padding=self.pos_conv.padding[1]) | |
| lam = lam.reshape(b, dv, dk, n).permute(0, 3, 2, 1) # [b,n,k,v] | |
| Yp = (Q.permute(0, 3, 1, 2) @ lam).permute(0, 2, 3, 1) # [b,h,v,n] | |
| Y = Yc + Yp | |
| out = Y.reshape(b, h * dv, n).permute(0, 2, 1) # 'b h v (hh ww) -> b (hh ww) c' | |
| return out.reshape(b, hh, ww, h * dv) # module contract: 'b h w c' | |
| vλ.MultiQuerySelfLambda.forward = forward | |
| class PrecomputedBN(torch.nn.Module): | |
| """BatchNorm replaced by per-channel scale/shift, constants computed at fp64 THEN cast. | |
| WHY: the fp16 export sits at 41.4 dB vs the golden while MLX fp16 manages rel 9.3e-04 on the | |
| same checkpoint. 25 of the 124 running_var tensors are below fp16's min-normal; evaluating | |
| (x-mean)·rsqrt(var+eps) in fp16 arithmetic mangles those channels. The COMPOSITE constants | |
| scale = γ/√(var+ε) and shift = β − mean·scale are fp16-representable even where var is not | |
| (γ/√(8e-07) ≈ 3000γ ≪ 65504), so fold the four tensors into two at full precision first. | |
| Numerically this is the same inference function — only the evaluation order changes. | |
| """ | |
| def __init__(self, bn: torch.nn.Module, spatial: bool): | |
| super().__init__() | |
| var = bn.running_var.data.double() | |
| mean = bn.running_mean.data.double() | |
| gamma = bn.weight.data.double() | |
| beta = bn.bias.data.double() | |
| scale = gamma / torch.sqrt(var + bn.eps) | |
| shift = beta - mean * scale | |
| shape = (1, -1, 1, 1) if spatial else (1, -1, 1) | |
| self.register_buffer("scale", scale.float().reshape(shape)) | |
| self.register_buffer("shift", shift.float().reshape(shape)) | |
| # ⚠️ 32 of the 124 "BatchNorms" are timm BatchNormAct2d — a subclass whose forward | |
| # appends drop + activation (ReLU here). isinstance(BatchNorm2d) matches it, and a | |
| # replacement that drops the activation diverges by rel ~1.0. The numeric gate caught | |
| # this; carry the epilogue through. | |
| self.act = getattr(bn, "act", None) or torch.nn.Identity() | |
| self.drop = getattr(bn, "drop", None) or torch.nn.Identity() | |
| def forward(self, x): | |
| return self.act(self.drop(x * self.scale + self.shift)) | |
| def patch_batchnorms_precomputed(module: torch.nn.Module) -> int: | |
| replaced = 0 | |
| for parent in module.modules(): | |
| for name, child in list(parent.named_children()): | |
| if isinstance(child, (torch.nn.BatchNorm2d, torch.nn.BatchNorm1d)): | |
| setattr(parent, name, | |
| PrecomputedBN(child, spatial=isinstance(child, torch.nn.BatchNorm2d))) | |
| replaced += 1 | |
| return replaced | |
| class MoebiusUNetWrapper(torch.nn.Module): | |
| """Export surface: `(sample, timestep, encoder_hidden_states) -> noise prediction`. | |
| The 20×3072 category table is deliberately left OUTSIDE the graph. Its lookup is a constant | |
| gather (CFG always indexes rows 10–19 then 0–9), so the projected conditioning is identical on | |
| every call — feeding it as an input keeps the graph free of an int64 embedding op, which is | |
| friendlier to the accelerator, and lets the host hoist the lookup out of the 19-step loop | |
| entirely. | |
| """ | |
| def __init__(self, unet: torch.nn.Module) -> None: | |
| super().__init__() | |
| self.model = unet | |
| n = patch_nearest_upsample(self.model) | |
| print(f"[export] patched {n} Upsample2D module(s) → repeat_interleave") | |
| def forward(self, sample, timestep, encoder_hidden_states): | |
| return self.model(sample, timestep=timestep, | |
| encoder_hidden_states=encoder_hidden_states).sample | |
| def main() -> None: | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--dtype", default="fp16", choices=["fp16", "fp32"]) | |
| ap.add_argument("--batch", type=int, default=2, help="2 = CFG-doubled, the production shape") | |
| ap.add_argument("--out", default=str(ROOT / "coreai/exports")) | |
| args = ap.parse_args() | |
| from coreai_torch import TorchConverter, get_decomp_table | |
| net, embedding = load_unet() | |
| wrapper = MoebiusUNetWrapper(net).eval() | |
| # Gate the λ einsum rewrite numerically BEFORE any cast: fp32 eager, pre- vs post-patch. | |
| # The rewrite is algebraically exact; this catches a transcription slip, not a design flaw. | |
| b = args.batch | |
| torch.manual_seed(0) | |
| probe = (torch.randn(b, 9, 64, 64), torch.full((b,), 900, dtype=torch.float32), | |
| torch.randn(b, 10, 3072)) | |
| with torch.no_grad(): | |
| pre = wrapper(*probe) | |
| patch_lambda_einsums() | |
| patch_self_lambda_forward() | |
| n_bn = patch_batchnorms_precomputed(wrapper) | |
| print(f"[export] replaced {n_bn} BatchNorms with fp64-precomputed scale/shift") | |
| with torch.no_grad(): | |
| post = wrapper(*probe) | |
| gap = (pre - post).abs().max().item() / (pre.abs().max().item() + 1e-12) | |
| print(f"[export] λ einsum rewrite gate: rel {gap:.3e} (fp32 eager, pre vs post)") | |
| if gap > 1e-5: | |
| raise SystemExit("[export] λ rewrite diverged from the original — refusing to export.") | |
| dtype = torch.float16 if args.dtype == "fp16" else torch.float32 | |
| if dtype == torch.float16: | |
| # UNIFORM fp16 — including BatchNorm statistics, which DIFFERS from the MLX side. | |
| # | |
| # convert_weights.py pins BN running stats to fp32 as a precaution against a running_var | |
| # rounding toward zero (rsqrt then explodes). That precaution is free on MLX. Here it is | |
| # not: mixing fp32 BatchNorm into an fp16 graph makes the lowering fail outright — | |
| # "failed to legalize unresolved materialization from tensor<*xf32> to | |
| # tensor<2x1280x16x16xf16>" inside the λ cross-attention, because norm_q/norm_v emit | |
| # fp32 into fp16 einsums and PyTorch's silent promotion has no lowering equivalent. | |
| # | |
| # So the precaution was MEASURED rather than carried over: across all 124 running_var | |
| # tensors the global minimum is 8.281e-07 — subnormal at fp16 but representable, and | |
| # ZERO tensors round to zero. Even under flush-to-zero the result is bounded by | |
| # eps (1/sqrt(1e-5) = 316), not infinite. Uniform fp16 is safe for THIS checkpoint; | |
| # re-measure for any sibling before assuming it transfers. | |
| wrapper = wrapper.half() | |
| sample = torch.randn(b, 9, 64, 64, dtype=dtype) | |
| timestep = torch.full((b,), 900, dtype=torch.float32) | |
| context = torch.randn(b, 10, 3072, dtype=dtype) | |
| print(f"[export] tracing — sample{tuple(sample.shape)} t{tuple(timestep.shape)} " | |
| f"ctx{tuple(context.shape)} dtype={args.dtype}") | |
| with torch.no_grad(): | |
| reference = wrapper(sample, timestep, context) | |
| print(f"[export] eager forward ok → {tuple(reference.shape)}") | |
| started = time.time() | |
| ep = torch.export.export(wrapper, args=(sample, timestep, context)) | |
| ep = ep.run_decompositions(get_decomp_table()) | |
| print(f"[export] torch.export + decompositions: {time.time() - started:.1f}s") | |
| started = time.time() | |
| program = ( | |
| TorchConverter() | |
| .add_exported_program( | |
| ep, | |
| input_names=["sample", "timestep", "encoder_hidden_states"], | |
| output_names=["noise_pred"], | |
| ) | |
| .to_coreai() | |
| ) | |
| program.optimize() | |
| print(f"[export] to_coreai + optimize: {time.time() - started:.1f}s") | |
| out = Path(args.out) / f"moebius-unet-{args.dtype}-b{b}.aimodel" | |
| out.parent.mkdir(parents=True, exist_ok=True) | |
| if out.exists(): | |
| shutil.rmtree(out) | |
| program.save_asset(out) # wants a Path, not a str | |
| size = sum(f.stat().st_size for f in out.rglob("*") if f.is_file()) / 1e6 | |
| print(f"[export] saved {out} ({size:.0f} MB)") | |
| # The conditioning is constant — bake it next to the asset so the runtime never recomputes it. | |
| np.save(Path(args.out) / "embedding_table.npy", embedding.float().numpy()) | |
| print(f"[export] wrote embedding_table.npy {tuple(embedding.shape)} (host-side constant gather)") | |
| if __name__ == "__main__": | |
| main() | |