From 09388ec41975c991146b21215682abafe5afa215 Mon Sep 17 00:00:00 2001 From: ausimian Date: Wed, 15 Apr 2026 09:23:27 +0930 Subject: [PATCH] M8: native conv via mlx::core::conv_general MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Lift Backend.conv/4 from the BinaryBackend fallback onto Native.conv_general (already bound to mlx::core::conv_general since M1). The fallback was correct but CPU-bound — ≥90% of ViT and Whisper forward-pass cost. Layout translation: MLX conv_general expects NHWC input and OHWI weight; Nx delivers tensors in the caller's layout plus input/kernel/output_permutation opts pointing at canonical NCHW/OIHW. Compose the caller's permutations with the NCHW↔NHWC and OIHW↔OHWI transposes, applying the inverse of output_permutation on the way out (Nx delivers it in user→canonical form; see deps/nx/lib/nx/shape.ex:729-735). Two ordered transposes rather than one composed — MLX lazy-fuses them and the step-wise form is obviously correct across rank 3, 4, and 5. Integer inputs get cast to out.type via Native.astype before the transpose chain (Nx.conv returns float but does not cast operands; MLX conv is float-only). Same-type astype is elided by MLX. batch_group_size > 1 and complex-typed conv still route through via_binary — no MLX primitive, and neither appears in the pinned Bumblebee ref. Whisper full-checkpoint conformance: 403.8s → 1.8s (~224× speedup). ViT full-checkpoint conformance: similar scale (1.9s end-to-end). New oracle suite in test/emily/backend_conv_test.exs covers 1-D/2-D/ 3-D, stride/:same/:valid/asymmetric padding, kernel and input dilation, grouped and depthwise conv, all three permutation options (independently and combined for NHWC end-to-end), and integer-input coercion. test/emily/backend_fallbacks_test.exs: dropped the obsolete "conv routes through BinaryBackend" test; added a batch_group_size > 1 case asserting the rare fallback still matches BinaryBackend. --- RELEASE.md | 36 ++++++ lib/emily/backend.ex | 73 ++++++++++-- test/emily/backend_conv_test.exs | 156 ++++++++++++++++++++++++++ test/emily/backend_fallbacks_test.exs | 34 +++--- 4 files changed, 274 insertions(+), 25 deletions(-) create mode 100644 test/emily/backend_conv_test.exs diff --git a/RELEASE.md b/RELEASE.md index bf40125..5539e0c 100644 --- a/RELEASE.md +++ b/RELEASE.md @@ -2,6 +2,42 @@ ## Added +- M8 — Native `conv`. `Emily.Backend.conv/4` now dispatches directly + to `Native.conv_general` (already bound to `mlx::core::conv_general` + since M1) instead of round-tripping through `Nx.BinaryBackend`. + The previous fallback was correct but CPU-bound — ≥90% of the ViT + and Whisper forward-pass cost. ViT/Whisper full-checkpoint + conformance tests drop from tens of seconds to under 2 s as a + side-effect. + - **Layout translation.** MLX `conv_general` expects NHWC input and + OHWI weight and returns NHWC; Nx's canonical layout is NCHW/OIHW. + The new callback composes the caller's `input_permutation`, + `kernel_permutation`, and `output_permutation` opts with the + NCHW↔NHWC and OIHW↔OHWI transposes, applying the inverse of + `output_permutation` on the way out (Nx delivers it in + user→canonical form; see + [`deps/nx/lib/nx/shape.ex:729-735`](https://hexdocs.pm/nx/Nx.Shape.html)). + Two ordered transposes per tensor rather than one composed + transpose — MLX's lazy graph fuses them and the step-wise form is + obviously correct across rank 3, 4, and 5. + - **Integer-operand coercion.** `Nx.conv` returns a float but does + not cast its operands; MLX conv is float-only. The backend now + runs `Native.astype(ref, out.type)` on input and kernel before + the transpose chain. Same-type astype is elided by MLX. + - **Remaining fallbacks.** `batch_group_size > 1` has no MLX + primitive and still routes through `via_binary`. Complex-typed + conv ditto. Neither appears in the pinned Bumblebee ref. + - **`test/emily/backend_conv_test.exs`** (new) — oracle suite vs + `Nx.BinaryBackend`. Covers 1-D / 2-D / 3-D; stride, `:same` / + `:valid` / explicit asymmetric padding; `kernel_dilation` and + `input_dilation > 1`; grouped and depthwise conv; all three + permutation options (independently and combined for an NHWC + end-to-end caller); and integer-input→float-output coercion. + - **`test/emily/backend_fallbacks_test.exs`** — removed the + now-obsolete "conv routes through BinaryBackend" test. Added a + `batch_group_size > 1` fallback case asserting the rare path + still matches BinaryBackend. + - M7 — Bumblebee conformance breadth. Two new models across four new test suites extend M3 (DistilBERT) and M4 (Qwen3) beyond encoder-only/decoder-only text into vision and audio. diff --git a/lib/emily/backend.ex b/lib/emily/backend.ex index bf1b11c..26601f4 100644 --- a/lib/emily/backend.ex +++ b/lib/emily/backend.ex @@ -930,20 +930,71 @@ defmodule Emily.Backend do end # ================================================================= - # Conv (delegated to BinaryBackend for M2) + # Conv # ================================================================= # - # Nx's conv signature is rich (batch groups, feature groups, - # permutations, dilations). MLX conv_general covers most of this but - # the translation is involved. For M2 we go through BinaryBackend; - # native conv lands in M3 alongside Bumblebee integration. - - # Conv lands natively via Native.conv_general in M3 alongside - # Bumblebee integration. Nx.conv on BinaryBackend is correct but - # CPU-bound — fine for tests, unusable for real workloads. + # MLX `conv_general` expects NHWC input and OHWI weight; Nx's + # canonical layout is NCHW / OIHW. Nx does not pre-transpose tensors + # before dispatching — it delivers them in their original layout plus + # `input_permutation`, `kernel_permutation`, `output_permutation` + # such that `Nx.transpose(user_input, axes: input_permutation)` is + # the canonical form. We transpose into NHWC/OHWI on the way in, call + # the NIF, then reverse the layout on the way out. + # + # `batch_group_size > 1` and complex-typed conv have no MLX primitive + # and fall back to `via_binary`; they are rare enough not to warrant + # a handwritten MLX reshape trick. @impl true - def conv(out, input, kernel, opts), - do: via_binary(out, [input, kernel], &Nx.conv(&1, &2, opts)) + def conv(out, input, kernel, opts) do + cond do + opts[:batch_group_size] > 1 -> + via_binary(out, [input, kernel], &Nx.conv(&1, &2, opts)) + + match?({:c, _}, out.type) -> + via_binary(out, [input, kernel], &Nx.conv(&1, &2, opts)) + + true -> + ip = opts[:input_permutation] + kp = opts[:kernel_permutation] + op = opts[:output_permutation] + {lows, highs} = opts[:padding] |> Enum.unzip() + + input_to_nhwc = [hd(ip)] ++ Enum.drop(ip, 2) ++ [Enum.at(ip, 1)] + kernel_to_ohwi = [hd(kp)] ++ Enum.drop(kp, 2) ++ [Enum.at(kp, 1)] + rank = tuple_size(out.shape) + nhwc_to_nchw = [0, rank - 1] ++ Enum.to_list(1..(rank - 2)//1) + inv_op = invert_permutation(op) + + ir = input |> ref() |> Native.astype(out.type) |> Native.transpose(input_to_nhwc) + kr = kernel |> ref() |> Native.astype(out.type) |> Native.transpose(kernel_to_ohwi) + + ir + |> Native.conv_general( + kr, + opts[:strides], + {lows, highs}, + opts[:kernel_dilation], + opts[:input_dilation], + opts[:feature_group_size], + false + ) + |> Native.transpose(nhwc_to_nchw) + |> Native.transpose(inv_op) + |> wrap(out) + end + end + + # Invert a 0-based permutation: given `perm` where position i holds + # j, produce `inv` where position j holds i. Used to reverse + # `output_permutation` — Nx delivers it in "user → canonical" form + # (see `deps/nx/lib/nx/shape.ex:729-735`), so we need the inverse to + # go "canonical → user". + defp invert_permutation(perm) do + perm + |> Enum.with_index() + |> Enum.sort() + |> Enum.map(&elem(&1, 1)) + end # ================================================================= # Unsupported / fallback callbacks diff --git a/test/emily/backend_conv_test.exs b/test/emily/backend_conv_test.exs new file mode 100644 index 0000000..e2ae80c --- /dev/null +++ b/test/emily/backend_conv_test.exs @@ -0,0 +1,156 @@ +defmodule Emily.BackendConvTest do + @moduledoc """ + Oracle tests for `Emily.Backend.conv/4`. Each case runs the same + `Nx.conv` on `Emily.Backend` and `Nx.BinaryBackend` and asserts the + outputs agree within a dtype-appropriate tolerance. + + Covers the layout mapping (Nx NCHW/OIHW ↔ MLX NHWC/OHWI), the + non-default permutation paths, grouping (including depthwise), + dilation, padding variants, and integer-operand `Nx.astype` coercion. + """ + + use ExUnit.Case, async: true + + import Emily.BackendGenerators, only: [assert_close: 2] + + defp emily(tensor), do: Nx.backend_transfer(tensor, Emily.Backend) + defp bin(tensor), do: Nx.backend_transfer(tensor, Nx.BinaryBackend) + + # Build a pair of (input, kernel) on BinaryBackend with deterministic + # values scaled so accumulation stays in the f32 sweet spot. + defp inputs(input_shape, kernel_shape) do + scale = 0.1 + + input = + input_shape + |> Nx.iota(type: {:f, 32}, backend: Nx.BinaryBackend) + |> Nx.multiply(scale) + + kernel = + kernel_shape + |> Nx.iota(type: {:f, 32}, backend: Nx.BinaryBackend) + |> Nx.multiply(scale) + + {input, kernel} + end + + defp run(input, kernel, opts) do + emily_result = Nx.conv(emily(input), emily(kernel), opts) + ref_result = Nx.conv(bin(input), bin(kernel), opts) + assert_close(emily_result, ref_result) + end + + describe "rank & default layout" do + test "1-D conv, stride 1, padding :valid" do + # {batch, channels, length} × {out, in, width} + {input, kernel} = inputs({2, 3, 8}, {4, 3, 3}) + run(input, kernel, []) + end + + test "2-D conv, default NCHW, stride 1, padding :valid" do + {input, kernel} = inputs({2, 3, 5, 5}, {4, 3, 3, 3}) + run(input, kernel, []) + end + + test "3-D conv (rank-5) smoke" do + {input, kernel} = inputs({1, 2, 4, 4, 4}, {3, 2, 2, 2, 2}) + run(input, kernel, []) + end + end + + describe "strides, padding, dilation" do + test "2-D conv with stride 2" do + {input, kernel} = inputs({1, 2, 6, 6}, {4, 2, 3, 3}) + run(input, kernel, strides: 2) + end + + test "2-D conv with padding :same" do + {input, kernel} = inputs({1, 2, 5, 5}, {3, 2, 3, 3}) + run(input, kernel, padding: :same) + end + + test "2-D conv with explicit asymmetric padding" do + {input, kernel} = inputs({1, 2, 5, 5}, {3, 2, 3, 3}) + run(input, kernel, padding: [{1, 2}, {0, 1}]) + end + + test "2-D conv with kernel_dilation > 1" do + {input, kernel} = inputs({1, 2, 7, 7}, {3, 2, 3, 3}) + run(input, kernel, kernel_dilation: 2) + end + + test "2-D conv with input_dilation > 1 (transposed conv)" do + {input, kernel} = inputs({1, 2, 4, 4}, {3, 2, 3, 3}) + run(input, kernel, input_dilation: 2) + end + end + + describe "grouping" do + test "feature_group_size: 2" do + # in_channels = 4 split into 2 groups; each group contributes to + # 3 of the 6 output filters. + {input, kernel} = inputs({1, 4, 5, 5}, {6, 2, 3, 3}) + run(input, kernel, feature_group_size: 2) + end + + test "depthwise conv (feature_group_size == in_channels)" do + # in_channels = 4 → 4 groups (each channel independent); + # kernel in-channels-per-group = 1; out_channels = 4. + {input, kernel} = inputs({1, 4, 5, 5}, {4, 1, 3, 3}) + run(input, kernel, feature_group_size: 4) + end + end + + describe "non-default permutations" do + test "input_permutation: [0, 3, 1, 2] (caller passes NHWC input)" do + # Start from canonical NCHW, transpose to NHWC, then tell Nx about it. + {nchw_input, kernel} = inputs({1, 3, 5, 5}, {4, 3, 3, 3}) + nhwc_input = Nx.transpose(nchw_input, axes: [0, 2, 3, 1]) + + run(nhwc_input, kernel, input_permutation: [0, 3, 1, 2]) + end + + test "kernel_permutation: [3, 2, 0, 1] (caller passes HWIO kernel)" do + {input, oihw_kernel} = inputs({1, 3, 5, 5}, {4, 3, 3, 3}) + # HWIO: [height, width, in, out]; pull axes so that original [0,1,2,3] + # (O, I, H, W) becomes [2, 3, 1, 0] (H, W, I, O). + hwio_kernel = Nx.transpose(oihw_kernel, axes: [2, 3, 1, 0]) + + run(input, hwio_kernel, kernel_permutation: [3, 2, 0, 1]) + end + + test "output_permutation: [0, 2, 3, 1] (caller wants NHWC output)" do + {input, kernel} = inputs({1, 3, 5, 5}, {4, 3, 3, 3}) + run(input, kernel, output_permutation: [0, 2, 3, 1]) + end + + test "all three permutations non-default (NHWC end-to-end)" do + {nchw_input, oihw_kernel} = inputs({1, 3, 5, 5}, {4, 3, 3, 3}) + nhwc_input = Nx.transpose(nchw_input, axes: [0, 2, 3, 1]) + hwio_kernel = Nx.transpose(oihw_kernel, axes: [2, 3, 1, 0]) + + run(nhwc_input, hwio_kernel, + input_permutation: [0, 3, 1, 2], + kernel_permutation: [3, 2, 0, 1], + output_permutation: [0, 2, 3, 1] + ) + end + end + + describe "dtype coercion" do + test "integer input is cast to float output dtype" do + # Nx.conv produces float output even for integer operands; backend + # must cast before dispatching to MLX (which is float-only for conv). + input = Nx.iota({1, 2, 4, 4}, type: {:s, 32}, backend: Emily.Backend) + kernel = Nx.iota({3, 2, 2, 2}, type: {:s, 32}, backend: Emily.Backend) + + ref_input = Nx.iota({1, 2, 4, 4}, type: {:s, 32}, backend: Nx.BinaryBackend) + ref_kernel = Nx.iota({3, 2, 2, 2}, type: {:s, 32}, backend: Nx.BinaryBackend) + + emily_result = Nx.conv(input, kernel) + ref_result = Nx.conv(ref_input, ref_kernel) + + assert_close(emily_result, ref_result) + end + end +end diff --git a/test/emily/backend_fallbacks_test.exs b/test/emily/backend_fallbacks_test.exs index d23c2d7..287fdab 100644 --- a/test/emily/backend_fallbacks_test.exs +++ b/test/emily/backend_fallbacks_test.exs @@ -68,20 +68,6 @@ defmodule Emily.Backend.FallbacksTest do end end - describe "convolution fallback" do - test "conv routes through BinaryBackend" do - # {batch=1, channels=1, height=3, width=3} input, {out=1, in=1, 2, 2} kernel. - input = Nx.iota({1, 1, 3, 3}, type: {:f, 32}, backend: Emily.Backend) - kernel = emily([[[[1.0, 0.0], [0.0, 1.0]]]]) - - result = Nx.conv(input, kernel) - - assert Nx.shape(result) == {1, 1, 2, 2} - # Diagonal kernel: 0+4=4, 1+5=6, 3+7=10, 4+8=12. - assert flat(result) == [4.0, 6.0, 10.0, 12.0] - end - end - describe "reduce fallbacks" do test "reduce with a custom accumulator function" do t = emily([1.0, 2.0, 3.0, 4.0]) @@ -248,5 +234,25 @@ defmodule Emily.Backend.FallbacksTest do result = Nx.cumulative_min(t, axis: 1) assert Nx.shape(result) == {2, 2, 2} end + + # Conv with `batch_group_size > 1`: MLX `conv_general` has no batch-group + # parameter, so Emily falls back to BinaryBackend for this rare path. The + # common `batch_group_size: 1` case is covered natively in + # `backend_conv_test.exs`. + test "conv with batch_group_size > 1 falls back" do + # Splitting batch=4 into 2 groups; each produces 1 output filter from + # its half of the input channels. + input = Nx.iota({4, 2, 3, 3}, type: {:f, 32}, backend: Emily.Backend) + kernel = Nx.iota({2, 2, 2, 2}, type: {:f, 32}, backend: Emily.Backend) + + result = Nx.conv(input, kernel, batch_group_size: 2) + + ref_input = Nx.iota({4, 2, 3, 3}, type: {:f, 32}, backend: Nx.BinaryBackend) + ref_kernel = Nx.iota({2, 2, 2, 2}, type: {:f, 32}, backend: Nx.BinaryBackend) + ref = Nx.conv(ref_input, ref_kernel, batch_group_size: 2) + + assert Nx.shape(result) == Nx.shape(ref) + assert flat(result) == Nx.to_flat_list(ref) + end end end