Skip to content

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
Edge0-AI:mainfrom
agourakis82:cuda-backend
Open

agourakis82 wants to merge 32 commits into
Edge0-AI:mainfrom
agourakis82:cuda-backend

Conversation

@agourakis82

@agourakis82 agourakis82 commented Sep 12, 2026 •

Copy link
Copy Markdown
Contributor

What

A torch backend for edge0, behind EDGE0_BACKEND=cuda, and the two model ports the
engines 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 NYI at 0.30.4, no GatherQMM at
0.31.1, SmallVector out of range inside core.eval at 0.32.x). docs/nvidia.md keeps
that investigation so the next person doesn't repeat it.

Nothing changes for MLX users: backends/mlx/ is untouched, the backend facade picks the
implementation, and EDGE0_BACKEND defaults 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 array
    ops, RMSNorm/QuantizedLinear/QuantizedEmbedding on the MLX quantization layout,
    gather_qmm, and a load_model that reads the published (MLX-format) checkpoints.
  • backends/cuda/_impl/bailing_hybrid.py, _impl/qwen3_5_moe.py — torch ports of the
    vendored MLX backbones, same module tree, parameter names and call signatures, so
    engine/ling.py and engine/qwen.py run unchanged.
  • engine/, prerouter/, streaming/ — the small changes that let shared code run on
    either backend (no direct mlx imports, 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):

Piece Against MLX
gather_qmm (2/4/8-bit affine, every broadcast shape the streaming layer uses), through both of its code paths mx.gather_qmm
routing, sampling, RMSNorm, gelu, the prerouter's batched head at their real call sites identical expert choices and sampler masks
StreamingSwitchGLU, all four paths, on layer 1 of the real edge0-8b checkpoint 1.2–1.5% of output scale (about two bf16 ulps)
bailing_hybrid port, real edge0-8b weights, every layer, chunked prefill + decode ≤ 1.5e-6 per layer, ≤ 1.7e-6 on the logits
qwen3_5_moe port, real edge0-35b weights (19.5 GB), all 40 layers, 0 missing / 0 unexpected params ≤ 2.4e-6 (GatedDeltaNet), ≤ 3.7e-6 (full attention), logits ≤ 2.4e-6, same argmax at every step
whole edge0-8b engine (LoRA, prerouter-staged decode, streaming) identical greedy tokens (pytest -m slow)
whole edge0-35b engine (LoRA, 33 prerouter heads, streaming), 32 greedy tokens 32/32 identical to MLX-CPU
edge0-8b backbone on an NVIDIA GB10 (sm_121, torch 2.14+cu130), float32, each layer fed MLX's own input ≤ 5.3e-7 per layer, 4.3e-7 on the logits, same argmax
whole edge0-8b engine on that GPU, 32 greedy tokens 31 of 32 identical to MLX
whole edge0-35b engine on that GPU, 32 greedy tokens 32/32 identical, 18.4 s

The 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:

  • edge0-8b step 31 is not a GPU artifact. Torch on a CPU sits at the same ~4.5% median
    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.
  • edge0-35b against MLX-Metal is 29/32, and Metal disagrees with MLX-CPU at exactly
    the same three positions, so that is MLX-GPU tie-breaking, not the port.

EDGE0_TORCH_DEVICE (cpu / mps / cuda) overrides device selection; running the
suite with mps is 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 / unique inside gather_qmm, each a device→host sync, one per expert, against
0.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:

ms/token
tier default (cache_slots=64) 688
+ batched decode gather 401
+ cache_slots=3072 384
+ EDGE0_TORCH_WEIGHT_CACHE=1 291

The GB10 timings in docs/nvidia.md predate this change and have not been re-measured.

Status and limits

  • Correctness is established for both engines on real weights; performance is not.
    gather_qmm and the quantized linears still dequantize on every call, and
    core.compile is eager. No fused dequantize+matmul, no graph capture.
  • Batched independent contexts and long-context behaviour are untested.
  • The two gather_qmm paths agree to float32 summation-order noise (~1e-5), not
    bit-for-bit.

Verified with EDGE0_8B_MODEL and EDGE0_35B_MODEL set: 122 passed, 2 skipped on this
branch merged with current main, on the torch CPU device. The real-weights 35b test
needs ~22 GB of RAM for both models; without EDGE0_35B_MODEL it skips.

Reviewing this in pieces

If the whole thing is too large to read at once, the natural seams are:

  1. the backend facade + backends/cuda/{core,nn,io,quant} and their tests;
  2. the two model ports;
  3. the shared-code changes in engine/, prerouter/, streaming/ (start here — the only
    part MLX users can feel);
  4. docs.

I can split it along those lines on request.

🤖 Generated with Claude Code

https://claude.ai/code/session_0146Hf9MXz9hHfRnq3MNXXWX

Demetrios Chiuratto Agourakis and others added 25 commits September 11, 2026 07:01
…=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
agourakis82 and others added 2 commits September 14, 2026 15:27
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
@agourakis82 agourakis82 changed the title backends/cuda: a torch backend — edge0-8b generates text on an NVIDIA GPU backends/cuda: a torch reference backend — edge0-8b and edge0-35b generate text on an NVIDIA GPU Sep 21, 2026
agourakis82 and others added 2 commits September 28, 2026 12:30
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
Copilot AI lite review requested due to automatic review settings September 28, 2026 15:36

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot review overview

🟡 Changes recommended

Unresolved critical correctness, dependency, hygiene, and security issues remain.

Review effort: Lite
Findings: 8 High severity · 2 Medium severity

Open (10)
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.

Comment thread python/src/edge0/backends/__init__.py
Comment thread src/edge0/backends/cuda/_impl/bailing_hybrid.py Outdated
Comment thread src/edge0/backends/cuda/core.py Outdated
Comment thread src/edge0/backends/cuda/io.py Outdated
Comment thread python/src/edge0/engine/ling.py
Comment thread python/src/edge0/engine/qwen.py
Comment thread python/src/edge0/prerouter/install.py
Comment thread python/tests/test_nvidia_backend_smoke.py
Comment thread README.md Outdated
Comment thread tests/test_nvidia_backend_smoke.py Outdated
agourakis82 and others added 3 commits October 2, 2026 06:09
- 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
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants