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

Check out the documentation for more information.

ragged-dot-tpu

MoE experts forward for torch_tpu, built on tokamax's Pallas grouped matmul (tokamax.ragged_dot, implementation="mosaic_tpu_v2") and wrapped with torch_tpu's jax_op. experts_forward is a drop-in transformers experts implementation for gated experts without bias, with a SiLU (e.g. Qwen3-MoE) or GELU-tanh (e.g. Gemma 4) gate, and supports expert parallelism. The whole experts forward runs as one op.

from kernels import get_kernel
from transformers import AutoModelForCausalLM
from transformers.integrations.moe import ALL_EXPERTS_FUNCTIONS

repo_id = "tengomucho/ragged-dot-tpu"
kernel = get_kernel(repo_id, version=1, trust_remote_code=[repo_id])
ALL_EXPERTS_FUNCTIONS.register("tokamax_ragged_dot", kernel.experts_forward)
model = AutoModelForCausalLM.from_pretrained(
    "Qwen/Qwen3-30B-A3B-Instruct-2507",
    experts_implementation="tokamax_ragged_dot",
)

Requires torch_tpu, jax and tokamax.

On a TPU v6e-8 (expert parallelism over 8 cores), against transformers' default grouped_mm:

  • Qwen3-30B-A3B-Instruct-2507 decodes 1.13x faster at batch size 1 and 21.5x faster at batch size 64 (988 against 52 tokens/s).
  • gemma-4-26B-A4B-it decodes 1.07x faster at batch size 1 and 12.8x faster at batch size 64 (162 against 13 tokens/s).
Downloads last month
-