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.

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