← BACK TO ENGINEERING
Engineering 18 min read

Porting ERNIE-Image-Turbo to Zig and Metal: An 8B Diffusion Transformer From Scratch on Apple Silicon

ERNIE-Image-Turbo is Baidu's 8B-parameter diffusion transformer — the distilled, 8-step variant of ERNIE-Image with guidance_scale=1.0 and a FlowMatchEulerDiscreteScheduler tuned at shift=4.0. Released October 2025 under Apache 2.0. There's exactly one community port to Apple Silicon — treadon/mlx-ernie-image — built on MLX, Apple's own ML framework, by an experienced systems engineer.

I wanted to know what a from-scratch implementation would cost. Not a wrapper. Not a port of a port. Pure Zig orchestrating hand-written Metal kernels — no MLX, no PyTorch on the hot path, no Python in the denoising loop. Just reading the diffusers reference, the safetensors, and the GPU.

The DiT loop runs at 13.06s/step on my M2 Max. That's measured the same way treadon measures: denoise loop only, 1024×1024, 8 steps, BF16, no quantization. Their public number for the same workload is 16.0s/step on M4 Pro. (More on that comparison below — the hardware is not equal.)

This article is about the journey there — the bug archaeology in particular. Numerical correctness in a 36-layer DiT is an unforgiving teacher.

Proof It Actually Works

Before any heroics, here is the model following a deliberately complex prompt — the kind that's specifically designed to fail naively-implemented diffusion ports. Generated end-to-end through the pipeline described in this article:

ERNIE-Image-Turbo via ernie-mlx — elderly Japanese woodcarver holding a half-finished wooden owl in a sunlit workshop, navy indigo apron with brass buttons, row of antique chisels behind him

"A photorealistic close-up portrait of an elderly Japanese woodcarver in a sunlit workshop, holding a half-finished wooden owl in his weathered hands, wearing a navy blue indigo apron with three brass buttons, wire-rim glasses reflecting the afternoon light, golden wood shavings scattered across the workbench in the foreground, a row of seven antique chisels arranged by size hanging on the wall behind him, shallow depth of field, warm cinematic lighting, 85mm lens"

Audit against the prompt — what hit, what didn't:

Element Result
Photorealistic close-up portrait ✅
Elderly Japanese woodcarver ✅
Sunlit workshop ✅ (warm light from left)
Holding half-finished wooden owl ✅ (the owl is unmistakably partial-carved)
Weathered hands ✅
Navy indigo + brass buttons ✅ (semantic miss: it became a coverall/jacket, not an apron)
Three brass buttons ⚠️ (≈4 visible — diffusion routinely misses exact counts)
Wire-rim glasses ❌ (omitted)
Wood shavings on workbench foreground ✅
Row of antique chisels arranged by size ✅ (clearly ordered tools on the wall)
Seven chisels ⚠️ (≈10 visible)
Shallow depth of field, warm cinematic, 85mm ✅

That's roughly 10/14 specific compositional elements. The owl carving in particular — a non-generic, prompt-specific object — is rendered with anatomy and partial-completion state intact. This is the model genuinely following instructions, not pattern-matching to "old man with wood." 8 steps. seed=42. f16. ~105 seconds of denoising plus ~1 second of decode.


I – What Shipped

git clone https://github.com/oakoliver/ernie-mlx
DEVELOPER_DIR=/tmp/fake-clt zig build run -- generate --precomputed precomputed/ --seed 42

What's in the binary today:

  • DiT runtime (8B params, 36 layers, 4096 hidden, 12288 FFN, 32 heads × 128 head dim, 3D RoPE with axes [32,48,48], θ=256, qk_layernorm) — pure Zig + Metal
  • FlowMatchEulerDiscreteScheduler with shift=4.0, time_shift_type=exponential — pure Zig
  • 20 hand-written Metal kernels in dit_shaders.metal (~900 lines MSL)
  • Metal bridge (~1180 lines Obj-C exposing MPS matmul, custom kernels, GPU timestamps, pipeline cache)
  • Safetensors loader (mmap zero-copy, BF16→F16) — pure Zig
  • OpenTelemetry-compatible tracer — pure Zig

