File size: 4,766 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
# /// script
# requires-python = ">=3.11"
# dependencies = [
#     "coreai-core==1.0.0b2",
#     "coreai-torch==0.4.1",
#     "diffusers",
#     "numpy",
# ]
#
# [tool.uv]
# index-url       = "https://pypi.org/simple"
# prerelease      = "allow"
# index-strategy  = "unsafe-best-match"
# ///
"""Export the Moebius VAE (AutoencoderKL, KL-f8) to CoreAI .aimodel assets.

Two assets, shaped for the pipeline's exact call pattern:
  * encoder, batch 2, [2,3,512,512] -> posterior MEAN [2,4,64,64]
    (one forward encodes image + masked_image together, as the pipeline does; the mean is the
    deterministic moment the oracle/MLX ports gate on — no sampling in the graph)
  * decoder, batch 1, [1,4,64,64] -> [1,3,512,512]

scaling_factor stays OUT of the graph (host-side scalar), matching oracle semantics.

`patch_nearest_upsample` is load-bearing here: the decoder carries 3 nearest-x2 Upsample2D
modules, exactly the op MPSGraph's segmenter rejects (routes to BNNS/CPU) — same fix as the UNet.

Run:  uv run coreai/export_vae.py
"""
import shutil
import time
from pathlib import Path

import torch

ROOT = Path(__file__).resolve().parent.parent
VAE_DIR = ROOT / "weights/PixelHacker/vae"
OUT = ROOT / "coreai/exports"


def patch_nearest_upsample(module: torch.nn.Module) -> int:
    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


class EncoderMean(torch.nn.Module):
    """image [b,3,512,512] -> posterior mean [b,4,64,64] (deterministic; sf applied host-side)."""

    def __init__(self, vae):
        super().__init__()
        self.encoder = vae.encoder
        self.quant_conv = vae.quant_conv

    def forward(self, image):
        moments = self.quant_conv(self.encoder(image))
        mean, _logvar = moments.chunk(2, dim=1)
        return mean


class Decoder(torch.nn.Module):
    """latents [b,4,64,64] (UNSCALED — divide by sf host-side first) -> image [b,3,512,512]."""

    def __init__(self, vae):
        super().__init__()
        self.post_quant_conv = vae.post_quant_conv
        self.decoder = vae.decoder

    def forward(self, latents):
        return self.decoder(self.post_quant_conv(latents))


def export(wrapper, example, name: str, dtype=torch.float16) -> None:
    from coreai_torch import TorchConverter, get_decomp_table

    # ⚠️ Eager sanity runs at fp32: torch's CPU fp16 conv path is `slow_conv2d` and a single
    # 512² encoder forward at fp16 ground for 20+ CPU-MINUTES before being killed (the
    # quantized-forward-on-CPU trap family). torch.export itself traces with fake tensors —
    # no real compute — so only this sanity call ever executes kernels.
    wrapper = wrapper.eval()
    with torch.no_grad():
        out = wrapper(*example)
    print(f"[export] {name}: eager fp32 ok {tuple(example[0].shape)} -> {tuple(out.shape)}")
    wrapper = wrapper.to(dtype)
    example = tuple(t.to(dtype) for t in example)

    started = time.time()
    ep = torch.export.export(wrapper, args=example)
    ep = ep.run_decompositions(get_decomp_table())
    program = (TorchConverter()
               .add_exported_program(ep, input_names=["x"], output_names=["out"])
               .to_coreai())
    program.optimize()
    path = OUT / f"{name}.aimodel"
    if path.exists():
        shutil.rmtree(path)
    program.save_asset(path)
    size = sum(f.stat().st_size for f in path.rglob("*") if f.is_file()) / 1e6
    print(f"[export] saved {path.name} ({size:.0f} MB, {time.time() - started:.1f}s)")


def main() -> None:
    from diffusers.models import AutoencoderKL

    vae = AutoencoderKL.from_pretrained(str(VAE_DIR)).eval()
    print(f"[export] vae scaling_factor={vae.config.scaling_factor}")
    n = patch_nearest_upsample(vae)
    print(f"[export] patched {n} Upsample2D module(s) -> repeat_interleave")

    # Encoder ships fp32: at fp16 it reads 45.6 dB (investigate) and produces NaN on the CPU
    # lane — the classic SD-VAE fp16 activation-range problem, and mixed precision does not
    # lower in a CoreAI graph (measured). One encode per image makes fp32's ~2x cost invisible.
    export(EncoderMean(vae), (torch.randn(2, 3, 512, 512),), "moebius-vae-encoder-fp32-b2",
           dtype=torch.float32)
    # Decoder ships fp16: 68.5 dB [PASS] vs the shared golden.
    export(Decoder(vae), (torch.randn(1, 4, 64, 64),), "moebius-vae-decoder-fp16-b1")


if __name__ == "__main__":
    main()