@@ -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