backends/cuda: a torch reference backend — edge0-8b and edge0-35b generate text on an NVIDIA GPU - #19
Open
agourakis82 wants to merge 32 commits into
Open
agourakis82 wants to merge 32 commits into
agourakis82 wants to merge 32 commits into
Conversation
…=cuda dispatch - core/nn/io/quant namespaces matching the documented contract - quant.gather_qmm: reference affine int4 dequant+matmul (numerics UNVERIFIED against mx.quantize -- no MLX/Mac available to check) - io.load_model: explicit NotImplementedError, not a silent stub -- see docstring for the two different paths edge0-35b and edge0-8b need - passes tests/test_repo_hygiene.py unmodified Also see: MLX itself ships a native CUDA backend (pip install mlx[cuda12]) with a real GatherQMM CUDA kernel and a cudaMallocManaged-backed allocator that preserves the unified-memory model on GPUs with concurrentManagedAccess. Testing EDGE0_BACKEND=mlx + mlx[cuda12] unmodified on real NVIDIA hardware is very likely a faster path than this reference backend for the shipped tiers; this commit is the fallback/reference, not the recommended first step.
MLX's own CUDA backend (pip install mlx[cuda12], manylinux x86_64 + aarch64 wheels at the exact version this repo already pins) implements every op the streaming/quant path depends on, including a real CUDA GatherQMM kernel. No EDGE0_BACKEND change, no framework code change -- just install the CUDA variant instead of Metal. Also corrects a now-stale assumption in test_repo_hygiene.py's docstring: MLX does ship Linux wheels (CPU-only and CUDA-backed) as of at least the version already pinned here. Benchmark table left as TBD pending a run on real NVIDIA hardware.
…real checkpoints)
_resolve_model_class dispatches on architectures/auto_map (config.json
has no top-level model_type field -- checked directly against the real
Edge0/Edge0-8B-A1B-preview config, an earlier model_type-keyed version
of this function would have silently mis-dispatched on it):
- edge0-35b: transformers.Qwen3_5MoeForCausalLM natively. Experts are
one stacked [num_experts, ...] tensor per projection.
- edge0-8b: BailingMoeV3ForCausalLM via trust_remote_code, using the
modeling_bailing_moe_v3.py/configuration_bailing_moe_v3.py the
checkpoint repo ships directly (confirmed via the HF API file
listing, then downloaded and actually imported -- needs einops,
fla, triton as extra deps, not currently in edge0's own dependency
list). Experts are nn.ModuleList of per-expert MLP modules, not one
stacked tensor -- structurally different from the 35b tier.
Both verified end-to-end on torch.device('meta') against the REAL
downloaded config + modeling files (not synthetic): construction
succeeds with zero real memory, and the Bailing tier's meta-device
param count (7,923,995,168) matches the model card's ~7.9B almost
exactly.
load_model loads every dense (non-expert) tensor for real from the
safetensors shards and leaves expert tensors on the meta device,
deliberately -- wiring them to streaming/layer.py's per-forward gather
is the next unit of work, blocked on quant.gather_qmm's packing being
verified first, and now also confirmed to need two different wiring
shapes (stacked-tensor gather for 35b, per-expert-module swap for 8b),
not one.
…kpoints; add MoESpec instances Downloaded config.json (+ safetensors index/header, no weight data) for BOTH real published checkpoints and tested against them directly -- found and fixed two bugs that synthetic testing missed: 1. edge0-35b's real architectures is Qwen3_5MoeForConditionalGeneration (vision+text wrapper), not Qwen3_5MoeForCausalLM as first assumed -- now correctly extracts just text_config and builds the text-only class. 2. edge0-35b's real on-disk keys all carry a language_model. prefix (confirmed via model.safetensors.index.json) that transformers' Qwen3_5MoeForCausalLM's own attribute names don't have -- every dense weight would have silently failed to load without stripping it. edge0-8b needs no such prefix -- confirmed the same way, not assumed identical. Added backends/cuda/model_specs.py: MoESpec for both tiers, reusing the exact family facts already verified in models/edge0_35b and models/edge0_8b, with block_path corrected for transformers' actual (unwrapped) attribute path while key_template keeps the real on-disk prefix. Tested block_of()/layer_of() resolution against both real meta-device model instances (34.66B / 7.92B params, matching both model cards almost exactly) -- both resolve to the correct block type with the expert attribute install_streaming_experts already checks for, unmodified.
….quantize output Ran a ground-truth dump on real hardware (DGX Spark, GB10, mlx 0.32.2, CPU execution -- packing is backend-independent) and checked it against this file's _unpack_affine_u4: LSB-first nibble order confirmed (dequant error 0.0 vs 4.15 in reverse order), unsigned codes 0..15, no zero-point, w = code*scale + bias, scale can be negative. No code change needed -- the implementation already matched; only the docstrings' certainty level changes, from UNVERIFIED to VERIFIED with the source of that verification stated. Also flagged, found by the same test and NOT yet fixed: gather_qmm's default calling convention returns a full [n_experts, n_x_rows, out] broadcast, not the one-row-per-expert [T, out] shape this file assumes -- streaming/layer.py's real calls use sorted_indices=True plus pre-expanded x, which the test didn't exercise. Dequant math is solid; the calling-convention shape for real streaming calls still isn't verified.
…ts vs real MLX The previous gather_qmm crashed whenever more than one expert was gathered: _unpack_affine_u4 reshaped n_rows*out*in codes to (out, in), so the function had never actually run (RuntimeError: shape '[16, 128]' is invalid for input of size 4096 with two experts). Checked against real MLX 0.30.4 on Metal (the pinned version): - output batch shape is broadcast(x.shape[:-2], rhs_indices.shape), and out[b] = x[b] @ W[rhs_indices[b]].T -- both streaming/layer.py call patterns (unsorted: x[...,1,1,D] vs idx[...,K]; sorted, after _gather_sort: x[T*K,1,D] vs idx[T*K]) are instances of this rule; - sorted_indices is a kernel hint only: flipping it changes nothing (max diff 0.0), so it is ignored here. gather_qmm now reshapes the gathered, dequantized experts to rhs_indices.shape + (rows, in) and lets torch.matmul broadcast, which is exactly that rule. Also generalizes the unpack to 2/4/8 bits and drops the dead duplicated block + second swiglu definition left by the previous commit. tests/test_cuda_backend.py compares against mx.gather_qmm directly: default convention, both streaming paths, other broadcast shapes, 2/4/8 bits, transpose=False, sorted_indices=True, bf16 checkpoint dtypes. 12/12 pass; against the previous quant.py 11/12 fail. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
- tests/test_nvidia_backend_smoke.py: test_gather_qmm_matches_dense_reference failed on Metal (rel_err=1.40) -- a [2, in] x broadcasts against both indices and returns [2, 2, out]. Use [2, 1, in] so each row goes through its own expert. _has_gpu() now evaluates an op instead of trusting default_device(), which never touches the driver. Dropped the claim that the op was "confirmed working" on CUDA 0.32.x -- it never ran there. - backends/cuda/nn.py: RMSNorm subclasses torch.nn.RMSNorm instead of wrapping it, so the parameter is named "weight" (as in MLX and the checkpoints) rather than "norm.weight". RMSNorm and gelu are now checked against mx.fast.rms_norm / mlx.nn.gelu. - backends/__init__.py, backends/cuda/backend.py: remove the advice to try EDGE0_BACKEND=mlx + mlx[cuda12] first because it "implements every op"; docs/nvidia.md shows three MLX releases failing on real hardware. - backends/cuda/io.py, nn.py: drop references to a "mapping doc" that is not in the repo, and the stale "MLX has no Linux wheels" remark. - docs/nvidia.md: backends/cuda/ is no longer an ImportError stub; say what exists and what is left. Full suite: 76 passed, 2 skipped (env-gated). Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
…swers as MLX The shared code reached past the backend facade into MLX-only array API, and the torch facade diverged from MLX at the call sites that exist. torch facade (backends/cuda/core.py), each checked against real MLX: - argpartition: returned k indices via topk(kth + 1), so the kth=-top_k calls in moe/routing.py and prerouter/ asked for topk(-7). Now a stable argsort -- a full permutation, which satisfies the partition property for any kth, as mx.argpartition's output does. - topk: returned torch's (values, indices) tuple; mx.topk returns only the values, ascending (sampling.py reads the threshold from the first entry). - take_along_axis / put_along_axis: numpy/MLX broadcast the indices over the non-axis dims, torch.gather/scatter do not. group_select_from_logits masks whole expert groups with [..., k, 1] indices, so only one expert per dropped group was being masked and the 8b router picked different experts. put_along_axis also broadcasts a scalar value. - softmax(precise=True), max, maximum, argmax, set_cache_limit added; compile is now eager (torch.compile re-traces on every new shape). Shared code: mlx.core has no astype/size functions (they are array methods), so backends/mlx/core.py re-exports mlx.core -- same objects -- and adds astype, size and index_add as thin calls to the MLX methods. sampling.py and moe/routing.py use core.astype / core.size, and a descending sort is the negation of an ascending sort of the negation (exact; torch has no negative-step slicing). quant.gather_sort and quant.scatter_unsort move the mlx_lm helpers behind the facade (MLX delegates to mlx_lm itself). tests/test_backend_parity.py runs the same cases under EDGE0_BACKEND=mlx and =cuda in subprocesses: both routers pick identical experts (scores within 1e-5) and the sampler masks match. Full suite: 79 passed. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
… out-of-bounds gather
layer.py reached into MLX directly -- the mlx_lm sort helpers, the .size
property and .at[] indexing -- so it could not run on torch arrays. It
now goes through the facade (quant.gather_sort / scatter_unsort,
core.size, core.index_add); on MLX these are the same calls.
Upstream bug, found while porting: the hot-stack prefill path sends
misses to row n_hot ("overflow zero row") of a stack that has only n_hot
rows. MLX gathers do not bounds-check, so every miss read past the end of
the stack -- the path was exact only while that memory happened to be
zeros. mx.gather_qmm with an out-of-range index returns garbage in
isolation (values around 600 in the probe). The hot backing now carries
an explicit all-zero row, as the staged path already does with its zero
slot; _build and materialize_hot account for it.
torch facade: expand_dims resolved negative axes one unsqueeze at a time,
so expand_dims(x[m, H], (-2, -3)) gave [1, m, 1, H] instead of
[m, 1, 1, H]. Negative axes now index the output, as in numpy/MLX.
tests/test_backend_parity.py::test_streaming_layer_real_experts (needs
EDGE0_8B_MODEL) runs layer 1 of the real edge0-8b checkpoint through the
exact (unsorted and sorted), whole-layer, hot (half the experts missing)
and staged paths, compiled and eager, on both backends. torch matches MLX
to 1.2-1.5% of the output scale on every path (about two bf16 ulps after
three chained 4-bit matmuls), and hot == exact on each backend.
MLX path unchanged: greedy decode of 32 tokens on the real edge0-8b
checkpoint gives the same token ids as upstream fbab5f8.
Full suite with EDGE0_8B_MODEL set: 84 passed.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
- .astype(...) -> core.astype(...) in prerouter/heads.py, stager.py, install.py and adapters/lora.py (torch tensors have no .astype). - load_safetensors comes from the active backend's io instead of edge0.backends.mlx.io in prerouter/install.py and adapters/lora.py. - prerouter/install.py imported the vendored MLX Qwen3Next block at module level only to class-patch it for the 35b tier; the import now lives in _patch_qwen_consume so the module imports under any backend. Porting that patch (and the plain-array weight assignment on nn.Linear, which torch rejects) belongs with the torch model port. MLX path unchanged: LoRA applied=153, prerouter 16 heads, and the same 32 greedy tokens on the real edge0-8b checkpoint. Suite: 84 passed. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
install_streaming_experts is now exercised end to end on the torch backend: a transformers Qwen3_5MoeForCausalLM, the edge0-35b MoESpec paths from backends/cuda/model_specs.py (language_model.-prefixed keys), and experts MLX-quantized into a real safetensors shard. With the twins installed the logits match the same model running its own dense experts on the weights the shard encodes (max relative error 6-8e-4 -- the reference dequantizes in bf16 -- and the same argmax at every position), for both the unsorted (T=6) and sorted (T=40) paths. What that needed: - transformers' Qwen3.5 block calls experts(hidden, top_k_index, top_k_weights) and expects the weighted sum back; the twin returns raw per-expert outputs like the vendored MLX blocks expect. backends/cuda/moe_blocks.py::TransformersExpertsAdapter bridges that, passed through a new wrap= argument (default: no wrapping, so the MLX engines are unaffected). - torch refuses to assign a non-Module over a registered child module; the swap falls back to replacing the child with a plain attribute. - block._edge0_resident = block registered the block as its own child module (torch's state_dict would recurse forever); the resident stashes are plain attributes now, which is all their readers need. - Layer discovery stopped only on AttributeError, but indexing past the last layer raises IndexError (ModuleList and plain lists alike), so num_layers=None never worked; the engines always pass it. Not covered by an adapter: the edge0-8b checkpoint's own modeling_bailing_moe_v3.py indexes and iterates self.experts as a ModuleList of 128 modules, so its block forward has to be replaced, not its experts (documented on the adapter). Suite: 84 passed. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
The previous load_model had only been checked by constructing models on
the meta device; no weight was ever read. Against the published
checkpoints it could not work:
- every edge0 checkpoint is an MLX checkpoint, quantized throughout:
392 non-expert tensors for edge0-35b and 236 for edge0-8b carry .scales
(embeddings, lm_head, attention, shared experts, the 35b routers).
_NP_DTYPES had no U32, so the first quantized tensor raised KeyError;
past that, a packed U32 [V, H/8] embedding would meet a [V, H] bf16
parameter with the same name.
- edge0-35b went through edge0's MLX qwen3_5.py sanitize: conv1d weights
are stored [C, k, 1] (torch wants [C, 1, k]) and five RMSNorm kinds are
stored as w + 1, because MLX computes x * w and transformers x * (1 + w).
Read from the published bytes with HTTP range requests: input_layernorm
mean 1.03, q_norm 1.33, model.norm 2.63 -- a transformers checkpoint is
near 0. Loaded as-is, every one of those norms becomes 2 + w, silently.
- building the whole model under torch.device("meta") also leaves
non-persistent buffers (rotary inv_freq) on meta, so no forward could run.
Now:
- modules whose weight has .scales become QuantizedLinear /
QuantizedEmbedding (backends/cuda/nn.py; 4-bit payload stays resident,
dequantized on the fly with the verified unpack). bits and group size
come from the shapes, which covers the per-path overrides in config.json
(8-bit routers). Quantized params of other modules (transformers' Qwen3.5
router is not an nn.Linear) are dequantized in place.
- _undo_mlx_qwen35_sanitize transposes conv1d back and subtracts 1 from
the shifted norms, only when conv1d is in MLX layout -- the same signal
the sanitize uses in the other direction.
- parameters are created on meta, buffers for real; dtype= casts dense
tensors; strict=True now means no missing and no unexpected non-expert
tensors, and is what the test uses.
tests: test_load_model_edge0_35b_format writes a small Qwen3.5-MoE
checkpoint in exactly the published edge0-35b format (ConditionalGeneration
config with text_config, language_model. prefix, 4-bit linears and
embeddings, 8-bit router, switch_mlp experts, w + 1 norms, MLX-layout
conv1d), loads it strict, streams the experts in, and matches the source
model on the same effective weights to 4e-7 relative (float32 noise); 37
quantized modules, no parameter left on meta. The reference now
dequantizes with float32 arithmetic -- with bf16 scales MLX rounds the
weight itself to bf16 -- which also brings the install test from 7e-4 to
5e-7; both tolerances tightened to 1e-5.
Not addressed: edge0-8b still needs the checkpoint's own
modeling_bailing_moe_v3.py, which imports fla (Triton kernels; no macOS
wheels) and iterates its experts as a ModuleList.
Suite with EDGE0_8B_MODEL set: 86 passed.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
engine/qwen.py and engine/ling.py imported the vendored MLX models and backends.mlx.io at module level, so on a machine without MLX the package could not even be imported, and on one with both backends an EDGE0_BACKEND=cuda run built MLX models and pushed their arrays through torch ops. - load_model / open_shards / load_tokenizer now come from the active backend's io (the same functions under MLX); the vendored model classes are imported inside _get_model_classes. - load_installed starts with require_mlx_backend(tier) (engine/base.py): on another backend it raises NotImplementedError naming what exists (ops, streaming layer, load_model) and what does not (an engine for the tier), instead of failing deep inside a forward. The engines themselves stay MLX-only: they drive the vendored models' per-layer callbacks, mlx-lm caches and the class-level prerouter patch, so a torch engine is a port, not a conditional import. tests: under EDGE0_BACKEND=cuda, importing edge0, both engines, the CLI, prerouter, adapters and streaming loads no mlx / mlx_lm module, and both engines refuse with the message. MLX path unchanged: pytest -m slow passes for edge0-8b on Metal, the same 32 greedy tokens, and decode time within run-to-run noise of upstream (6.2 s vs 6.3 s mean over three alternating runs). Suite: 87 passed. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
…s left Replaces the "streaming layer and engines still reach into MLX" summary with the current state: a table of each piece and the MLX reference it is checked against (with the measured agreement), what load_model handles in the published checkpoints, and the remaining work -- torch engines, why edge0-8b fits the transformers route poorly (remote code, fla/Triton, ModuleList experts) versus porting the vendored model, a real-weight edge0-35b run, and performance. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
… by layer against MLX backends/cuda/_impl/bailing_hybrid.py ports the vendored MLX Ling 3.0 backbone line by line -- KDA (short convs, safe-gate decay, delta-rule recurrence as mlx-lm's gated_delta_ops), MLA (interleaved RoPE, head-wise output gate), dense MLP, sigmoid group-limited router, prerouter head, sparse MoE, decoder layer, model -- with the same module tree, parameter names and call signatures, plus torch KVCache / ArraysCache. It replaces the checkpoint's own modeling_bailing_moe_v3.py (remote code, fla/Triton, ModuleList experts). Routed experts come from install_streaming_experts. Two places where torch semantics differ from MLX and are handled explicitly: - mx.fast's "causal" mask is aligned to the end of the keys (query i of L sees keys up to Lk - L + i); torch's is_causal aligns to the start, which is wrong once a cache holds earlier tokens. _sdpa builds the mask. - dtype promotion: MLX's RoPE runs in float32 and its concatenate promotes, so the hidden state becomes float32 at the first MLA layer and stays there. The port follows the same promotion (a first draft cast back to bf16 -- no error, just a different model). backends/cuda/io.py: load_model's engine path (get_model_classes, as the engines call it) builds the vendored pair from config.json plus overrides and returns (model, config) like mlx-lm, calling get_model_classes(config=...) as mlx-lm does. backends/cuda/nn.py: Linear accepts a plain tensor assigned to weight/bias, as mlx.nn.Linear does and prerouter/install.py relies on. Verification (real edge0-8b checkpoint): - load accounting: 0 missing parameters, 0 left on meta, 236 quantized modules (all 236 quantized non-expert tensors of the checkpoint); the only unused tensors are the in-checkpoint prerouter copy (mlp.pregate), which edge0 loads from its own file. - teacher-forced layer by layer against the MLX model on the MLX CPU device, float32, over a chunked prefill (7 + 5 tokens, so the second chunk runs over a cache) and decode steps: every layer <= 1.5e-6 relative, free-running logits <= 1.7e-6, same argmax at every step. Against MLX on the GPU the MLA/dense layers look worse (up to 1.5e-3): measured against float64, that is MLX's Metal float32 matmul/SDPA on this machine (M5 Max: 7.5e-4 from float64; MLX on the CPU and torch: 2e-7), not the port. - negative control: with torch's start-aligned is_causal the MLA layers jump to 0.2-0.3 and the free-running logits to 0.21 on the second chunk. - bf16 (production dtype): per layer <= 2e-2 (1-3 bf16 ulps), same argmax over 7 steps, and the same hidden-state dtype after every layer as MLX. Tests: piece by piece against MLX without a checkpoint (KDA recurrence vs gated_delta_ops, RoPE with an offset, short conv carrying state, router, causal SDPA with an offset vs float64), and test_bailing_port_matches_mlx_layer_by_layer on the real checkpoint (EDGE0_8B_MODEL; per-layer and logits < 1e-5, load accounting). Suite with EDGE0_8B_MODEL set: 93 passed. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
engine/ling.py now picks the active backend's port of the bailing backbone
(_get_model_classes) and accepts EDGE0_BACKEND=cuda; the engine itself is
unchanged -- LoRA, prerouter-staged decode, streaming experts and sampling
all run through the torch port and backend. require_backend(tier,
supported) replaces the MLX-only guard; edge0-35b stays MLX-only.
Result on the real edge0-8b checkpoint: 8 greedy tokens on torch (CPU)
identical to the MLX engine's ("The sky appears blue because sunlight
consists of"), with the same build log (lora applied=153 not_found=0,
prerouter installed: 16 heads).
Memory, found by running it (each fix keeps the per-element arithmetic;
the full parity suite still passes):
- gather_qmm dequantized one copy per (token, expert) pair -- 248 copies
per projection in a 31-token prefill, several GB per MoE layer. It now
dequantizes each distinct expert once and applies it to the rows routed
to it.
- QuantizedLinear dequantized its whole weight per call (edge0-8b's
lm_head: ~1 GB float plus 1.9 GB of int64 codes). It now works in
chunks of 4096 output rows, and _dequantize unpacks in int32 (the mask
keeps only the low bits, which sign extension never reaches).
Peak RSS of a 32-token run: 14.2 GB -> 7.3 GB.
Longer runs on this machine are still killed by macOS jetsam for low swap
(swap 8.8 of 10 GB in use by other workloads): system memory pressure,
not the engine.
Tests: the engine-guard test now expects edge0-8b to resolve to the torch
port and edge0-35b to refuse; test_edge0_8b_engine_same_tokens_on_both_backends
(slow) greedy-decodes through the engine on both backends and requires
identical tokens. docs/nvidia.md updated. Suite: 93 passed.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_017vBR5ebDS5kubR1mdvfLN6
…runs on torch backends/cuda/_impl/qwen3_5_moe.py ports the vendored MLX Qwen3.5-MoE backbone (qwen3_5_moe.py + the qwen3_5.py / qwen3_next.py pieces, mlx-lm 0.31.0) line by line: GatedDeltaNet (depthwise conv, scalar per-head decay g = exp(-exp(A_log) softplus(a + dt_bias)) in a's dtype, state in q's dtype, k heads repeated to the value heads, gated RMSNorm with the gating in float32), gated full attention with GQA, q/k norms and partial RoPE (mx.fast.rope's half-split layout on 25% of head_dim), softmax top-k MoE block with a gated shared expert, decoder layer, text model with the engine's layer callbacks, and the sanitize -- same module tree, parameter names and call signatures, so engine/qwen.py runs unchanged. It computes the MLX way on the MLX checkpoint layout (RMSNorm weight as the absolute multiplier, conv1d [C, k, 1]), so the published edge0-35b tensors are used as stored: no conversion on load, which is where the transformers route needed the w + 1 / conv-layout fix-ups. engine/qwen.py picks the backend's port and accepts EDGE0_BACKEND=cuda; prerouter/install.py patches the active backend's MoE block class. Verification, without the 23 GB checkpoint -- MLX builds a small edge0-35b-shaped model, casts it to bf16, quantizes it with its own predicate and saves it in its own format (the published format, including the per-path 8-bit overrides in config.json): - pieces vs MLX on the CPU device: gated_delta_update (with GQA repeat), partial RoPE at offsets 0 and 7 vs mx.fast.rope, gated norm, depthwise conv, gated GQA attention through the KV cache over prefill, a second chunk and decode steps. - the model, teacher-forced layer by layer and free-running, float32: <= 2.6e-7 per layer, <= 4.1e-7 on the logits, same argmax at every step; 0 missing / 0 unexpected tensors, router loaded at 8 bits. - the whole engine (streaming, staged decode, class-level prerouter patch with heads in the real file format; LoRA off): identical greedy tokens, per-step logits within bf16 noise (0.35-2.5%). Token equality alone was too coarse -- the same tokens come out with the prerouter disabled -- so the test also requires that, at the steps where the prerouter acts (it moves the logits 5-11% there), torch is at least twice as close to MLX with the prerouter as to MLX without it. - negative controls: the transformers norm convention x * (1 + w) sends the per-layer error to 1.1 and changes the tokens; the wrong RoPE layout or offset gives errors of 4-5 against mx.fast.rope. Suite with EDGE0_8B_MODEL set: 99 passed. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017vBR5ebDS5kubR1mdvfLN6
… the driver The caveat said an aborted CUDA process could leave the GB10 driver returning 999 until a reboot. On that machine the cause was a host cgroup device policy denying /dev/nvidia-uvm and /dev/nvidia-caps (EPERM when opened as a user), attached a few minutes after boot. Say that, and give the check that tells the two apart. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017vBR5ebDS5kubR1mdvfLN6
…ole autograd graph) MLX arrays carry no autograd state; torch parameters default to requires_grad=True. With no no_grad around the engine's forward, torch recorded the graph of every forward and kept everything saved for backward alive until the logits were dropped -- each dequantized expert of gather_qmm included. On Linux the edge0-8b engine's 31-token prefill grew 1.2 GB/s and took a 121 GB, swapless host into the OOM killer. load_model now freezes the model it returns, and the facade Linear and RMSNorm freeze themselves on construction (the prerouter heads are built after load_model). Grad mode is thread-local in torch, so freezing the parameters rather than disabling grad in one thread holds for any caller. Checked: - on the DGX Spark (Linux aarch64, torch CPU), under a cgroup memory cap: before, anonymous memory hit a 24 GB cap in 20 s whatever glibc malloc was told (MALLOC_ARENA_MAX=2, a fixed mmap threshold); with grad off it stays at 1.1-1.25 GB. The 32 greedy tokens of a chat prompt are identical to MLX on Apple Silicon on the same checkpoint (same hashes). - the qwen35 engine parity case now reports how many steps' logits carry a graph: 0 with this change, all 10 without it (negative control). Suite with EDGE0_8B_MODEL set: 99 passed, 2 skipped. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017vBR5ebDS5kubR1mdvfLN6
…ers to DEVICE EDGE0_TORCH_DEVICE overrides the device choice (cpu, mps, cuda). On Apple Silicon, "mps" puts torch on the GPU next to the MLX reference, where a tensor left on the host fails loudly as it would on CUDA -- the closest check of device placement available without a CUDA device. It found one real bug: load_model moved the checkpoint tensors to DEVICE but not the buffers modules compute at init and the checkpoint lacks (transformers' rotary inv_freq), so the transformers route failed on any device but the CPU. They are now moved one by one -- not model.to(DEVICE): the streamed expert parameters are still on meta and cannot be copied. The parity worker's cases fed host tensors (torch.tensor, from_numpy) to device models; they now put inputs, and the transformers reference model, on the backend's device and bring results back with .cpu(). With EDGE0_TORCH_DEVICE=mps: the whole suite (99 passed, 2 skipped) and the slow engine test pass, and the edge0-8b engine gives the same 32 greedy tokens as MLX. On the CPU: 99 passed, 2 skipped. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017vBR5ebDS5kubR1mdvfLN6
… has no uint32 index) The edge0-8b engine on a GB10 died in QuantizedEmbedding: "index_cuda not implemented for UInt32". CUDA has no index kernel for uint32; the CPU and MPS backends do, so nothing on this Mac could catch it. The packed rows are read as int32 by _dequantize anyway, so the gather now runs on the int32 view. Same arithmetic, same results. With that, edge0-8b runs on the GPU (DGX Spark GB10, sm_121, torch 2.14+cu130, as a Slurm job): - every layer fed MLX's own input for that layer, float32: 5.3e-7 max per layer (median 2.6e-7), 4.3e-7 on the logits, same argmax, against MLX on the Apple CPU device. TF32 off, float32_matmul_precision "highest". - the whole engine, 32 greedy tokens: 31 of 32 identical to MLX, diverging at step 31. Not a GPU artifact -- torch on the CPU sits at the same 4.5% median per-step logit distance from MLX (the staged prerouter drops non-staged experts), and the top-2 margin at that step is below that noise. - 12.6 s for the 32 tokens, against 58 s on the Spark's Grace CPU. Suite: 99 passed, 2 skipped, on the CPU and with EDGE0_TORCH_DEVICE=mps. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017vBR5ebDS5kubR1mdvfLN6
…eights
The 35b port was checked only on a small model MLX wrote in the published
format. This adds the same layer-by-layer harness the 8b has, on the real
23 GB checkpoint, gated on EDGE0_35B_MODEL.
case_qwen35_real_parity loads both the vendored MLX model and the torch
port from the checkpoint, MLX on the CPU device (its Metal float32 matmul
runs at ~7e-4 from float64 on this machine, the CPU path at 2e-7), both
float32, and feeds every torch layer MLX's input to that layer; separately
the whole torch model runs free with its own caches, over a chunked
prefill (7 + 5 tokens, so the second chunk runs over a non-empty cache)
and decode steps.
Measured on the real weights: 40 layers, 0 missing parameters, 0
unexpected tensors, 40 streaming twins; <= 2.4e-6 per GatedDeltaNet layer,
<= 3.7e-6 per gated full-attention layer, <= 2.4e-6 on the free-running
logits, same argmax at every step. Thresholds in the test are 1e-5.
The checkpoint itself was verified before use: all four shards byte-exact
against their own safetensors headers and all 1757 index keys present with
no extras (the index's total_size is 0.9 GB larger than the tensor bytes,
which is publisher metadata, not truncation).
Engine end to end (LoRA 310 targets, prerouter 33 heads, streamed
experts), greedy, same prompt: the torch backend and MLX agree on all 32
tokens when MLX runs on its CPU device. Against MLX on Metal, 29 of 32
match and the walk splits at step 28 on a near-tie ("limited to exactly
two" vs "specifically limited to two"); MLX-Metal disagrees with
MLX-CPU at exactly the same three positions, so the flip is MLX's GPU
precision, not the port.
Suite with both checkpoints present: 100 passed.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
Both remaining gaps are closed, so the status table gains three rows and "What is left" loses two bullets: - the 35b backbone on the real 23 GB checkpoint, layer by layer against the MLX model on its CPU device: <= 2.4e-6 (GatedDeltaNet), <= 3.7e-6 (gated full attention), <= 2.4e-6 on the logits, same argmax, 40 layers with no missing or unexpected tensor; - the 35b engine on the real weights (LoRA 310 targets, prerouter 33 heads, streamed experts): identical 32 greedy tokens to MLX-CPU. MLX on Metal differs from BOTH torch and MLX-CPU at the same three positions, so that flip is its GPU precision, not the port; - the 35b engine on the GB10 (torch 2.14+cu130, as a Slurm job): identical 32 tokens to MLX-CPU and to torch on the Mac's CPU, 18.4 s. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
The 35b GPU evidence was engine-level (identical 32 greedy tokens). This adds the per-layer check the 8b already had on that GPU: each torch layer fed MLX's own input for that layer, on the real weights, float32. Against MLX on the Apple CPU device: 1.6e-6 max per layer (median 2.7e-8), 5.4e-7 across the 30 GatedDeltaNet layers and 1.6e-6 across the 10 gated full-attention ones, 5.1e-7 on the logits, same argmax. No missing or unexpected tensor, 40 streamed expert layers. TF32 off, float32_matmul_precision "highest". Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017vBR5ebDS5kubR1mdvfLN6
…on a GB10) The backend is a correctness reference and decodes edge0-8b at 316 ms a token on a GB10. A profile says the GPU is idle most of that -- 63 ms of CUDA against 643 ms of CPU across three steps -- so the cost is launches and copies: - 1681 host-to-device copies a token: the streaming layer rebuilds expert bundles from the mmap every step. The shared LRU already covers this; cache_slots=3072 instead of the tier's 64 keeps them resident: 206 ms a token, 1.9 GB. Config only, no code change. - 235 quantized-linear calls a token (103 ms) dequantizing the same attention and lm_head weights each time. EDGE0_TORCH_WEIGHT_CACHE=1 keeps the dequantized weight: 124 ms a token, 6.0 GB resident. Together 2.5x (3.2 -> 8.1 tok/s), same greedy tokens. Off by default: the point of the quantized layers is that the 4-bit payload is what stays resident, which is exactly right on the Mac this runs on. test_quantized_linear_weight_cache_is_exact pins the cached path to the chunked one bit-for-bit on the CPU, on a projection wider than one chunk, and both to MLX (on its CPU stream). On CUDA the unchunked weight changes cuBLAS's split, so the two differ within the engine's own noise: teacher-forced they choose the same token at all 32 steps and sit the same distance from MLX (4.8% vs 5.3% median per step). Suite: 100 passed, 3 skipped, with the cache off and on. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017vBR5ebDS5kubR1mdvfLN6
Keeps PR Edge0-AI#19 mergeable against upstream ba4597b. Resolving the merge needed more than the one textual conflict: - streaming/layer.py: the only textual conflict, a comment; upstream wording taken. - prerouter/stager.py: upstream's batched-head path called ``feats.astype(dtype)``, an MLX-only array method; it is now ``core.astype`` so the stager runs on the torch backend. - backends/cuda/core.py: add ``einsum``, which the stager now uses; listed in the backend contract. Checked against MLX by ``test_stacked_head_einsum`` (the stager's own concatenate/astype/einsum sequence). - backends/cuda/_impl/qwen3_5_moe.py: mirror upstream's ``QWEN_HIDDEN_CLIP`` (per-layer clamp + NaN scrub in the vendored MLX qwen3_5), which the edge0-35b engine now turns on at 1000 by default. ``test_qwen35_port_hidden_clip_matches_mlx`` sets the clip below the model's peak activation so the clamp is proven to engage, then checks the torch port against MLX with it off and on. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
The torch backend's decode was CPU-bound, and the CPU was mostly waiting. Profiling edge0-8b decode on Metal (`EDGE0_TORCH_DEVICE=mps`): of an 8.2 s run, `nonzero` took 2.4 s, `tolist` 1.9 s and `unique` 1.4 s -- all inside `gather_qmm`, against 0.07 s for the matmuls they were arranging. None of them is arithmetic; each is a device->host sync. The loop paid one per expert to ask which experts a call needs and where each one's rows are. Decode never needs to ask: one token routes to K experts, so the call is a handful of rows. Gather those rows' experts as one batch and run a single `bmm` -- no sync, and a repeated expert only costs a duplicate dequantize. The per-expert loop still runs above that size (prefill), where a dequantized copy per row is what blew memory up in the first place, but now behind ONE host transfer of the index list instead of a `nonzero` per expert (faster at every size measured, on CPU and Metal alike). The cap is a row count AND a byte budget (`EDGE0_TORCH_GATHER_BATCH`, `EDGE0_TORCH_GATHER_BYTES`), because one row of a 35b expert is far bigger than an 8b one. edge0-8b decode on Metal, same 12 greedy tokens throughout: 688 ms/token (tier default) -> 401 batched -> 384 with `cache_slots=3072` -> 291 with `EDGE0_TORCH_WEIGHT_CACHE=1`. 2.4x. The two paths agree to float32 matmul-reordering noise (~1e-5 absolute), not bit-for-bit: `bmm` and `matmul` sum in a different order. Every gather_qmm case in tests/test_cuda_backend.py now runs through BOTH paths and checks them against each other as well as against MLX -- the shapes there are all decode-sized, so the loop would otherwise have gone untested. The GB10 figures in docs/nvidia.md predate this and still need re-measuring: the Spark's GPU has been fenced off since the 13-Sep reboot (`/dev/nvidia-uvm` EPERM, inside the Slurm job as well as on the host). Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX
Upstream Edge0-AI#102 turned on incr_stack + incr_writeback + warm_willneed in the staged_k4 preset (edge0-35b). The in-place slot writes had never run under torch; the streaming parity case now stages two overlapping sets and checks the persistent stack against the rebuilt staged stack on MLX and torch, before and after churn. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01TvZvF47Ebb9fpuX3kXL3yo
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Unresolved critical correctness, dependency, hygiene, and security issues remain.
Review effort: Lite
Findings: 8
Open (10)
Package and document the CUDA backend's torch dependency · New Forward prerouter_cache through Ling8BEngine · New Normalize negative axes before indexing and reshaping · New Require explicit opt-in for tokenizer remote code execution · New Avoid MLX imports that violate the repository hygiene test · New Move the MLX import out of shared CUDA code · New Prevent dynamic MLX imports from breaking hygiene checks · New Gate gather_qmm smoke tests on MLX CUDA support · New Update requirements and install docs for the CUDA backend · New Make the end-to-end smoke test perform generation · New
What changed in this PR
Adds a torch-backed CUDA reference backend for edge0-8b and edge0-35b while preserving MLX as the default.
Changes:
- Adds CUDA core, NN, quantization, I/O, and model implementations.
- Makes shared engine, streaming, prerouter, routing, sampling, and LoRA code backend-aware.
- Adds parity tests, NVIDIA smoke tests, and CUDA documentation.
| File | Summary |
|---|---|
tests/test_repo_hygiene.py |
Updates repository hygiene checks. |
tests/test_nvidia_backend_smoke.py |
Adds NVIDIA/MLX smoke coverage. |
tests/test_cuda_backend.py |
Tests CUDA backend operations and model ports. |
tests/test_backend_parity.py |
Adds cross-backend parity tests. |
tests/backend_parity_worker.py |
Runs parity cases in subprocesses. |
src/edge0/streaming/layer.py |
Routes streaming operations through the backend facade. |
src/edge0/streaming/install.py |
Supports torch module replacement and adapters. |
src/edge0/sampling.py |
Makes sampling backend-neutral. |
src/edge0/prerouter/stager.py |
Uses backend-compatible casting. |
src/edge0/prerouter/install.py |
Dispatches prerouter loading and patching by backend. |
src/edge0/prerouter/heads.py |
Uses backend-compatible casting. |
src/edge0/moe/routing.py |
Uses backend-compatible routing operations. |
src/edge0/engine/qwen.py |
Dispatches Qwen loading by backend. |
src/edge0/engine/ling.py |
Dispatches Ling loading by backend. |
src/edge0/engine/base.py |
Adds backend capability checks. |
src/edge0/backends/mlx/quant.py |
Exposes MLX quantization helpers. |
src/edge0/backends/mlx/core.py |
Adds MLX facade compatibility helpers. |
src/edge0/backends/mlx/backend.py |
Assembles the MLX backend facade. |
src/edge0/backends/cuda/quant.py |
Implements torch quantized operations. |
src/edge0/backends/cuda/nn.py |
Implements torch neural-network modules. |
src/edge0/backends/cuda/moe_blocks.py |
Adapts streamed experts to MoE blocks. |
src/edge0/backends/cuda/model_specs.py |
Defines CUDA MoE specifications. |
src/edge0/backends/cuda/io.py |
Loads CUDA models, tokenizers, and tensors. |
src/edge0/backends/cuda/core.py |
Implements the torch array-operation facade. |
src/edge0/backends/cuda/backend.py |
Defines the CUDA backend implementation. |
src/edge0/backends/cuda/_impl/qwen3_5_moe.py |
Adds the torch Qwen model port. |
src/edge0/backends/cuda/_impl/bailing_hybrid.py |
Adds the torch Bailing/Ling model port. |
src/edge0/backends/cuda/_impl/__init__.py |
Initializes CUDA model ports. |
src/edge0/backends/cuda/__init__.py |
Exports CUDA backend components. |
src/edge0/backends/__init__.py |
Enables backend selection. |
src/edge0/adapters/lora.py |
Makes LoRA loading backend-aware. |
README.md |
Links NVIDIA support documentation. |
docs/nvidia.md |
Documents CUDA findings, validation, and limitations. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
- load_tokenizer / load_model / _resolve_model_class take trust_remote_code (default False); the Bailing generic path refuses to run checkpoint-supplied Python unless it is passed explicitly. - pyproject: optional `cuda` extra declaring torch>=2.14. - README: requirements and install instructions no longer say CUDA is on the roadmap; document the NVIDIA install path. - tests/test_remote_code_policy.py keeps the opt-in from regressing. Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01QDHVY8E8jmoH9wiyCM8yT6
- cuda/core.py take(): the result shape put the index shape first, which is only right for axis=0; numpy/MLX place it at the axis. Normalize negative axes too. All current callers use axis 0 or None, so this was latent. tests/test_cuda_core_take.py checks every axis vs numpy. - bailing_hybrid.Model.forward forwards before_layer_cb, async_eval_per_layer and prerouter_cache to BailingModel. - smoke tests: gather_qmm runs on Metal only (unsupported on pinned MLX CUDA); the tokenizer-only check is no longer named or marked as an end-to-end generation test. Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01QDHVY8E8jmoH9wiyCM8yT6
Upstream moved the Python framework from src/ and tests/ to python/src/ and python/tests/. The cuda backend package and the files this branch adds (mlx/core.py, parity, cuda, take, smoke and remote-code tests) move with it. README: take the new structure (Requirements and Install now live under "Python (macOS - Apple Silicon)", Design is "###") and re-apply this branch's wording: Torch backend in Requirements and in the backend isolation bullet, optional `cuda` extra, and an NVIDIA / CUDA install subsection. The upstream Roadmap entry for the CUDA backend is left as is. Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01QDHVY8E8jmoH9wiyCM8yT6
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.


