Download export_vae.py from coreai-community/Moebius-CoreAI: direct link, hf CLI and curl.
- Browser
- Download file 4.77 kB
-
https://huggingface.co/coreai-community/Moebius-CoreAI/resolve/main/export_vae.py
- Command line
-
hf download hf://coreai-community/Moebius-CoreAI/export_vae.py
-
curl -L -o export_vae.py https://huggingface.co/coreai-community/Moebius-CoreAI/resolve/main/export_vae.py
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() | |