File size: 16,988 Bytes
09cb542
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
# /// 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()