Issues / #1357

#1357 mtp: the per-token drafter router call is missing the n_expert == 512 guard its three sibling call sites carry

closed · @rwkeyes · 1 comments · View on GitHub

NVIDIA / CUDAModels & quants

Description

While porting this engine I hit a guard asymmetry in the MTP drafter's router call and wanted to flag it. Found by reading, not by reproducing a crash.

## The callee's documented requirement

`include/strata/kernels/native_router.hpp:10-16` states the contract explicitly:

> Pinned CUDA topk-moe contract for ONE token, **512 experts**, top 10, softmax, no selection bias, lower normalization clamp 2^-14, and scale 1. **Reads 512 finite F32 logits**; writes 10 I32 IDs and 10 F32 weights.

So `native_router_top10` requires a 512-wide logits row, and the CUDA implementation hardcodes that: `src/kernels/cuda/native_router.cu:49` does `logits += (size_t) blockIdx.x * 512;`, and `:163` validates the spans as `512 * 4` bytes.

## The guard asymmetry

Three of the four call sites test the expert count first. One does not:

| file:line | guard |
|---|---|
| `src/core/layer.cpp:372` | `native_router_enabled() && g.n_expert == 512 && k == 10` |
| `src/core/verify.cpp:1110` | `dec_batch && n > 1 && w_router != nullptr && native_router_enabled() && NE == 512 && K == 10` |
| `src/core/mtp.cpp:877` | `T > 1 && native_router_enabled() && g.n_expert == 512 && K == 10` |
| **`src/core/mtp.cpp:883`** | **`if (native_router_enabled())`** — nothing else |

```
877:  if (T > 1 && native_router_enabled() && g.n_expert == 512 && K == 10) {
878:      ...
879:      native_router_top10_multi(logits_, ids_, w_, T, cs);
      } else {
880-882:   (loop over t)
883:      if (native_router_enabled()) native_router_top10(logits_ + t * g.n_expert, ids_ + t * K, w_ + t * K, cs);
```

The `else` at `:883` is reached precisely when the `:877` test fails — that is, **when the model does not have 512 experts** — and it then passes a row of width `g.n_expert` to a callee that reads 512 floats from it. The correct guard sits eight lines above, which makes this look like an oversight rather than intent.

## What it does

For `g.n_expert < 512` the call reads past the end of each row: for every token but the last, into the next token's logits; for the last row, past the allocation. Wherever that memory is mapped, the values are another token's logits, so the router's selection and weights are silently wrong for the drafter. That makes this a correctness defect, not only a bounds one. A 256-expert model — which the drafter otherwise supports, since `mtp.cpp:296` sizes its expert store from `g.n_expert` — over-reads 256 floats (1 KB) per row.

## The SYCL tree has it identically

`sycl/src/core/mtp.cpp:679`:

```c
if (native_router_enabled()) native_router_top10(logits_ + t * g.n_expert, ids_ + t * K, w_ + t * K, cs);
```

No `n_expert == 512` guard there either.

## Scope, stated honestly

This is latent, and I want to be precise about that: the header itself says the path is an experiment that is **off by default** (and it does not modify the legacy router), and I could not find any caller of `native_router_set_enabled` outside the router's own translation units. So a stock session will not reach it. Reaching it takes `native_router_enabled()` true, a model whose `n_expert != 512`, and the MTP drafter enabled (`--mtp-*`). I could not run that configuration — I found this by reading the source and checking the callee's contract, not by hitting a failure.

## Suggested fix

Mirror the `:877` condition on the `else`, and use the geometry-aware entry, which is documented as handling exactly this and as producing identical results where both apply (`include/strata/kernels/router_top10.hpp:17,20-21` — `router_top10(logits, n_tokens, n_expert, k, ids, weights, stream)`, with the ARM fast path being "the default there for 64 < n_expert <= 512, k <= 32; bitwise the same results"):

```c
} else if (native_router_enabled() && g.n_expert == 512 && K == 10) {
    native_router_top10(logits_ + t * g.n_expert, ids_ + t * K, w_ + t * K, cs);
} else {
    router_top10(logits_ + t * g.n_expert, 1, (int) g.n_expert, (int) K, ids_ + t * K, w_ + t * K, cs);
}
```

Relatedly, `src/kernels/native_multi_parity.cpp` currently exercises `native_router_top10_multi` against the single-token call for `n = 1..19`; a non-512 case there would have caught this, since the single-token call is the one that over-reads.

Happy to send a patch if that would help. Context for how I hit it: I maintain a Vulkan port of these kernels for Intel Arc hardware, and needed the geometry-aware router entry to implement this path correctly.

Related on strata.com

Editorial links to help you install, pick models, or read release notes — not part of the upstream thread.