Moebius-CoreAI / export_vae.py
xocialize's picture
Moebius 0.22B diffusion inpainting fp16/fp32 — first diffusion pipeline in coreai-community
09cb542 verified
Raw
History Blame Contribute Delete
4.77 kB
# /// 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()