What
A torch backend for
edge0, behindEDGE0_BACKEND=cuda, and the two model ports theengines need. With it, both engines — edge0-8b and edge0-35b — generate text on an NVIDIA
GPU, checked against MLX at every level.
MLX's own CUDA backend does not get there at any version available today: the first
forward pass fails, differently at each version (
QMM NYIat 0.30.4, noGatherQMMat0.31.1,
SmallVector out of rangeinsidecore.evalat 0.32.x).docs/nvidia.mdkeepsthat investigation so the next person doesn't repeat it.
Nothing changes for MLX users:
backends/mlx/is untouched, the backend facade picks theimplementation, and
EDGE0_BACKENDdefaults to MLX.Scope. This is a slow, eager, correctness-first reference, not the CUDA offload
redesign discussed in #18 and #20, and it keeps the existing streaming/mmap structure.
Its intended role is a second implementation to diff a faster CUDA path against, layer by
layer, on any machine with torch. Whether and how it fits is the question in #114; it can
be split if that is easier to review (see below).
Layout
backends/cuda/{core,nn,io,quant}.py— the backend contract on torch: MLX-named arrayops,
RMSNorm/QuantizedLinear/QuantizedEmbeddingon the MLX quantization layout,gather_qmm, and aload_modelthat reads the published (MLX-format) checkpoints.backends/cuda/_impl/bailing_hybrid.py,_impl/qwen3_5_moe.py— torch ports of thevendored MLX backbones, same module tree, parameter names and call signatures, so
engine/ling.pyandengine/qwen.pyrun unchanged.engine/,prerouter/,streaming/— the small changes that let shared code run oneither backend (no direct
mlximports, backend-dispatched model classes,class-level prerouter patch on the active backend). These are the only edits to code
MLX users run.
How it is checked
MLX on the CPU is the reference throughout: on some Apple GPUs MLX's float32
matmul/SDPA sits ~7e-4 from float64, while MLX-CPU and torch both sit ~2e-7, so a GPU
reference makes a correct port look worse than it is. The tests run each case under both
backends in subprocesses (
tests/test_backend_parity.py,tests/test_cuda_backend.py):gather_qmm(2/4/8-bit affine, every broadcast shape the streaming layer uses), through both of its code pathsmx.gather_qmmRMSNorm,gelu, the prerouter's batched head at their real call sitesStreamingSwitchGLU, all four paths, on layer 1 of the real edge0-8b checkpointbailing_hybridport, real edge0-8b weights, every layer, chunked prefill + decodeqwen3_5_moeport, real edge0-35b weights (19.5 GB), all 40 layers, 0 missing / 0 unexpected paramspytest -m slow)sm_121, torch 2.14+cu130), float32, each layer fed MLX's own inputThe tests carry negative controls (the transformers norm convention, a wrong RoPE layout,
the prerouter disabled) so they are known to be discriminative rather than merely green.
Two notes on the token comparisons:
per-step logit distance from MLX — the staged prerouter path drops experts it did not
stage — and the top-2 margin at that step is below that noise.
the same three positions, so that is MLX-GPU tie-breaking, not the port.
EDGE0_TORCH_DEVICE(cpu/mps/cuda) overrides device selection; running thesuite with
mpsis what catches host/device mistakes on a machine that has no CUDA.Speed
Still not a fast path, but no longer needlessly slow. Profiling edge0-8b decode on the
torch backend showed the CPU mostly waiting: of an 8.2 s run, 5.7 s were
nonzero/tolist/uniqueinsidegather_qmm, each a device→host sync, one per expert, against0.07 s for the matmuls they were arranging. Decode-sized calls now gather the rows'
experts as one batch and run a single
bmm(the per-expert loop remains for prefill,behind one host transfer; capped by rows and by bytes so a 35b expert cannot blow memory).
edge0-8b decode on Metal (
EDGE0_TORCH_DEVICE=mps), same greedy tokens throughout:cache_slots=64)cache_slots=3072EDGE0_TORCH_WEIGHT_CACHE=1The GB10 timings in
docs/nvidia.mdpredate this change and have not been re-measured.Status and limits
gather_qmmand the quantized linears still dequantize on every call, andcore.compileis eager. No fused dequantize+matmul, no graph capture.gather_qmmpaths agree to float32 summation-order noise (~1e-5), notbit-for-bit.
Verified with
EDGE0_8B_MODELandEDGE0_35B_MODELset: 122 passed, 2 skipped on thisbranch merged with current
main, on the torch CPU device. The real-weights 35b testneeds ~22 GB of RAM for both models; without
EDGE0_35B_MODELit skips.Reviewing this in pieces
If the whole thing is too large to read at once, the natural seams are:
backends/cuda/{core,nn,io,quant}and their tests;engine/,prerouter/,streaming/(start here — the onlypart MLX users can feel);
I can split it along those lines on request.
🤖 Generated with Claude Code
https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX