YAML Metadata Warning:empty or missing yaml metadata in repo card
Check out the documentation for more information.
masked-mha-runtime
FlashRT native SM110 masked FP16/BF16 MHA for fixed-shape CUDA Graph runtimes. It masks padded logits during softmax, supports fused-QKV token strides, and keeps caller-owned logits/output buffers stable across graph replay.
from kernels import get_kernel
import torch
ops = get_kernel("flashrt/masked-mha-runtime", version=1)
logits = ops.allocate_workspace(q, k)
out = torch.empty_like(q, memory_format=torch.contiguous_format)
ops.forward_static(q, k, v, logits=logits, out=out)
Public functions are allocate_workspace, forward_static, and forward.
Inputs use (sequence, heads, head_dim). The production GROOT gate covers
DiT (41, 32, 48), ViT/LLM sequence lengths 277/1024, padded boundaries
1025/2048, FP16 and BF16, fused strides, and bitwise CUDA Graph replay.
See CARD.md for the complete contract. Source provenance is FlashRT
commit 24df793f4fa2d50780aea03b644208c6e0cb4162.
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support