@@ -82,3 +82,299 @@ kernel void kernel_mul_mat_f16_f32_l4(
8282 }
8383 }
8484}
85+
86+ // Each subgroup produces DR_NDST outputs, assumes ne11 == 1
87+ #define MUL_MAT_F16_F32_L4_DR_NDST 4
88+
89+ #ifdef ADRENO_GPU
90+ REQD_SUBGROUP_SIZE_64
91+ #endif
92+ kernel void kernel_mul_mat_f16_f32_l4_dr (
93+ global char * src0 ,
94+ ulong offset0 ,
95+ global char * src1 ,
96+ ulong offset1 ,
97+ global float * dst ,
98+ ulong offsetd ,
99+ int ne00 ,
100+ int ne01 ,
101+ int ne02 ,
102+ ulong nb00 ,
103+ ulong nb01 ,
104+ ulong nb02 ,
105+ ulong nb03 ,
106+ int ne10 ,
107+ int ne11 ,
108+ int ne12 ,
109+ ulong nb10 ,
110+ ulong nb11 ,
111+ ulong nb12 ,
112+ ulong nb13 ,
113+ int ne0 ,
114+ int ne1 ,
115+ int r2 ,
116+ int r3
117+ ) {
118+ src0 = (global char * )((global char * )src0 + offset0 );
119+ src1 = (global char * )((global char * )src1 + offset1 );
120+ dst = (global float * )((global char * )dst + offsetd );
121+
122+ const int r0_base = get_group_id (0 ) * MUL_MAT_F16_F32_L4_DR_NDST ;
123+ const int im = get_group_id (2 );
124+
125+ const int i12 = im % ne12 ;
126+ const int i13 = im / ne12 ;
127+
128+ // assume ne11 == 1
129+ const ulong offset_src1 = i12 * nb12 + i13 * nb13 ;
130+ global float4 * y4 = (global float4 * )(src1 + offset_src1 );
131+
132+ global half4 * x4 [MUL_MAT_F16_F32_L4_DR_NDST ];
133+ float sumf [MUL_MAT_F16_F32_L4_DR_NDST ];
134+
135+ const ulong k_head_off = (i12 /r2 )* nb02 + (i13 /r3 )* nb03 ;
136+
137+ #pragma unroll
138+ for (int n = 0 ; n < MUL_MAT_F16_F32_L4_DR_NDST ; ++ n ) {
139+ int r0 = r0_base + n ;
140+ int r0c = r0 < ne01 ? r0 : 0 ;
141+ ulong off = (ulong )r0c * nb01 + k_head_off ;
142+ x4 [n ] = (global half4 * )(src0 + off );
143+ sumf [n ] = 0.0f ;
144+ }
145+
146+ const int n_chunks = ne00 / 4 ;
147+ const int sg_size = get_max_sub_group_size ();
148+ const int lid = get_sub_group_local_id ();
149+
150+ for (int i = lid ; i < n_chunks ; i += sg_size ) {
151+ float4 q = y4 [i ];
152+ #pragma unroll
153+ for (int n = 0 ; n < MUL_MAT_F16_F32_L4_DR_NDST ; ++ n ) {
154+ float4 k = convert_float4 (x4 [n ][i ]);
155+ sumf [n ] = mad (k .s0 , q .s0 , sumf [n ]);
156+ sumf [n ] = mad (k .s1 , q .s1 , sumf [n ]);
157+ sumf [n ] = mad (k .s2 , q .s2 , sumf [n ]);
158+ sumf [n ] = mad (k .s3 , q .s3 , sumf [n ]);
159+ }
160+ }
161+
162+ #pragma unroll
163+ for (int n = 0 ; n < MUL_MAT_F16_F32_L4_DR_NDST ; ++ n ) {
164+ float reduced = sub_group_reduce_add (sumf [n ]);
165+ int r0 = r0_base + n ;
166+ if (lid == 0 && r0 < ne01 ) {
167+ dst [im * ne1 * ne0 + r0 ] = reduced ;
168+ }
169+ }
170+ }
171+
172+ // Kernels for decoding, Adreno only for now
173+ #define MUL_MAT_F16_F32_L4_DR_LS_R2_MAX 8
174+
175+ #ifdef ADRENO_GPU
176+ #pragma OPENCL EXTENSION cl_qcom_subgroup_shuffle : enable
177+ #define sub_group_shuffle_xor (val , mask ) qcom_sub_group_shuffle_xor((val), (mask), CLK_SUB_GROUP_SHUFFLE_WIDTH_WAVE_SIZE_QCOM, 0.0f)
178+
179+ REQD_SUBGROUP_SIZE_64
180+ kernel void kernel_mul_mat_f16_f32_l4_dr_ls (
181+ global char * src0 ,
182+ ulong offset0 ,
183+ global char * src1 ,
184+ ulong offset1 ,
185+ global float * dst ,
186+ ulong offsetd ,
187+ int ne00 ,
188+ int ne01 ,
189+ int ne02 ,
190+ ulong nb00 ,
191+ ulong nb01 ,
192+ ulong nb02 ,
193+ ulong nb03 ,
194+ int ne10 ,
195+ int ne11 ,
196+ int ne12 ,
197+ ulong nb10 ,
198+ ulong nb11 ,
199+ ulong nb12 ,
200+ ulong nb13 ,
201+ int ne0 ,
202+ int ne1 ,
203+ int r2 ,
204+ int r3
205+ ) {
206+ src0 = (global char * )((global char * )src0 + offset0 );
207+ src1 = (global char * )((global char * )src1 + offset1 );
208+ dst = (global float * )((global char * )dst + offsetd );
209+
210+ const int r0_base = get_group_id (0 ) * 2 ;
211+ const int kv_grp = get_group_id (2 ); // KV head group; im = kv_grp*r2 + q
212+
213+ const int i12_kv = kv_grp % ne02 ;
214+ const int i13_kv = kv_grp / ne02 ;
215+
216+ const int lid = get_sub_group_local_id ();
217+ const int subhalf = lid >> 5 ; // 0 or 1 (which K row in the WG)
218+ const int intra = lid & 31 ; // 0..31 (lane within the half)
219+
220+ const int r0 = r0_base + subhalf ;
221+ const int r0c = r0 < ne01 ? r0 : 0 ; // clamp OOB to row 0; skip write below
222+
223+ // K row pointer for this lane (one K row per half-wave).
224+ const ulong k_off = (ulong )r0c * nb01 + (ulong )i12_kv * nb02 + (ulong )i13_kv * nb03 ;
225+ global half4 * x4 = (global half4 * )(src0 + k_off );
226+
227+ global float4 * y4 [MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ];
228+ #pragma unroll
229+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
230+ const int i12_q = i12_kv * r2 + q ;
231+ const ulong q_off = (ulong )i12_q * nb12 + (ulong )i13_kv * nb13 ;
232+ y4 [q ] = (global float4 * )(src1 + q_off );
233+ }
234+
235+ float partial [MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ];
236+ #pragma unroll
237+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
238+ partial [q ] = 0.0f ;
239+ }
240+
241+ const int n_chunks = ne00 / 4 ;
242+
243+ for (int i = intra ; i < n_chunks ; i += 32 ) {
244+ float4 k = convert_float4 (x4 [i ]);
245+
246+ #pragma unroll
247+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
248+ if (q < r2 ) {
249+ float4 v = y4 [q ][i ];
250+ partial [q ] = mad (k .s0 , v .s0 , partial [q ]);
251+ partial [q ] = mad (k .s1 , v .s1 , partial [q ]);
252+ partial [q ] = mad (k .s2 , v .s2 , partial [q ]);
253+ partial [q ] = mad (k .s3 , v .s3 , partial [q ]);
254+ }
255+ }
256+ }
257+
258+ // half-wave reduction
259+ #pragma unroll
260+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
261+ if (q < r2 ) {
262+ partial [q ] += sub_group_shuffle_xor (partial [q ], 1u );
263+ partial [q ] += sub_group_shuffle_xor (partial [q ], 2u );
264+ partial [q ] += sub_group_shuffle_xor (partial [q ], 4u );
265+ partial [q ] += sub_group_shuffle_xor (partial [q ], 8u );
266+ partial [q ] += sub_group_shuffle_xor (partial [q ], 16u );
267+ }
268+ }
269+
270+ if (intra == 0 && r0 < ne01 ) {
271+ #pragma unroll
272+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
273+ if (q < r2 ) {
274+ const int im = i12_kv * r2 + q + i13_kv * ne12 ;
275+ dst [im * ne1 * ne0 + r0 ] = partial [q ];
276+ }
277+ }
278+ }
279+ }
280+
281+ REQD_SUBGROUP_SIZE_64
282+ kernel void kernel_mul_mat_f16_f32_l4_dr_lq (
283+ global char * src0 ,
284+ ulong offset0 ,
285+ global char * src1 ,
286+ ulong offset1 ,
287+ global float * dst ,
288+ ulong offsetd ,
289+ int ne00 ,
290+ int ne01 ,
291+ int ne02 ,
292+ ulong nb00 ,
293+ ulong nb01 ,
294+ ulong nb02 ,
295+ ulong nb03 ,
296+ int ne10 ,
297+ int ne11 ,
298+ int ne12 ,
299+ ulong nb10 ,
300+ ulong nb11 ,
301+ ulong nb12 ,
302+ ulong nb13 ,
303+ int ne0 ,
304+ int ne1 ,
305+ int r2 ,
306+ int r3
307+ ) {
308+ src0 = (global char * )((global char * )src0 + offset0 );
309+ src1 = (global char * )((global char * )src1 + offset1 );
310+ dst = (global float * )((global char * )dst + offsetd );
311+
312+ const int r0_base = get_group_id (0 ) * 4 ;
313+ const int kv_grp = get_group_id (2 );
314+
315+ const int i12_kv = kv_grp % ne02 ;
316+ const int i13_kv = kv_grp / ne02 ;
317+
318+ const int lid = get_sub_group_local_id ();
319+ const int subq = lid >> 4 ; // 0..3 (which K row)
320+ const int intra = lid & 15 ; // 0..15 (lane within quarter)
321+
322+ const int r0 = r0_base + subq ;
323+ const int r0c = r0 < ne01 ? r0 : 0 ;
324+
325+ const ulong k_off = (ulong )r0c * nb01 + (ulong )i12_kv * nb02 + (ulong )i13_kv * nb03 ;
326+ global half4 * x4 = (global half4 * )(src0 + k_off );
327+
328+ global float4 * y4 [MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ];
329+ #pragma unroll
330+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
331+ const int i12_q = i12_kv * r2 + q ;
332+ const ulong q_off = (ulong )i12_q * nb12 + (ulong )i13_kv * nb13 ;
333+ y4 [q ] = (global float4 * )(src1 + q_off );
334+ }
335+
336+ float partial [MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ];
337+ #pragma unroll
338+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
339+ partial [q ] = 0.0f ;
340+ }
341+
342+ const int n_chunks = ne00 / 4 ;
343+
344+ for (int i = intra ; i < n_chunks ; i += 16 ) {
345+ float4 k = convert_float4 (x4 [i ]);
346+
347+ #pragma unroll
348+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
349+ if (q < r2 ) {
350+ float4 v = y4 [q ][i ];
351+ partial [q ] = mad (k .s0 , v .s0 , partial [q ]);
352+ partial [q ] = mad (k .s1 , v .s1 , partial [q ]);
353+ partial [q ] = mad (k .s2 , v .s2 , partial [q ]);
354+ partial [q ] = mad (k .s3 , v .s3 , partial [q ]);
355+ }
356+ }
357+ }
358+
359+ // quarter-wave reduction
360+ #pragma unroll
361+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
362+ if (q < r2 ) {
363+ partial [q ] += sub_group_shuffle_xor (partial [q ], 1u );
364+ partial [q ] += sub_group_shuffle_xor (partial [q ], 2u );
365+ partial [q ] += sub_group_shuffle_xor (partial [q ], 4u );
366+ partial [q ] += sub_group_shuffle_xor (partial [q ], 8u );
367+ }
368+ }
369+
370+ if (intra == 0 && r0 < ne01 ) {
371+ #pragma unroll
372+ for (int q = 0 ; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX ; ++ q ) {
373+ if (q < r2 ) {
374+ const int im = i12_kv * r2 + q + i13_kv * ne12 ;
375+ dst [im * ne1 * ne0 + r0 ] = partial [q ];
376+ }
377+ }
378+ }
379+ }
380+ #endif // ADRENO_GPU
0 commit comments