liangsu9988's picture
Promote latest kernel artifacts to main
a1d71b9 verified
|
Raw
History Blame Contribute Delete
2.88 kB

transformer-layout-primitives

Generic BF16 layout, RoPE, and text-state primitives for transformer pipelines.

Hub repo: flashrt/transformer-layout-primitives

Public API

  • fill_neginf_bf16(dst) -> dst
  • add_bias_bf16_(data, bias) -> data
  • repeat_interleave_heads_bf16(src, repeat) -> bf16
  • text_gather_bf16(src, batch, seq) -> bf16
  • text_scatter_bf16(dst, src, batch, seq) -> dst
  • rope_rotate_half_bf16_(x, cos, sin) -> x
  • qk_rmsnorm_rope_bf16_(qk, weight, cos, sin, eps=1e-6) -> qk
  • qk_pair_rmsnorm_rope_bf16(q, k, q_weight, k_weight, cos, sin, eps=1e-6) -> (q, k)
  • gather_rows_bf16(src, row_indices, out=None) -> bf16
  • scatter_rows_bf16(src, row_indices, rows, out=None) -> bf16

The package is intentionally model-neutral. It exposes Tensor APIs for common transformer integration gaps: head repeat for GQA/MQA, first/last token gather and scatter, bias add, RoPE rotate-half, and fused Q/K RMSNorm+RoPE. The pair API handles different Q and KV head counts in one launch.

Example

from kernels import get_kernel
import torch

ops = get_kernel("flashrt/transformer-layout-primitives", version=1)

q = torch.randn((128, 32, 128), device="cuda", dtype=torch.bfloat16)
weight = torch.ones((128,), device="cuda", dtype=torch.bfloat16)
cos = torch.randn((128, 128), device="cuda", dtype=torch.bfloat16)
sin = torch.randn((128, 128), device="cuda", dtype=torch.bfloat16)
ops.qk_rmsnorm_rope_bf16_(q, weight, cos, sin)

k = torch.randn((128, 8, 128), device="cuda", dtype=torch.bfloat16)
q, k = ops.qk_pair_rmsnorm_rope_bf16(
    q, k, weight, weight, cos, sin
)

Shape contract

  • All tensors are contiguous CUDA BF16 tensors.
  • repeat_interleave_heads_bf16: src is (seq, heads, head_dim).
  • text_gather_bf16: src is flattened (batch * seq, dim) and returns first and last token rows as (2 * batch, dim).
  • text_scatter_bf16: writes (2 * batch, dim) rows back to first and last positions in (batch * seq, dim).
  • RoPE functions use rotate-half layout with cos/sin shaped (seq, head_dim) or (rows, head_dim).
  • qk_pair_rmsnorm_rope_bf16 accepts Q (rows, q_heads, head_dim) and K (rows, kv_heads, head_dim). head_dim must be even and in [8, 256]. Q and K may have different head counts.
  • gather_rows_bf16 and scatter_rows_bf16 use contiguous CUDA int64 row indices. Scatter indices must be unique.
  • Indexed-row validation includes the Cosmos3-Edge production layout 128 -> 60 rows with hidden size 2048, including exact CUDA Graph replay.

Validation

Correctness is tested against PyTorch BF16/FP32 reference formulas with exact checks for pure layout operations and strict BF16 tolerances for math ops. The pair API must additionally match the two staged single-tensor kernels exactly.

See benchmarks/RESULTS.md for current local RTX 5090 source benchmark data.