Skip to content

PR #19 (torch reference backend): how does it fit the planned CUDA work, and how would you like it reviewed? #114

Description

@agourakis82

PR #19 adds a torch backend behind EDGE0_BACKEND=cuda (33 files, +4.3k/−82, CI green, currently mergeable, no review yet). This issue is to ask how it fits with the CUDA work you already have planned, and to make it cheap to review or to decline.

Where I think it fits

Per the replies on #18 and #20, you are building a CUDA path yourselves, and it needs a different offload design than the unified-memory mmap one (not a port). #19 is deliberately not that. It is a slow, eager, correctness-first reference: every layer of both models runs on torch and is compared against MLX. Its most useful role may be as an oracle for your CUDA implementation — a second implementation you can diff against layer by layer, on any machine that has torch, instead of only on a Mac. It also answers, with logs, the question in #107 for anyone who wants to run the models today (slowly), and it documents why MLX-CUDA does not work at the versions I could test (docs/nvidia.md: QMM NYI at 0.30.4, no GatherQMM at 0.31.1, SmallVector out of range at 0.32.x).

Nothing changes for MLX users: backends/mlx/ is untouched and EDGE0_BACKEND defaults to MLX.

What it was checked against

MLX on the CPU is the reference. On the Apple GPU, MLX's float32 matmul/SDPA sits ~7e-4 from float64, while MLX-CPU and torch both sit ~2e-7, so a GPU reference makes the port look worse than it is.

result
edge0-8b, real weights, every layer, chunked prefill + decode ≤ 1.5e-6 per layer, ≤ 1.7e-6 logits
edge0-35b, real 19.5 GB weights, 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-35b engine (LoRA, 33 prerouter heads, streaming), 32 greedy tokens 32/32 identical to MLX-CPU
same run vs MLX-Metal 29/32 — Metal disagrees with MLX-CPU at exactly the same 3 positions, so that is MLX-GPU tie-breaking, not the port
GB10 (sm_121, torch 2.14+cu130): 8b per layer / 35b engine ≤ 5.3e-7 per layer; 35b 32/32 identical in 18.4 s
full suite on the branch merged with current main (both checkpoints) 122 passed, 2 skipped

The tests carry negative controls (wrong norm convention, wrong RoPE layout, prerouter off) so they are known to fail when they should.

What it is not

  • Not fast. Decode on the torch backend is CPU-bound. The largest single cost was a device→host sync per expert inside gather_qmm; batching that for decode-sized calls took edge0-8b from 688 → 401 ms/token on Metal (EDGE0_TORCH_DEVICE=mps), 291 with the optional caches, same tokens. The GB10 timings in docs/nvidia.md predate that change and I have not been able to re-measure them yet.
  • Not the offload redesign. It keeps the existing streaming/mmap structure and just runs it on torch tensors.
  • Batched independent contexts and long-context behaviour are untested.

Questions

  1. Is a torch reference backend something you would want in-tree, given your own CUDA work? If not, would you rather it live as a separate repo or as tests/ tooling only?
  2. If yes, is the PR too big to review as one? I can split it along its natural seams: (a) the backend facade + backends/cuda/{core,nn,io,quant} and their tests; (b) the two model ports; (c) the small shared-code changes (engine/, prerouter/, streaming/ no longer importing mlx directly); (d) docs.
  3. Anything in (c) you would want done differently? Those are the only edits to code that MLX users run, so they matter most.

If you only have ten minutes

  1. src/edge0/backends/__init__.py and the diff under engine/, prerouter/, streaming/ — the only changes MLX users can feel.
  2. tests/test_backend_parity.py — how each case runs under both backends in subprocesses.
  3. docs/nvidia.md — the investigation and the numbers above.

Run it: EDGE0_8B_MODEL=... EDGE0_35B_MODEL=... python -m pytest -q (the real-weights 35b test needs ~22 GB RAM for both models; without EDGE0_35B_MODEL it skips).

Related: #18, #20, #107.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions