Revert "Precompute RoPE cos/sin in Python, avoid scalar kernel input"

This reverts commit 97d7f03032.
This commit is contained in:
dmcc73
2026-03-15 14:57:02 +00:00
parent 058ccfba24
commit f94e7e5b95
@@ -101,11 +101,13 @@ def _gen_fused_qk_norm_rope_source(H_q=16, H_kv=2, D=256, rope_dims=64):
ushort partner = (ushort)(slid ^ {PARTNER_XOR}u);
int cos_base = (int)(slid & {FIRST_HALF - 1}u) * N_READS;
// Load precomputed cos/sin values (computed in Python from position * inv_freq)
// Compute cos/sin on-device from precomputed inv_freq
// angle = position * inv_freq[d], where inv_freq[d] = theta^(-d/{ROPE_HALF})
float cos_arr[{N_READS}], sin_arr[{N_READS}];
for (int i = 0; i < N_READS; i++) {{
cos_arr[i] = rope_cos[cos_base + i];
sin_arr[i] = rope_sin[cos_base + i];
float angle = (float)position * inv_freq[cos_base + i];
cos_arr[i] = metal::fast::cos(angle);
sin_arr[i] = metal::fast::sin(angle);
}}
// Exchange normalized values with partner thread via simd_shuffle
@@ -148,7 +150,7 @@ def _get_kernel(H_q, H_kv, D, rope_dims):
_kernel_cache[key] = mx.fast.metal_kernel(
name="fused_qk_norm_rope",
input_names=["queries", "keys", "q_norm_w", "k_norm_w",
"rope_cos", "rope_sin"],
"inv_freq", "position"],
output_names=["q_out", "k_out"],
source=_gen_fused_qk_norm_rope_source(H_q, H_kv, D, rope_dims),
)
@@ -166,7 +168,7 @@ def fused_qk_norm_rope(queries, keys, q_norm_weight, k_norm_weight,
q_norm_weight: [D] bf16 — RMSNorm learned weight for queries.
k_norm_weight: [D] bf16 — RMSNorm learned weight for keys.
inv_freq: [rope_dims/2] f32 — precomputed theta^(-d/half_dims).
cache_offset: int|float — sequence position for RoPE angles.
cache_offset: int — sequence position for RoPE angles.
H_q: int — number of query heads.
H_kv: int — number of key/value heads.
D: int — head dimension.
@@ -183,15 +185,11 @@ def fused_qk_norm_rope(queries, keys, q_norm_weight, k_norm_weight,
q_flat = queries.reshape(B, H_q * D)
k_flat = keys.reshape(B, H_kv * D)
# Precompute cos/sin in Python (avoids passing scalar to Metal kernel)
angles = float(cache_offset) * inv_freq # (rope_dims/2,) f32
rope_cos = mx.cos(angles)
rope_sin = mx.sin(angles)
pos = mx.array(cache_offset, dtype=mx.int32)
n_heads = H_q + H_kv
results = kern(
inputs=[q_flat, k_flat, q_norm_weight, k_norm_weight, rope_cos, rope_sin],
inputs=[q_flat, k_flat, q_norm_weight, k_norm_weight, inv_freq, pos],
output_shapes=[(B * H_q * D,), (B * H_kv * D,)],
output_dtypes=[mx.bfloat16, mx.bfloat16],
grid=(n_heads * 32, 1, B),