YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.

grouped-moe-gemv

FlashRT native CUDA grouped expert projection kernels for Blackwell decode and small verify batches. Version 2 adds W4A4 with device-side top-k routing while preserving the version 1 W4A16 APIs.

Hardware Backends

  • SM110 (Jetson AGX Thor): W4A16 decode and grouped expert GEMV use the FlashRT edge backend validated by FlashRT PR #169. This target is compiled independently with FLASHRT_W4A16_EDGE_UNROLL=2; the SM120 value remains 4.
  • SM120/SM121: W4A16 and block-scaled-MMA W4A4 paths are available.
  • W4A4 is intentionally rejected on SM110 because that implementation uses the SM120 block-scaled MMA path. It never silently falls back or launches an incompatible cubin.

Functions

  • w4a16_decode_gemv_bf16(x_bf16, weight_packed, sfb, alpha=1.0, out=None)
  • grouped_w4a16_gemv_bf16(activations, weight_stack, sfb_stack, alpha_stack, expert_idx, n, w_stride=None, sfb_stride=None, out=None)
  • quantize_activations_nvfp4_bf16(activations, packed=None, sfa=None)
  • quantize_weights_nvfp4_bf16(weights, packed=None, sfb=None)
  • grouped_w4a4_gemv_bf16(activations_packed, weight_stack, sfa, sfb_stack, alpha_stack, expert_idx, out=None)
  • grouped_w4a4_gemv_from_bf16(activations, weight_stack, sfb_stack, alpha_stack, expert_idx, packed=None, sfa=None, out=None)

The grouped API runs one BF16-activation x NVFP4-weight GEMV per routed slot. It is intended for static routed expert batches where the caller already owns packed weights and swizzled scale-factor buffers.

On SM120/SM121, the W4A4 API accepts packed activations [M,K/2], expert weights [E,N,K/2], and a contiguous device routing tensor [M,top_k]. It emits [M,top_k,N] in one grouped compute launch. For down projections with a different activation per routed pair, flatten to M=routed_pairs, top_k=1.

K must be divisible by 16 and N by 8. Target K%64==0 shapes use tuned SM120 paths; the remaining K%16 shapes use a fixed-order SIMT contract path. No atomics, host synchronization, or dynamic workspace are used by the native ops. Pass packed, sfa, and out buffers to the composed helper for allocation-free CUDA Graph capture.

Example

from kernels import get_kernel
import torch

try:
    moe = get_kernel(
        "flashrt/grouped-moe-gemv", version=2, trust_remote_code=True
    )
except TypeError:  # kernels==0.12.x compatibility
    moe = get_kernel("flashrt/grouped-moe-gemv", version=2)

M, TOP_K, E, N, K = 7, 8, 8, 1024, 2048
x = torch.randn(M, K, device="cuda", dtype=torch.bfloat16)
expert_idx = torch.randint(E, (M, TOP_K), device="cuda", dtype=torch.int32)

def sf_bytes(rows, dim):
    return ((rows + 127) // 128) * (((dim // 16) + 3) // 4) * 512

# Do this once while loading the checkpoint, not in the inference hot path.
weights_bf16 = torch.randn(E, N, K, device="cuda", dtype=torch.bfloat16)
weights_packed = torch.empty(E, N, K // 2, device="cuda", dtype=torch.uint8)
weight_sfs = torch.empty(E, sf_bytes(N, K), device="cuda", dtype=torch.uint8)
for expert in range(E):
    moe.quantize_weights_nvfp4_bf16(
        weights_bf16[expert],
        packed=weights_packed[expert],
        sfb=weight_sfs[expert],
    )
weight_alpha = torch.ones(E, device="cuda", dtype=torch.float32)

packed = torch.empty(M, K // 2, device="cuda", dtype=torch.uint8)
sfa = torch.empty(sf_bytes(M, K), device="cuda", dtype=torch.uint8)
out = torch.empty(M, TOP_K, N, device="cuda", dtype=torch.bfloat16)
y = moe.grouped_w4a4_gemv_from_bf16(
    x, weights_packed, weight_sfs, weight_alpha, expert_idx,
    packed=packed, sfa=sfa, out=out,
)

The example buffers may be larger than the minimum; wrappers validate storage. For production code, derive SF sizes from the checkpoint packer metadata.

Dispatch guidance

Use the built-artifact benchmark to dispatch rather than selecting only by dtype. On the tested cu128 artifact W4A16 wins gate-up, W4A4 wins down verify, and down decode is effectively tied kernel-only. A fused/upstream FP4 producer removes the standalone quantization charge, but callers still should not assume lower precision is automatically faster.

Validation

python grouped-moe-gemv/tests/test_grouped_moe_gemv.py --backend source --mode full
python grouped-moe-gemv/benchmarks/benchmark.py --backend source
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support