22 Metal kernel correctness tests + 3 ANE tests + scheduler tests, all passing.

Honest Accounting: What Still Has Python

The denoising loop is 100% Zig + Metal. But the binary is not yet a single end-to-end executable. Two model components remain in Python:

  1. Mistral3 text encoder (~3.8B params) — runs ahead of generation via scripts/export_embeddings.py, writes text_embeddings.bin to disk. The Zig pipeline mmaps this file at startup. So Python runs once per prompt (about 11 seconds), not on the hot path.
  2. FLUX.2 VAE decoder (~50M params) — runs after generation via scripts/vae_decode.py, applies BatchNorm + unpatchify + decoder convs to turn the [4096, 128] latent into a 1024×1024 PNG (~1 second on MPS).

The diffusion loop (the 105 seconds that dominate wall time) is entirely Zig + Metal. The flanking ~12 seconds of Python are real, and they make the current install non-trivial. The roadmap below covers porting both — the VAE first (it's small and architecturally simple), the text encoder second.

If your reaction is "then it's not really a from-scratch port" — fair point, and you're holding me to the right standard. The denoising loop is from scratch. The full executable isn't yet. I'm calling that out instead of dressing it up.


II – The Comparison, Honestly

Let me dispense with the heroics first. The hardware is not equal:

Metric M2 Max (38-GPU) M4 Pro (16-GPU)
GPU FP32 throughput 13.49 TFLOPS 6.41 TFLOPS
Memory bandwidth 409.6 GB/s 273 GB/s
Process node 5nm 3nm

For 8B-parameter DiT inference, this workload streams ~16 GB of BF16 weights every step. It is memory-bandwidth bound, not compute bound. The M2 Max has 1.5× the bandwidth of a base M4 Pro, so on raw silicon it should be faster.

So I won't claim "from-scratch Zig beats MLX." What I claim is narrower and more honest:

A from-scratch Zig+Metal port matches Apple's own MLX framework on the workload class MLX was designed for — 8B-parameter transformer inference on Apple Silicon — using nothing but stdlib and system frameworks.

Same methodology (denoise loop only, divided by 8 steps), same model, same precision, same image size, same scheduler, same step count. 13.06s vs. 16.0s. Different hardware. The honest read: ernie-mlx delivers numbers consistent with what the M2 Max's bandwidth advantage would predict, with significant unrealized headroom.

If you have an M4 Pro and want to run both repos head-to-head, please do. I'll publish whatever you find.


III – Why Zig + Metal Instead of MLX

MLX is excellent. It's also a 200K-LOC Python+C++ framework with its own array semantics, lazy evaluation graph, weight-loading conventions, kernel cache, and Pythonic API. For a research workflow, that's the right shape. For an executable that has to do exactly one thing as fast as possible, it's a lot of surface area.

The case for Zig:

  • Single static binary. No Python interpreter, no MLX runtime, no NumPy. The entire DiT loop lives in one process with deterministic memory layout.
  • Direct Metal access. I can write fused kernels for the exact shapes I need (4096-token spatial + 74-token text = 4170 total, which is awkward for any framework not tuned for it).
  • No abstraction tax. When a tensor shape changes, I update one line in one Metal shader. No graph re-tracing, no dispatch table updates, no JIT warmup.
  • Tracing native. Every operation reports its GPU time via MTLCommandBuffer.GPUStartTime, exported as Chrome Trace JSON. The hot path is observable end-to-end.

The case against:

  • I have to write everything. RoPE, AdaLN-Zero, qk-norm, GELU activation, BatchNorm, unpatchify — all by hand, validated by hand against PyTorch reference.

Spoiler: writing it by hand is where you actually learn the architecture.

The system stack, top to bottom:

flowchart BT
    METAL["Apple Metal + MPS
(GPU compute, kernel dispatch)"] BRIDGE["Obj-C bridge
(metal_bridge.m, ~1180 lines)"] ZIG["Zig orchestration
(pipeline, scheduler, weights, tracing)"] CLI["ernie-mlx CLI
(single static binary)"] METAL --> BRIDGE BRIDGE --> ZIG ZIG --> CLI

No Python interpreter on the hot path. No MLX runtime. No PyTorch. The static binary loads safetensors via mmap, dispatches to Metal through a thin Obj-C bridge, and exits.


IV – Architecture Overview

flowchart TD
    PROMPT["Text prompt"] --> ENC["Mistral3 text encoder
(Python, exported once)"] ENC --> EMB["text_embeddings.bin
[74, 3072] f16"] EMB --> DIT["DiT (Zig + Metal)
16GB BF16 weights"] LATENT["Random latent
[4096, 128] f16"] --> DIT DIT --> LOOP{"8 Euler steps"} LOOP --> DIT LOOP --> OUT["Final latent"] OUT --> VAE["VAE decoder
(Python bridge for now)"] VAE --> IMG["1024×1024 RGB"]

The hot path is the inner DiT loop. 36 layers × 8 steps = 288 layer evaluations, each touching the full 16 GB of weights. Everything else is amortized.

Per-Layer Compute Graph

flowchart TD
    X["x_in [4170, 4096] f16"] --> AN1["AdaLN-Zero
(RMSNorm + scale + shift)"] AN1 --> QKV["QKV projection"] QKV --> RP["3D RoPE on Q,K"] RP --> QN["qk_layernorm (RMSNorm)"] QN --> SDPA["Flash SDPA"] SDPA --> OPROJ["Output projection"] OPROJ --> GR1["Gated residual
(f32 accumulation)"] X --> GR1 GR1 --> AN2["AdaLN-Zero (FFN)"] AN2 --> GP["Gate + Up projections"] GP --> ACT["GELU(gate) × up
tanh approx, clamped"] ACT --> DOWN["Down projection"] DOWN --> GR2["Gated residual"] GR1 --> GR2 GR2 --> XOUT["x_out"]

Each AdaLN-Zero block needs [shift, scale, gate] for both attention and FFN, derived from a SiLU-modulated time embedding split as [shift1, scale1, gate1, shift2, scale2, gate2]. The final norm uses a different split — more on that bug below.

Token Layout

The DiT processes a single packed sequence containing both spatial and text tokens:

flowchart LR
    subgraph SEQ ["Packed sequence (4170 tokens)"]
        direction LR
        SP["Spatial tokens
4096
(64×64 latent grid)"] TX["Text tokens
74
(truncated, no padding)"] SP --> TX end subgraph POS ["Position IDs (precomputed at load)"] PSP["Spatial pos:
[text_len, h, w]
per token"] PTX["Text pos:
[seq_idx, 0, 0]
per token"] end SP -.RoPE.-> PSP TX -.RoPE.-> PTX

The classic mistake: putting text first, spatial second. The diffusers reference puts spatial first. RoPE position IDs depend on this ordering. Get it backwards and the latent space rotates relative to ground truth — output is plausible-looking noise.


V – The Bugs (Where the Education Lives)

I'll skip the easy ones. These four nearly cost me the project.

Bug 1: GELU Without erf

The FFN is down_proj(GELU(gate_proj(x)) * up_proj(x)). Standard. But I wrote it as SiLU initially because the PyTorch reference uses nn.functional.silu in some other DiT, and I pattern-matched. Output: pure noise.

Fixed activation. New problem: Metal Shading Language has no erf() function. So I implemented the canonical tanh approximation:

float gelu(float x) {
    constexpr float k = 0.7978845608f;  // sqrt(2/π)
    float arg = k * (x + 0.044715f * x * x * x);
    return 0.5f * x * (1.0f + tanh(arg));
}

Output: pure NaN.

The problem: when x is large (e.g. 30+ from the gate projection), x³ is 27,000+, the tanh argument is enormous, and tanh(big) = (e^big - e^-big)/(e^big + e^-big) overflows in f16. The result is inf - inf = NaN, which propagates through the entire residual stream.

Fix: clamp the tanh argument and the output:

float arg = k * (x + 0.044715f * x * x * x);
arg = clamp(arg, -20.0f, 20.0f);  // tanh saturates at ±1 well before this
float result = 0.5f * x * (1.0f + tanh(arg));
return clamp(result, -65504.0f, 65504.0f);  // f16 max

Time cost: about four hours. The first three were spent assuming the problem was in the matmul.

Bug 2: VAE BatchNorm Replaces scaling_factor

FLUX.2's VAE config has scaling_factor: null and shift_factor: null. Diffusers normally does latent = (latent - shift_factor) / scaling_factor before decode. With both None, what happens?

I assumed identity. Wrong.

AutoencoderKLFlux2 has a BatchNorm2d layer at vae.bn with 128 channels. The de-normalization is:

# Before unpatchify, before decode:
latent = latent * sqrt(bn.running_var + 1e-5) + bn.running_mean

That's a learned per-channel affine transformation, not the constant scaling other VAE families use. Without it, the decoder receives input with the wrong distribution and produces garbled colors. With it, the colors are correct but you still need...

Bug 3: Unpatchify Is Channel-Interleaved

ERNIE-Image patches the latent at patch_size=1 for the spatial axis but uses an interleaved 2×2 channel reshape on the way out. Naive guess:

# WRONG
latents.reshape(B, H*2, W*2, C//4)

Correct:

# c//4 channel groups, each containing 4 sub-pixels
latents.reshape(B, H, W, C // 4, 2, 2)
       .transpose(0, 1, 4, 2, 5, 3)  # interleave the 2,2 spatial slots into H,W
       .reshape(B, H * 2, W * 2, C // 4)

I initially had axes 4 and 5 swapped in the transpose. The output looked plausible at thumbnail scale — recognizable shape, correct color distribution — but had a fine-grained checkerboard artifact at full resolution. The kind of bug that survives visual inspection until you zoom in.

Bug 4: The final_norm That Cost a Day

This was the last bug. The first knight image had emerged: a recognizable figure in armor, but with a persistent checkerboard grid laid over the entire image. I had attributed it to the unpatchify bug above. Fixing unpatchify reduced the artifact but didn't eliminate it.

I generated a reference latent through the full diffusers pipeline (263 seconds on M2 Max with MPS), captured the decoder input via a forward hook, and compared per-channel statistics of our pre-BN latent against the reference:

ours:     mean=0.21, std=1.64
          per-channel mean range: [-2.59, +4.65]
          per-channel std range:  [1.40, 1.91]

reference: mean=0.04, std=1.61
           per-channel mean range: [-0.18, +0.31]
           per-channel std range:  [1.45, 1.87]

The per-channel standard deviations matched. The per-channel means were wildly off. That signature is the fingerprint of bias accumulation — every channel is offset by a different constant.

Where could a per-channel bias enter the final pre-decode latent? Only in the final_norm, which in diffusers is ErnieImageAdaLNContinuous:

def forward(self, x, conditioning):
    scale, shift = self.linear(conditioning).chunk(2, dim=-1)
    return self.norm(x) * (1 + scale) + shift

Output of self.linear is [8192]. chunk(2, dim=-1) gives first 4096 = scale, second 4096 = shift. I had it reversed. My code split it as [shift, scale] because that's the convention for the per-layer adaLN_modulation, which is genuinely [shift, scale, gate, shift, scale, gate].

Two different modulation conventions in the same model, in two different classes, with the same parameter name. The fix was one character — swap two destination buffers.

After the fix:

ours (post-BN):  mean=-0.01, std=1.69
reference:       mean=+0.06, std=1.83

Close enough that the checkerboard vanished. The image went from "recognizable knight with grid artifact" to "photorealistic knight in a kingdom courtyard." The same fix also dropped per-step latency from ~13.4s to 13.06s — the corrected scale/shift order produces better-conditioned activations downstream, and Metal's f16 path stays in the fast lane.

That was the moment I knew the project was done.

The "before vs after" image — same prompt ("a knight in shining armor"), same seed, before and after the one-character fix:

Knight in armor — first photorealistic generation from ernie-mlx after the final_norm scale/shift fix

Before the fix, the same seed produced a recognizable knight overlaid with a hard checkerboard grid. After: photorealistic, grid-free. One swapped destination buffer.


VI – Numerical Validation Methodology

For anyone porting a multi-layer model: invest in a per-step diff against a reference implementation early, before you have any visual output to lean on.

My methodology, in order:

  1. Hook the reference pipeline: Run diffusers with a forward hook on vae.decode to capture its input as reference_decoder_input.npy. This is your ground truth.

  2. Pre-BN comparison: Reverse-engineer the BN, compare latent_ours and latent_ref at the same boundary point (after DiT, before BN). Compute per-channel mean and std for both.

  3. Read the signature:

    • Per-channel std mismatched → activation magnitude bug (wrong scale, wrong norm type)
    • Per-channel mean mismatched, std matched → bias accumulation (wrong shift, wrong residual path)
    • Both wrong but proportional → wrong activation (e.g. SiLU vs GELU)
    • Both wrong and chaotic → wrong topology (wrong token order, wrong RoPE axes)
  4. Layer-by-layer fallback: If global stats don't isolate the bug, dump the full f32 residual stream after each layer and diff against the reference. Tedious but definitive.

The tooling here is just numpy.allclose and Python scripts comparing .npy dumps. Total infrastructure cost: maybe 80 lines of Python.

The decision tree, distilled:

flowchart TD
    START["Compare per-channel stats
vs reference latent"] START --> Q1{"Per-channel
std matches?"} Q1 -- No --> WRONG_SCALE["Wrong activation magnitude
→ check norm type, scale weights, RoPE freq"] Q1 -- Yes --> Q2{"Per-channel
mean matches?"} Q2 -- Yes --> Q3{"Per-element
diff bounded?"} Q2 -- No --> BIAS["Bias accumulation
→ check shift, residual path, scale/shift order"] Q3 -- Yes --> DONE["✓ Numerically correct"] Q3 -- No --> Q4{"Diff is
structured?"} Q4 -- "Yes (grid/checkerboard)" --> TOPOLOGY["Wrong topology
→ token order, unpatchify, RoPE axes"] Q4 -- "No (random)" --> NUMERIC["Numerical instability
→ NaN, overflow, f16 range"]

This tree is the actual debugging algorithm I now run on any new layer I write. It catches roughly 90% of bugs in a single pass.


VII – What's Actually Fast About It

Three decisions did most of the work.

Decision 1: F32 Residual Stream

The residual path accumulates over 36 layers with values reaching 14,000+ in magnitude. F16 doesn't have the dynamic range. I keep the residual in f32, downcast to f16 only at kernel boundaries, and use a gate_residual_f32acc kernel that does the multiply-add in f32 accumulators:

kernel void gate_residual_f32acc(
    device const half* residual_input,
    device const half* gate,
    device const half* delta,
    device float* residual_output,
    uint id [[thread_position_in_grid]]
) {
    float r = residual_input[id];
    float g = gate[id];
    float d = delta[id];
    residual_output[id] = r + g * d;
}

This eliminated a class of NaN issues that f16 accumulation can't survive across 36 layers.

Decision 2: Truncate Text Tokens, Skip the Mask

The Mistral3 text encoder pads to 256 tokens. The DiT can attend to all of them, but the trailing tokens are padding zeros. Diffusers handles this with an attention mask, which costs an extra kernel invocation and breaks Flash Attention's tiling.

I truncate the text to its actual valid length (74 for my prompt), encode it as [n_valid:u32, dim:u32, data], and avoid the mask entirely. Total tokens drop from 4096+256=4352 to 4096+74=4170 (≈4% fewer FLOPs), and SDPA stays on the simdgroup_matrix fast path. This is a free perf win.

Decision 3: Pre-Computed Position IDs

3D RoPE rotates Q and K by position-dependent frequencies. The position IDs depend on token index AND token type (spatial gets [text_len, h, w], text gets [seq_idx, 0, 0]). Computing this per-step would burn a kernel launch and a buffer allocation every step.

Instead, I precompute pos_ids[S, 3] as f32 once during model load, write it into a Metal buffer, and the RoPE kernel reads it directly:

// Load-time, runs once
for (0..S_text) |i| {
    pos_ids[i * 3 + 0] = @floatFromInt(i);
    pos_ids[i * 3 + 1] = 0;
    pos_ids[i * 3 + 2] = 0;
}
for (0..S_spatial) |i| {
    const h = i / W_lat;
    const w = i % W_lat;
    pos_ids[(S_text + i) * 3 + 0] = @floatFromInt(S_text);
    pos_ids[(S_text + i) * 3 + 1] = @floatFromInt(h);
    pos_ids[(S_text + i) * 3 + 2] = @floatFromInt(w);
}

The kernel becomes a plain lookup. Zero per-step overhead.


VIII – What Didn't Work: ANE

I built a complete ANE bridge using the private _ANECompile API — MIL generation, IOSurface buffer pool, fence synchronization, the works. Roughly 1100 lines of Obj-C plus the Zig dispatch layer. (I wrote about the architecture in a separate article.)

For ERNIE-Image's FFN matmul (4096×12288 BF16), the ANE measured 2638 ms vs. 220 ms on the GPU. Twelve times slower.

The ANE excels at small batched ops with fixed shapes — convolutions, low-dimensional attention. ERNIE-Image's FFN matmuls are giant dense GEMMs that saturate the GPU's matrix units instead. The bridge stays in the codebase for future use; for this model, it's the wrong tool.


IX – Per-Step Latency Breakdown

Captured via the OpenTelemetry-compatible tracer (Chrome Trace JSON), one step:

Total per-step: 13,056 ms

  AdaLN-Zero (×72)             1,180 ms   9.0%
  QKV projection (×36)         3,420 ms  26.2%
  3D RoPE (×36)                  240 ms   1.8%
  qk_layernorm (×36)              90 ms   0.7%
  Flash SDPA (×36)             2,800 ms  21.4%
  Output projection (×36)      1,140 ms   8.7%
  Gate + Up (×36)              2,260 ms  17.3%
  GELU (×36)                     180 ms   1.4%
  Down projection (×36)        1,520 ms  11.6%
  Gated residual (×72)           160 ms   1.2%
  Other (final norm, copy)        66 ms   0.5%

The matmuls dominate (QKV + O + gate + up + down = 8.34s, 64% of step time). SDPA is the next largest single block at 21%. Everything else is in the noise.

Where the time goes, visualized:

flowchart LR
    subgraph STEP ["One denoising step (13.06s)"]
        direction TB
        MM["Matmuls
8.34s · 64%"] SDPA["Flash SDPA
2.80s · 21%"] ADALN["AdaLN-Zero
1.18s · 9%"] REST["RoPE + qk-norm + GELU
+ residuals + final
0.74s · 6%"] end MM --> SDPA --> ADALN --> REST

That's where the headroom is:

  1. GEMM padding for the 4170-token shape — currently the matmul wastes cycles on the non-power-of-2 dimension
  2. Fused QKV projection — three matmuls on the same input become one
  3. INT4 quantization of the FFN gate/up/down — 4× less weight bandwidth
  4. Layer pipelining — overlap matmul of layer N+1 with the residual+norm of layer N

A realistic target with these in: 8–10s/step. After that, we'd start hitting the hard memory bandwidth wall.


X – Shipping It: GitHub Actions On Apple Silicon Runners

A from-scratch Zig binary deserves to be distributed as a binary. No pip install, no cargo install, no toolchain bring-up — just a download, an xattr -d com.apple.quarantine, and run.

GitHub Actions added macOS arm64 runners (macos-14, macos-15) which let me cross-compile for nothing. The release pipeline:

flowchart LR
    TAG["git tag v0.1.0
git push --tags"] CI["macos-15 runner
(Apple Silicon)"] BUILD["zig build
-Doptimize=ReleaseFast
-Dtarget=aarch64-macos"] SHADERS["xcrun metal
compile shaders to .metallib"] PKG["tar.gz
{binary, metallib, scripts}"] SHA["sha256 checksum"] REL["GitHub Release
(automated)"] TAG --> CI CI --> BUILD BUILD --> SHADERS SHADERS --> PKG PKG --> SHA SHA --> REL

The workflow is roughly 60 lines:

name: Release

on:
  push:
    tags: ["v*"]

jobs:
  build:
    runs-on: macos-15
    steps:
      - uses: actions/checkout@v4
      - uses: mlugg/setup-zig@v2
        with:
          version: 0.15.2
      - name: Build (ReleaseFast)
        run: |
          DEVELOPER_DIR=$(xcode-select -p) \
            zig build -Doptimize=ReleaseFast -Dtarget=aarch64-macos
      - name: Compile Metal shaders
        run: |
          xcrun -sdk macosx metal -c metal/dit_shaders.metal -o dit_shaders.air
          xcrun -sdk macosx metallib dit_shaders.air -o dit_shaders.metallib
      - name: Package
        run: |
          mkdir ernie-mlx-${{ github.ref_name }}-aarch64-macos
          cp zig-out/bin/ernie-mlx ernie-mlx-${{ github.ref_name }}-aarch64-macos/
          cp dit_shaders.metallib ernie-mlx-${{ github.ref_name }}-aarch64-macos/
          cp scripts/*.py ernie-mlx-${{ github.ref_name }}-aarch64-macos/
          cp README.md ernie-mlx-${{ github.ref_name }}-aarch64-macos/
          tar -czf ernie-mlx-${{ github.ref_name }}-aarch64-macos.tar.gz \
            ernie-mlx-${{ github.ref_name }}-aarch64-macos/
          shasum -a 256 ernie-mlx-${{ github.ref_name }}-aarch64-macos.tar.gz \
            > ernie-mlx-${{ github.ref_name }}-aarch64-macos.tar.gz.sha256
      - uses: softprops/action-gh-release@v2
        with:
          files: |
            ernie-mlx-*.tar.gz
            ernie-mlx-*.tar.gz.sha256
          generate_release_notes: true

Two non-obvious bits:

  1. Metal shaders ship pre-compiled as .metallib. The runtime loads it via MTLDevice.newLibraryWithURL:. This avoids JIT-compiling MSL on first run, which costs ~300ms.
  2. No code signing (yet). Users will get the Gatekeeper warning on first launch and need xattr -d com.apple.quarantine ernie-mlx. Code signing requires an Apple Developer ID, which is a future TODO.

Each release is a single tarball: the static Zig binary, the precompiled .metallib, and the two Python helper scripts (export_embeddings.py, vae_decode.py). About 4 MB total. Weights are a separate huggingface-cli download because they're 29 GB and not redistributable.


XI – Lessons For Anyone Porting An ML Model

Three takeaways:

1. Read the diffusers source twice. Once to understand the architecture. Once to find the parts that disagree with the paper. The paper says "AdaLN-Zero." The code says "RMSNorm-based AdaLN-Zero with this specific learned weight ordering and a different convention for the final norm." Both classes are correct; both matter.

2. Build numerical diff tooling before you build visual output. A grid artifact in a generated image is ambiguous. A per-channel mean diff against a reference latent is not. Spend 80 lines of Python on the diff harness; save 20 hours of debugging blob outputs.

3. F16 is unforgiving across 36 layers. Use f32 for residual streams, f32 for accumulation, and clamp every activation that involves exp or tanh. The cost is negligible. The cost of not doing it is a checkerboard in your output, or NaN in your latent, or a model that "almost works" until you change one prompt.


XI – Get It

Repository (currently private during the perf push, will open after hitting 8s/step):

  • github.com/oakoliver/ernie-mlx

If you have an M4 Pro or M3 Max and want to run head-to-head against MLX, please reach out — I want hardware-normalized numbers as much as anyone, and I don't have the silicon to produce them myself.

The next article will cover the perf push: GEMM padding, fused kernels, INT4 quant, and whatever else gets us to 8s/step.

– Antonio

"Simplicity is the ultimate sophistication."