Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 21 additions & 11 deletions kernels/attention/fused_rope_cache_kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,12 +73,19 @@ def build_fused_rope_cache_module(
raise ValueError(f"dtype_str must be 'bf16' or 'f16', got {dtype_str!r}")
half_dim = rotary_dim // 2

# VEC_WIDTH: elements per thread. Use ceil division so vecs_per_head never
# exceeds WARP_SIZE for the fixed one-thread-per-vector mapping below.
# For D=64: VEC_WIDTH=1 -> vecs_per_head=64 (full wavefront, 16-bit loads).
# For D=96: VEC_WIDTH=2 -> vecs_per_head=48 (fits within one wavefront).
# For D=128: VEC_WIDTH=2 -> vecs_per_head=64 (32-bit loads, unchanged).
VEC_WIDTH = max(1, (head_dim + WARP_SIZE - 1) // WARP_SIZE)
# Select the smallest native copy width that fits one head in a wave and
# evenly divides each NeoX half. On wave32 this maps D=64/96/128/256 to
# VEC_WIDTH=2/4/4/8; on wave64 it maps them to 1/2/2/4.
VEC_WIDTH = next(
(
width
for width in (1, 2, 4, 8)
if head_dim % width == 0 and half_dim % width == 0 and head_dim // width <= WARP_SIZE
),
None,
)
if VEC_WIDTH is None:
raise ValueError(f"unsupported head_dim={head_dim} for wave{WARP_SIZE}")

vecs_per_half = half_dim // VEC_WIDTH
vecs_per_head = head_dim // VEC_WIDTH
Expand Down Expand Up @@ -149,10 +156,8 @@ def store_vec(val, div_tensor, idx, atom=None):
fx.copy(atom or copy_atom, r, fx.slice(div_tensor, (None, idx)))

# Helper: get the rotary-pair element via ds_bpermute (LDS cross-lane shuffle).
# For NeoX RoPE, the pair of thread tid is tid XOR vecs_per_half.
# ds_bpermute: thread tid reads the VGPR value held by thread (pair_byte_addr/4).
# pair_byte_addr = (tid XOR vecs_per_half) * 4.
# Handles VEC_WIDTH=1 (vector<1xbf16/f16>, 16-bit) and VEC_WIDTH=2 (vector<2xbf16/f16>, 32-bit).
# Handles native vectors from one to eight bf16/f16 elements.
def ds_bpermute_pair(vec_val, pair_byte_addr):
"""Return the copy of vec_val held by the rotary-pair thread, via ds_bpermute."""
if const_expr(VEC_WIDTH == 1):
Expand Down Expand Up @@ -190,9 +195,14 @@ def ds_bpermute_pair(vec_val, pair_byte_addr):
is_first_half = tid < vecs_per_half
cos_vec_idx = tid % vecs_per_half if reuse_freqs_front_part else tid

# Pair lane for ds_bpermute: tid XOR vecs_per_half (symmetric, works for both halves).
# XOR is the cheapest symmetric mapping when the half-wave span is
# a power of two. Other spans, such as D=96 on wave32 (12 lanes per
# half), require explicit addition/subtraction.
if const_expr((vecs_per_half & (vecs_per_half - 1)) == 0):
pair_lane = tid ^ vecs_per_half
else:
pair_lane = is_first_half.select(tid + vecs_per_half, tid - vecs_per_half)
# pair_byte_addr = pair_lane * 4 (ds_bpermute address unit is bytes, VGPR = 4 bytes).
pair_lane = tid ^ vecs_per_half
pair_byte_addr = pair_lane * 4

# --- Shared cos/sin (loaded once, used by both Q and K) ---
Expand Down
20 changes: 20 additions & 0 deletions tests/kernels/test_fused_rope_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -625,6 +625,26 @@ def test_f16(num_tokens, flash_layout):
assert passed, f"FAILED (T={num_tokens} flash={flash_layout}): {errs}"


# ===========================================================================
# Category 4b: Non-power-of-two head dimension
# ===========================================================================


@pytest.mark.parametrize("dtype_str", ["bf16", "f16"])
@pytest.mark.parametrize("flash_layout", [True, False], ids=["flash", "nonflash"])
def test_head_dim_96(dtype_str, flash_layout):
"""D=96 exercises a non-power-of-two NeoX half-wave span."""
passed, errs = run_test(
num_tokens=32,
head_dim=96,
num_q_heads=32,
num_kv_heads=32,
flash_layout=flash_layout,
dtype_str=dtype_str,
)
assert passed, f"FAILED (dtype={dtype_str} flash={flash_layout}): {errs}"


# ===========================================================================
# Category 5: pos_dtype — i32 vs i64 (stride-2 indexing)
# ===========================================================================
Expand Down