Skip to content

Commit bc38689

Browse files
authored
Add interleave mode for split qkv norm rope (#503)
* Add interleave mode for split qkv norm rope * lint checks
1 parent 5c5158b commit bc38689

1 file changed

Lines changed: 160 additions & 52 deletions

File tree

python/sgl_kernel_npu/sgl_kernel_npu/norm/split_qkv_rmsnorm_rope.py

Lines changed: 160 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ def split_qkv_rmsnorm_rope_kernel(
3131
HALF_ROPE_DIM: tl.constexpr,
3232
PASS_DIM: tl.constexpr,
3333
DO_PARTIAL: tl.constexpr,
34+
DO_HALF: tl.constexpr,
3435
):
3536
row_pid = tl.program_id(0)
3637
col_pid = tl.program_id(1)
@@ -74,6 +75,28 @@ def split_qkv_rmsnorm_rope_kernel(
7475
sc_offsets = row_idx * ROPE_DIM + tl.arange(0, ROPE_DIM)
7576
sin = (tl.load(sin_ptr + sc_offsets)).reshape(1, ROPE_DIM)
7677
cos = (tl.load(cos_ptr + sc_offsets)).reshape(1, ROPE_DIM)
78+
if not DO_HALF:
79+
sin_half = al.extract_slice(
80+
sin, offsets=(0, 0), sizes=(1, HALF_ROPE_DIM), strides=(1, 1)
81+
)
82+
cos_half = al.extract_slice(
83+
cos, offsets=(0, 0), sizes=(1, HALF_ROPE_DIM), strides=(1, 1)
84+
)
85+
sin = tl.zeros((1, ROPE_DIM), dtype=sin_half.dtype)
86+
sin = al.insert_slice(
87+
sin, sin_half, offsets=(0, 0), sizes=(1, HALF_ROPE_DIM), strides=(1, 2)
88+
)
89+
sin = al.insert_slice(
90+
sin, sin_half, offsets=(0, 1), sizes=(1, HALF_ROPE_DIM), strides=(1, 2)
91+
)
92+
93+
cos = tl.zeros((1, ROPE_DIM), dtype=cos_half.dtype)
94+
cos = al.insert_slice(
95+
cos, cos_half, offsets=(0, 0), sizes=(1, HALF_ROPE_DIM), strides=(1, 2)
96+
)
97+
cos = al.insert_slice(
98+
cos, cos_half, offsets=(0, 1), sizes=(1, HALF_ROPE_DIM), strides=(1, 2)
99+
)
77100
if DO_PARTIAL:
78101
rot_x = al.extract_slice(
79102
normalized_values,
@@ -89,33 +112,63 @@ def split_qkv_rmsnorm_rope_kernel(
89112
)
90113
else:
91114
rot_x = normalized_values
92-
x1 = al.extract_slice(
93-
rot_x,
94-
offsets=(0, 0),
95-
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
96-
strides=(1, 1),
97-
)
98-
x2 = al.extract_slice(
99-
rot_x,
100-
offsets=(0, HALF_ROPE_DIM),
101-
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
102-
strides=(1, 1),
103-
)
115+
if DO_HALF:
116+
x1 = al.extract_slice(
117+
rot_x,
118+
offsets=(0, 0),
119+
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
120+
strides=(1, 1),
121+
)
122+
x2 = al.extract_slice(
123+
rot_x,
124+
offsets=(0, HALF_ROPE_DIM),
125+
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
126+
strides=(1, 1),
127+
)
128+
else:
129+
x1 = al.extract_slice(
130+
rot_x,
131+
offsets=(0, 0),
132+
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
133+
strides=(1, 2),
134+
)
135+
x2 = al.extract_slice(
136+
rot_x,
137+
offsets=(0, 1),
138+
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
139+
strides=(1, 2),
140+
)
104141
cat_x = tl.zeros((Q_BLOCK_SIZE // HEAD_DIM, ROPE_DIM), dtype=tl.bfloat16)
105-
cat_x = al.insert_slice(
106-
cat_x,
107-
-x2,
108-
offsets=(0, 0),
109-
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
110-
strides=(1, 1),
111-
)
112-
cat_x = al.insert_slice(
113-
cat_x,
114-
x1,
115-
offsets=(0, HALF_ROPE_DIM),
116-
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
117-
strides=(1, 1),
118-
)
142+
if DO_HALF:
143+
cat_x = al.insert_slice(
144+
cat_x,
145+
-x2,
146+
offsets=(0, 0),
147+
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
148+
strides=(1, 1),
149+
)
150+
cat_x = al.insert_slice(
151+
cat_x,
152+
x1,
153+
offsets=(0, HALF_ROPE_DIM),
154+
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
155+
strides=(1, 1),
156+
)
157+
else:
158+
cat_x = al.insert_slice(
159+
cat_x,
160+
-x2,
161+
offsets=(0, 0),
162+
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
163+
strides=(1, 2),
164+
)
165+
cat_x = al.insert_slice(
166+
cat_x,
167+
x1,
168+
offsets=(0, 1),
169+
sizes=(Q_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
170+
strides=(1, 2),
171+
)
119172
roped_q = cat_x * sin + rot_x * cos
120173
if DO_PARTIAL:
121174
normalized_values = al.insert_slice(
@@ -182,6 +235,29 @@ def split_qkv_rmsnorm_rope_kernel(
182235
sin = (tl.load(sin_ptr + sc_offsets)).reshape(1, ROPE_DIM)
183236
cos = (tl.load(cos_ptr + sc_offsets)).reshape(1, ROPE_DIM)
184237

238+
if not DO_HALF:
239+
sin_half = al.extract_slice(
240+
sin, offsets=(0, 0), sizes=(1, HALF_ROPE_DIM), strides=(1, 1)
241+
)
242+
cos_half = al.extract_slice(
243+
cos, offsets=(0, 0), sizes=(1, HALF_ROPE_DIM), strides=(1, 1)
244+
)
245+
sin = tl.zeros((1, ROPE_DIM), dtype=sin_half.dtype)
246+
sin = al.insert_slice(
247+
sin, sin_half, offsets=(0, 0), sizes=(1, HALF_ROPE_DIM), strides=(1, 2)
248+
)
249+
sin = al.insert_slice(
250+
sin, sin_half, offsets=(0, 1), sizes=(1, HALF_ROPE_DIM), strides=(1, 2)
251+
)
252+
253+
cos = tl.zeros((1, ROPE_DIM), dtype=cos_half.dtype)
254+
cos = al.insert_slice(
255+
cos, cos_half, offsets=(0, 0), sizes=(1, HALF_ROPE_DIM), strides=(1, 2)
256+
)
257+
cos = al.insert_slice(
258+
cos, cos_half, offsets=(0, 1), sizes=(1, HALF_ROPE_DIM), strides=(1, 2)
259+
)
260+
185261
if DO_PARTIAL:
186262
rot_x = al.extract_slice(
187263
normalized_values,
@@ -197,33 +273,63 @@ def split_qkv_rmsnorm_rope_kernel(
197273
)
198274
else:
199275
rot_x = normalized_values
200-
x1 = al.extract_slice(
201-
rot_x,
202-
offsets=(0, 0),
203-
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
204-
strides=(1, 1),
205-
)
206-
x2 = al.extract_slice(
207-
rot_x,
208-
offsets=(0, HALF_ROPE_DIM),
209-
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
210-
strides=(1, 1),
211-
)
276+
if DO_HALF:
277+
x1 = al.extract_slice(
278+
rot_x,
279+
offsets=(0, 0),
280+
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
281+
strides=(1, 1),
282+
)
283+
x2 = al.extract_slice(
284+
rot_x,
285+
offsets=(0, HALF_ROPE_DIM),
286+
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
287+
strides=(1, 1),
288+
)
289+
else:
290+
x1 = al.extract_slice(
291+
rot_x,
292+
offsets=(0, 0),
293+
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
294+
strides=(1, 2),
295+
)
296+
x2 = al.extract_slice(
297+
rot_x,
298+
offsets=(0, 1),
299+
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
300+
strides=(1, 2),
301+
)
212302
cat_x = tl.zeros((KV_BLOCK_SIZE // HEAD_DIM, ROPE_DIM), dtype=tl.bfloat16)
213-
cat_x = al.insert_slice(
214-
cat_x,
215-
-x2,
216-
offsets=(0, 0),
217-
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
218-
strides=(1, 1),
219-
)
220-
cat_x = al.insert_slice(
221-
cat_x,
222-
x1,
223-
offsets=(0, HALF_ROPE_DIM),
224-
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
225-
strides=(1, 1),
226-
)
303+
if DO_HALF:
304+
cat_x = al.insert_slice(
305+
cat_x,
306+
-x2,
307+
offsets=(0, 0),
308+
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
309+
strides=(1, 1),
310+
)
311+
cat_x = al.insert_slice(
312+
cat_x,
313+
x1,
314+
offsets=(0, HALF_ROPE_DIM),
315+
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
316+
strides=(1, 1),
317+
)
318+
else:
319+
cat_x = al.insert_slice(
320+
cat_x,
321+
-x2,
322+
offsets=(0, 0),
323+
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
324+
strides=(1, 2),
325+
)
326+
cat_x = al.insert_slice(
327+
cat_x,
328+
x1,
329+
offsets=(0, 1),
330+
sizes=(KV_BLOCK_SIZE // HEAD_DIM, HALF_ROPE_DIM),
331+
strides=(1, 2),
332+
)
227333
roped_k = cat_x * sin + rot_x * cos
228334
if DO_PARTIAL:
229335
normalized_values = al.insert_slice(
@@ -281,6 +387,7 @@ def split_qkv_rmsnorm_rope(
281387
k_weight=None,
282388
q_bias=None,
283389
k_bias=None,
390+
is_neox_style=True,
284391
):
285392
_, num_vectorcore = get_device_properties()
286393

@@ -329,6 +436,7 @@ def split_qkv_rmsnorm_rope(
329436
rope_dim // 2,
330437
head_dim - rope_dim,
331438
DO_PARTIAL=(head_dim != rope_dim),
439+
DO_HALF=is_neox_style,
332440
)
333441

334442
return q_output, k_output, v_output

0 commit comments

Comments
 (0)