1#![allow(non_snake_case)]
35
36use core::ffi::c_void;
37
38use teeny_core::dtype::Float;
39use teeny_macros::kernel;
40use teeny_triton::triton::{
41 types::{AddOffsets, Comparison, Tensor},
42 *,
43};
44
45use super::flash_attn2::FlashAttention2Forward;
46
47#[kernel]
58pub fn psa_pack_qkv<T: Triton, D: Float, const KEY_DIM: i32>(
59 qkv_ptr: T::Pointer<D>,
60 out_ptr: T::Pointer<D>,
61 qkv_h: i32, H: i32,
63 W: i32,
64 B: i32,
65 num_heads: i32,
66) where
67 T::I32Tensor: Tensor<i32, 1>,
68 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
69 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
70{
71 let pid = T::program_id(Axis::X); let BH: i32 = B * num_heads;
73 let N: i32 = H * W;
74
75 let section: i32 = pid / (BH * N);
76 let bh: i32 = (pid / N) % BH;
77 let n: i32 = pid % N;
78 let b: i32 = bh / num_heads;
79 let h: i32 = bh % num_heads;
80
81 let d = T::arange(0, KEY_DIM);
82
83 let chan_base: i32 = h * 4 * KEY_DIM + section * KEY_DIM;
85 let src_off = (d + chan_base) * (H * W) + (b * qkv_h * H * W + n);
86
87 let x = T::load(qkv_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
88
89 let dst_base: i32 = section * BH * N * KEY_DIM + bh * N * KEY_DIM + n * KEY_DIM;
91 let dst_off = d + dst_base;
92
93 T::store(out_ptr.add_offsets(dst_off), x, None, &[], None, None);
94}
95
96#[kernel]
105pub fn psa_extract_v_nchw<T: Triton, D: Float, const KEY_DIM: i32>(
106 qkv_ptr: T::Pointer<D>,
107 v_ptr: T::Pointer<D>,
108 qkv_h: i32,
109 c: i32, H: i32,
111 W: i32,
112 num_heads: i32,
113) where
114 T::I32Tensor: Tensor<i32, 1>,
115 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
116 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
117{
118 let pid = T::program_id(Axis::X); let N: i32 = H * W;
120 let bh: i32 = pid / N;
121 let n: i32 = pid % N;
122 let b: i32 = bh / num_heads;
123 let h: i32 = bh % num_heads;
124
125 let d = T::arange(0, KEY_DIM);
126
127 let src_lo_base: i32 = h * 4 * KEY_DIM + 2 * KEY_DIM;
128 let src_hi_base: i32 = h * 4 * KEY_DIM + 3 * KEY_DIM;
129 let src_off_lo = (d + src_lo_base) * (H * W) + (b * qkv_h * H * W + n);
130 let src_off_hi = (d + src_hi_base) * (H * W) + (b * qkv_h * H * W + n);
131
132 let x_lo = T::load(qkv_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
133 let x_hi = T::load(qkv_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
134
135 let dst_lo_base: i32 = h * 2 * KEY_DIM;
136 let dst_hi_base: i32 = h * 2 * KEY_DIM + KEY_DIM;
137 let dst_off_lo = (d + dst_lo_base) * (H * W) + (b * c * H * W + n);
138 let dst_off_hi = (d + dst_hi_base) * (H * W) + (b * c * H * W + n);
139
140 T::store(v_ptr.add_offsets(dst_off_lo), x_lo, None, &[], None, None);
141 T::store(v_ptr.add_offsets(dst_off_hi), x_hi, None, &[], None, None);
142}
143
144#[kernel]
153pub fn psa_merge_attn_nchw<T: Triton, D: Float, const KEY_DIM: i32>(
154 lo_ptr: T::Pointer<D>,
155 hi_ptr: T::Pointer<D>,
156 out_ptr: T::Pointer<D>,
157 c: i32,
158 H: i32,
159 W: i32,
160 num_heads: i32,
161) where
162 T::I32Tensor: Tensor<i32, 1>,
163 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
164 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
165{
166 let pid = T::program_id(Axis::X); let N: i32 = H * W;
168 let bh: i32 = pid / N;
169 let n: i32 = pid % N;
170 let b: i32 = bh / num_heads;
171 let h: i32 = bh % num_heads;
172
173 let d = T::arange(0, KEY_DIM);
174
175 let src_base: i32 = bh * N * KEY_DIM + n * KEY_DIM;
177 let src_off = d + src_base;
178
179 let x_lo = T::load(lo_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
180 let x_hi = T::load(hi_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
181
182 let dst_lo_base: i32 = h * 2 * KEY_DIM;
183 let dst_hi_base: i32 = h * 2 * KEY_DIM + KEY_DIM;
184 let dst_off_lo = (d + dst_lo_base) * (H * W) + (b * c * H * W + n);
185 let dst_off_hi = (d + dst_hi_base) * (H * W) + (b * c * H * W + n);
186
187 T::store(out_ptr.add_offsets(dst_off_lo), x_lo, None, &[], None, None);
188 T::store(out_ptr.add_offsets(dst_off_hi), x_hi, None, &[], None, None);
189}
190
191#[kernel]
201pub fn psa_pack_qkv_backward<T: Triton, D: Float, const KEY_DIM: i32>(
202 d_packed_ptr: T::Pointer<D>, d_qkv_ptr: T::Pointer<D>, qkv_h: i32,
205 H: i32,
206 W: i32,
207 B: i32,
208 num_heads: i32,
209) where
210 T::I32Tensor: Tensor<i32, 1>,
211 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
212 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
213{
214 let pid = T::program_id(Axis::X); let BH = B * num_heads;
216 let N = H * W;
217 let section = pid / (BH * N);
218 let bh = (pid / N) % BH;
219 let n = pid % N;
220 let b = bh / num_heads;
221 let h = bh % num_heads;
222
223 let d = T::arange(0, KEY_DIM);
224
225 let src_base = section * BH * N * KEY_DIM + bh * N * KEY_DIM + n * KEY_DIM;
227 let dx = T::load(d_packed_ptr.add_offsets(d + src_base), None, None, &[], None, None, None, false);
228
229 let chan_base = h * 4 * KEY_DIM + section * KEY_DIM;
231 let dst_off = (d + chan_base) * (H * W) + (b * qkv_h * H * W + n);
232 T::atomic_add(d_qkv_ptr.add_offsets(dst_off), dx, None, None, None);
233}
234
235#[kernel]
244pub fn psa_extract_v_backward<T: Triton, D: Float, const KEY_DIM: i32>(
245 d_v_ptr: T::Pointer<D>, d_qkv_ptr: T::Pointer<D>, qkv_h: i32,
248 c: i32,
249 H: i32,
250 W: i32,
251 num_heads: i32,
252) where
253 T::I32Tensor: Tensor<i32, 1>,
254 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
255 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
256{
257 let pid = T::program_id(Axis::X); let N = H * W;
259 let bh = pid / N;
260 let n = pid % N;
261 let b = bh / num_heads;
262 let h = bh % num_heads;
263
264 let d = T::arange(0, KEY_DIM);
265
266 let v_lo_src_base = h * 2 * KEY_DIM;
268 let v_hi_src_base = h * 2 * KEY_DIM + KEY_DIM;
269 let src_off_lo = (d + v_lo_src_base) * (H * W) + (b * c * H * W + n);
270 let src_off_hi = (d + v_hi_src_base) * (H * W) + (b * c * H * W + n);
271 let dx_lo = T::load(d_v_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
272 let dx_hi = T::load(d_v_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
273
274 let qkv_lo_base = h * 4 * KEY_DIM + 2 * KEY_DIM;
276 let qkv_hi_base = h * 4 * KEY_DIM + 3 * KEY_DIM;
277 let dst_off_lo = (d + qkv_lo_base) * (H * W) + (b * qkv_h * H * W + n);
278 let dst_off_hi = (d + qkv_hi_base) * (H * W) + (b * qkv_h * H * W + n);
279 T::atomic_add(d_qkv_ptr.add_offsets(dst_off_lo), dx_lo, None, None, None);
280 T::atomic_add(d_qkv_ptr.add_offsets(dst_off_hi), dx_hi, None, None, None);
281}
282
283#[kernel]
292pub fn psa_merge_attn_backward<T: Triton, D: Float, const KEY_DIM: i32>(
293 d_merged_ptr: T::Pointer<D>, d_lo_ptr: T::Pointer<D>, d_hi_ptr: T::Pointer<D>, c: i32,
297 H: i32,
298 W: i32,
299 num_heads: i32,
300) where
301 T::I32Tensor: Tensor<i32, 1>,
302 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
303 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
304{
305 let pid = T::program_id(Axis::X); let N = H * W;
307 let bh = pid / N;
308 let n = pid % N;
309 let b = bh / num_heads;
310 let h = bh % num_heads;
311
312 let d = T::arange(0, KEY_DIM);
313
314 let src_lo_base = h * 2 * KEY_DIM;
316 let src_hi_base = h * 2 * KEY_DIM + KEY_DIM;
317 let src_off_lo = (d + src_lo_base) * (H * W) + (b * c * H * W + n);
318 let src_off_hi = (d + src_hi_base) * (H * W) + (b * c * H * W + n);
319 let dx_lo = T::load(d_merged_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
320 let dx_hi = T::load(d_merged_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
321
322 let dst_base = bh * N * KEY_DIM + n * KEY_DIM;
324 let dst_off = d + dst_base;
325 T::store(d_lo_ptr.add_offsets(dst_off), dx_lo, None, &[], None, None);
326 T::store(d_hi_ptr.add_offsets(dst_off), dx_hi, None, &[], None, None);
327}
328
329pub struct PsaPackQkvRuntimeOp<D: Float + Send + Sync + 'static> {
333 fwd: PsaPackQkv<D>,
334 bwd: PsaPackQkvBackward<D>,
335 num_heads: usize,
336}
337
338impl<D: Float + Send + Sync + 'static> PsaPackQkvRuntimeOp<D> {
339 pub fn new(key_dim: i32, num_heads: usize) -> Self {
341 Self { fwd: PsaPackQkv::<D>::new(key_dim), bwd: PsaPackQkvBackward::<D>::new(key_dim), num_heads }
342 }
343
344 pub fn kernel_name(&self) -> &str { self.fwd.name }
346 pub fn forward_source(&self) -> &str { &self.fwd.source }
348 pub fn backward_source(&self) -> &str { &self.bwd.source }
350}
351
352impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaPackQkvRuntimeOp<D> {
353 fn n_activation_inputs(&self) -> usize { 1 }
354
355 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
356
357 fn pack_args(
358 &self,
359 inputs: &[(teeny_core::model::RawPtr, &[usize])],
360 _params: &[teeny_core::model::RawPtr],
361 output: teeny_core::model::RawPtr,
362 _output_shape: &[usize],
363 _output_row_stride: i32,
364 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
365 ) {
366 let b = inputs[0].1[0] as i32;
368 let qkv_h = inputs[0].1[1] as i32;
369 let h = inputs[0].1[2] as i32;
370 let w = inputs[0].1[3] as i32;
371 visitor.visit_ptr(inputs[0].0);
372 visitor.visit_ptr(output);
373 visitor.visit_i32(qkv_h);
374 visitor.visit_i32(h);
375 visitor.visit_i32(w);
376 visitor.visit_i32(b);
377 visitor.visit_i32(self.num_heads as i32);
378 }
379
380 fn block(&self) -> [u32; 3] { [128, 1, 1] }
381
382 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
383 [(output_shape[0] * output_shape[1] * output_shape[2] * output_shape[3]) as u32, 1, 1]
386 }
387
388 #[cfg(feature = "training")]
389 fn has_backward(&self) -> bool { true }
390
391 #[cfg(feature = "training")]
393 fn pack_backward_args(
394 &self,
395 inputs: &[(teeny_core::model::RawPtr, &[usize])],
396 _params: &[teeny_core::model::RawPtr],
397 _output: teeny_core::model::RawPtr,
398 output_shape: &[usize],
399 grad_output: teeny_core::model::RawPtr,
400 _grad_output_row_stride: i32,
401 grad_inputs: &[teeny_core::model::RawPtr],
402 _grad_params: &[teeny_core::model::RawPtr],
403 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
404 ) {
405 let b = inputs[0].1[0] as i32;
407 let qkv_h = inputs[0].1[1] as i32;
408 let h = inputs[0].1[2] as i32;
409 let w = inputs[0].1[3] as i32;
410 let _ = output_shape;
411 visitor.visit_ptr(grad_output); visitor.visit_ptr(grad_inputs[0]); visitor.visit_i32(qkv_h);
414 visitor.visit_i32(h);
415 visitor.visit_i32(w);
416 visitor.visit_i32(b);
417 visitor.visit_i32(self.num_heads as i32);
418 }
419
420 #[cfg(feature = "training")]
421 fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
422
423 #[cfg(feature = "training")]
425 fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
426 [(output_shape[0] * output_shape[1] * output_shape[2] * output_shape[3]) as u32, 1, 1]
427 }
428}
429
430pub struct PsaExtractVRuntimeOp<D: Float + Send + Sync + 'static> {
434 fwd: PsaExtractVNchw<D>,
435 bwd: PsaExtractVBackward<D>,
436 num_heads: usize,
437}
438
439impl<D: Float + Send + Sync + 'static> PsaExtractVRuntimeOp<D> {
440 pub fn new(key_dim: i32, num_heads: usize) -> Self {
442 Self { fwd: PsaExtractVNchw::<D>::new(key_dim), bwd: PsaExtractVBackward::<D>::new(key_dim), num_heads }
443 }
444
445 pub fn kernel_name(&self) -> &str { self.fwd.name }
447 pub fn forward_source(&self) -> &str { &self.fwd.source }
449 pub fn backward_source(&self) -> &str { &self.bwd.source }
451}
452
453impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaExtractVRuntimeOp<D> {
454 fn n_activation_inputs(&self) -> usize { 1 }
455
456 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
457
458 fn pack_args(
459 &self,
460 inputs: &[(teeny_core::model::RawPtr, &[usize])],
461 _params: &[teeny_core::model::RawPtr],
462 output: teeny_core::model::RawPtr,
463 output_shape: &[usize],
464 _output_row_stride: i32,
465 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
466 ) {
467 let qkv_h = inputs[0].1[1] as i32;
469 let h = inputs[0].1[2] as i32;
470 let w = inputs[0].1[3] as i32;
471 let c = output_shape[1] as i32;
472 visitor.visit_ptr(inputs[0].0);
473 visitor.visit_ptr(output);
474 visitor.visit_i32(qkv_h);
475 visitor.visit_i32(c);
476 visitor.visit_i32(h);
477 visitor.visit_i32(w);
478 visitor.visit_i32(self.num_heads as i32);
479 }
480
481 fn block(&self) -> [u32; 3] { [128, 1, 1] }
482
483 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
484 let bh = output_shape[0] * self.num_heads;
486 let n = output_shape[2] * output_shape[3];
487 [(bh * n) as u32, 1, 1]
488 }
489
490 #[cfg(feature = "training")]
491 fn has_backward(&self) -> bool { true }
492
493 #[cfg(feature = "training")]
495 fn pack_backward_args(
496 &self,
497 inputs: &[(teeny_core::model::RawPtr, &[usize])],
498 _params: &[teeny_core::model::RawPtr],
499 _output: teeny_core::model::RawPtr,
500 output_shape: &[usize],
501 grad_output: teeny_core::model::RawPtr,
502 _grad_output_row_stride: i32,
503 grad_inputs: &[teeny_core::model::RawPtr],
504 _grad_params: &[teeny_core::model::RawPtr],
505 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
506 ) {
507 let qkv_h = inputs[0].1[1] as i32;
509 let h = inputs[0].1[2] as i32;
510 let w = inputs[0].1[3] as i32;
511 let c = output_shape[1] as i32;
512 visitor.visit_ptr(grad_output); visitor.visit_ptr(grad_inputs[0]); visitor.visit_i32(qkv_h);
515 visitor.visit_i32(c);
516 visitor.visit_i32(h);
517 visitor.visit_i32(w);
518 visitor.visit_i32(self.num_heads as i32);
519 }
520
521 #[cfg(feature = "training")]
522 fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
523
524 #[cfg(feature = "training")]
526 fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
527 let bh = output_shape[0] * self.num_heads;
528 let n = output_shape[2] * output_shape[3];
529 [(bh * n) as u32, 1, 1]
530 }
531}
532
533pub struct PsaMergeAttnRuntimeOp<D: Float + Send + Sync + 'static> {
537 fwd: PsaMergeAttnNchw<D>,
538 bwd: PsaMergeAttnBackward<D>,
539 num_heads: usize,
540}
541
542impl<D: Float + Send + Sync + 'static> PsaMergeAttnRuntimeOp<D> {
543 pub fn new(key_dim: i32, num_heads: usize) -> Self {
545 Self { fwd: PsaMergeAttnNchw::<D>::new(key_dim), bwd: PsaMergeAttnBackward::<D>::new(key_dim), num_heads }
546 }
547
548 pub fn kernel_name(&self) -> &str { self.fwd.name }
550 pub fn forward_source(&self) -> &str { &self.fwd.source }
552 pub fn backward_source(&self) -> &str { &self.bwd.source }
554}
555
556impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaMergeAttnRuntimeOp<D> {
557 fn n_activation_inputs(&self) -> usize { 2 }
558
559 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
560
561 fn pack_args(
562 &self,
563 inputs: &[(teeny_core::model::RawPtr, &[usize])],
564 _params: &[teeny_core::model::RawPtr],
565 output: teeny_core::model::RawPtr,
566 output_shape: &[usize],
567 _output_row_stride: i32,
568 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
569 ) {
570 let c = output_shape[1] as i32;
573 let h = output_shape[2] as i32;
574 let w = output_shape[3] as i32;
575 visitor.visit_ptr(inputs[0].0);
576 visitor.visit_ptr(inputs[1].0);
577 visitor.visit_ptr(output);
578 visitor.visit_i32(c);
579 visitor.visit_i32(h);
580 visitor.visit_i32(w);
581 visitor.visit_i32(self.num_heads as i32);
582 }
583
584 fn block(&self) -> [u32; 3] { [128, 1, 1] }
585
586 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
587 let bh = output_shape[0] * self.num_heads;
589 let n = output_shape[2] * output_shape[3];
590 [(bh * n) as u32, 1, 1]
591 }
592
593 #[cfg(feature = "training")]
594 fn has_backward(&self) -> bool { true }
595
596 #[cfg(feature = "training")]
598 fn pack_backward_args(
599 &self,
600 _inputs: &[(teeny_core::model::RawPtr, &[usize])],
601 _params: &[teeny_core::model::RawPtr],
602 _output: teeny_core::model::RawPtr,
603 output_shape: &[usize],
604 grad_output: teeny_core::model::RawPtr,
605 _grad_output_row_stride: i32,
606 grad_inputs: &[teeny_core::model::RawPtr],
607 _grad_params: &[teeny_core::model::RawPtr],
608 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
609 ) {
610 let c = output_shape[1] as i32;
612 let h = output_shape[2] as i32;
613 let w = output_shape[3] as i32;
614 visitor.visit_ptr(grad_output); visitor.visit_ptr(grad_inputs[0]); visitor.visit_ptr(grad_inputs[1]); visitor.visit_i32(c);
618 visitor.visit_i32(h);
619 visitor.visit_i32(w);
620 visitor.visit_i32(self.num_heads as i32);
621 }
622
623 #[cfg(feature = "training")]
624 fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
625
626 #[cfg(feature = "training")]
628 fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
629 let bh = output_shape[0] * self.num_heads;
630 let n = output_shape[2] * output_shape[3];
631 [(bh * n) as u32, 1, 1]
632 }
633}
634
635#[kernel]
650pub fn psa_fa2_backward<T: Triton, D: Float, const HEAD_DIM: i32>(
651 q_ptr: T::Pointer<D>, k_ptr: T::Pointer<D>, v_ptr: T::Pointer<D>, o_ptr: T::Pointer<D>, do_ptr: T::Pointer<D>, l_ptr: T::Pointer<D>, dq_ptr: T::Pointer<D>, dk_ptr: T::Pointer<D>, dv_ptr: T::Pointer<D>, N: i32, softmax_scale: f32,
662) where
663 T::I32Tensor: Tensor<i32, 1>,
664 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
665 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
666{
667 let pid_n = T::program_id(Axis::X); let pid_bh = T::program_id(Axis::Y); let row_base = pid_bh * N * HEAD_DIM + pid_n * HEAD_DIM;
671 let l_base = pid_bh * N + pid_n;
672 let bh_base = pid_bh * N * HEAD_DIM;
673 let l_bh = pid_bh * N;
674
675 let d = T::arange(0, HEAD_DIM);
676 let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
677
678 let q_vec = T::load(q_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
680 let o_vec = T::load(o_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
681 let do_vec = T::load(do_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
682
683 let d_n = T::sum(o_vec * do_vec, Some(0), false); let l_n_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_base), None, None, &[], None, None, None, false);
686 let l_n = T::sum(l_n_raw, Some(0), false); let mut dq_acc = T::zeros::<D>(&[HEAD_DIM]);
690 for k_row in 0..N {
691 let kv_row_base = bh_base + k_row * HEAD_DIM;
692 let k_vec = T::load(k_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
693 let v_vec = T::load(v_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
694 let qk = T::sum(q_vec * k_vec, Some(0), false) * scale_t;
695 let p = T::exp(qk - l_n);
696 let do_dot_v = T::sum(do_vec * v_vec, Some(0), false);
697 let ds = p * (do_dot_v - d_n);
698 dq_acc = dq_acc + ds * k_vec * scale_t;
699 }
700 T::atomic_add(dq_ptr.add_offsets(d + row_base), dq_acc, None, None, None);
701
702 let k_vec_n = T::load(k_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
704 let v_vec_n = T::load(v_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
705 let mut dk_acc = T::zeros::<D>(&[HEAD_DIM]);
706 let mut dv_acc = T::zeros::<D>(&[HEAD_DIM]);
707 for q_row in 0..N {
708 let q_row_base_m = bh_base + q_row * HEAD_DIM;
709 let l_row_base_m = l_bh + q_row;
710 let q_vec_m = T::load(q_ptr.add_offsets(d + q_row_base_m), None, None, &[], None, None, None, false);
711 let o_vec_m = T::load(o_ptr.add_offsets(d + q_row_base_m), None, None, &[], None, None, None, false);
712 let do_vec_m = T::load(do_ptr.add_offsets(d + q_row_base_m), None, None, &[], None, None, None, false);
713 let l_m_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_row_base_m), None, None, &[], None, None, None, false);
714 let l_m = T::sum(l_m_raw, Some(0), false);
715 let d_m = T::sum(o_vec_m * do_vec_m, Some(0), false);
716 let qk = T::sum(q_vec_m * k_vec_n, Some(0), false) * scale_t;
717 let p = T::exp(qk - l_m);
718 dv_acc = dv_acc + p * do_vec_m;
719 let do_dot_v_m = T::sum(do_vec_m * v_vec_n, Some(0), false);
720 let ds_m = p * (do_dot_v_m - d_m);
721 dk_acc = dk_acc + ds_m * q_vec_m * scale_t;
722 }
723 T::atomic_add(dk_ptr.add_offsets(d + row_base), dk_acc, None, None, None);
724 T::store(dv_ptr.add_offsets(d + row_base), dv_acc, None, &[], None, None);
725}
726
727use std::any::Any;
735use std::sync::Arc;
736use teeny_core::{
737 graph::{CustomOp, Shape},
738 model::RuntimeOp,
739};
740
741pub struct PsaPackQkvOp<D: Float + Send + Sync + 'static> {
745 inner: Arc<PsaPackQkvRuntimeOp<D>>,
746 num_heads: usize,
747}
748
749impl<D: Float + Send + Sync + 'static> PsaPackQkvOp<D> {
750 pub fn new(key_dim: i32, num_heads: usize) -> Self {
752 Self { inner: Arc::new(PsaPackQkvRuntimeOp::<D>::new(key_dim, num_heads)), num_heads }
753 }
754}
755
756impl<D: Float + Send + Sync + 'static> CustomOp for PsaPackQkvOp<D> {
757 fn name(&self) -> &str { "psa_pack_qkv" }
758
759 fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
767 let s = input_shapes[0]; let nh = self.num_heads;
769 vec![
770 s[0], Some(4), Some(nh), s[2].and_then(|h| s[3].map(|w| h * w)), s[1].map(|qkv_h| qkv_h / (nh * 4)), ]
776 }
777
778 fn as_any(&self) -> &dyn Any { self }
779
780 fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
781 Some((
782 self.inner.kernel_name().to_string(),
783 self.inner.forward_source().to_string(),
784 "entry_point".to_string(),
785 Arc::clone(&self.inner) as Arc<dyn RuntimeOp>,
786 ))
787 }
788
789 fn lower_backward_source(&self) -> String {
790 self.inner.backward_source().to_string()
791 }
792}
793
794pub struct PsaExtractVOp<D: Float + Send + Sync + 'static>(Arc<PsaExtractVRuntimeOp<D>>);
798
799impl<D: Float + Send + Sync + 'static> PsaExtractVOp<D> {
800 pub fn new(key_dim: i32, num_heads: usize) -> Self {
802 Self(Arc::new(PsaExtractVRuntimeOp::<D>::new(key_dim, num_heads)))
803 }
804}
805
806impl<D: Float + Send + Sync + 'static> CustomOp for PsaExtractVOp<D> {
807 fn name(&self) -> &str { "psa_extract_v_nchw" }
808
809 fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
810 let s = input_shapes[0];
811 vec![s[0], s[1].map(|qkv_h| qkv_h / 2), s[2], s[3]]
812 }
813
814 fn as_any(&self) -> &dyn Any { self }
815
816 fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
817 Some((
818 self.0.kernel_name().to_string(),
819 self.0.forward_source().to_string(),
820 "entry_point".to_string(),
821 Arc::clone(&self.0) as Arc<dyn RuntimeOp>,
822 ))
823 }
824
825 fn lower_backward_source(&self) -> String {
826 self.0.backward_source().to_string()
827 }
828}
829
830pub struct PsaMergeAttnOp<D: Float + Send + Sync + 'static> {
837 inner: Arc<PsaMergeAttnRuntimeOp<D>>,
838 num_heads: usize,
839 h: usize,
840 w: usize,
841}
842
843impl<D: Float + Send + Sync + 'static> PsaMergeAttnOp<D> {
844 pub fn new(key_dim: i32, num_heads: usize, h: usize, w: usize) -> Self {
846 Self { inner: Arc::new(PsaMergeAttnRuntimeOp::<D>::new(key_dim, num_heads)), num_heads, h, w }
847 }
848}
849
850impl<D: Float + Send + Sync + 'static> CustomOp for PsaMergeAttnOp<D> {
851 fn name(&self) -> &str { "psa_merge_attn_nchw" }
852
853 fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
854 let lo = input_shapes[0];
856 let nh = self.num_heads;
857 vec![
858 lo[0], lo[3].map(|kd| nh * 2 * kd), Some(self.h),
861 Some(self.w),
862 ]
863 }
864
865 fn as_any(&self) -> &dyn Any { self }
866
867 fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
868 Some((
869 self.inner.kernel_name().to_string(),
870 self.inner.forward_source().to_string(),
871 "entry_point".to_string(),
872 Arc::clone(&self.inner) as Arc<dyn RuntimeOp>,
873 ))
874 }
875
876 fn lower_backward_source(&self) -> String {
877 self.inner.backward_source().to_string()
878 }
879}
880
881pub struct FlashAttn2PsaOp<D: Float + Send + Sync + 'static>(Arc<FlashAttn2PsaRuntimeOp<D>>);
885
886impl<D: Float + Send + Sync + 'static> FlashAttn2PsaOp<D> {
887 pub fn new_lo(key_dim: i32) -> Self {
889 Self(Arc::new(FlashAttn2PsaRuntimeOp::<D>::new_lo(key_dim)))
890 }
891
892 pub fn new_hi(key_dim: i32) -> Self {
894 Self(Arc::new(FlashAttn2PsaRuntimeOp::<D>::new_hi(key_dim)))
895 }
896}
897
898impl<D: Float + Send + Sync + 'static> CustomOp for FlashAttn2PsaOp<D> {
899 fn name(&self) -> &str { "flash_attention2_forward" }
900
901 fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
902 let s = input_shapes[0];
904 vec![s[0], s[2], s[3], s[4]]
905 }
906
907 fn as_any(&self) -> &dyn Any { self }
908
909 fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
910 Some((
911 self.0.kernel_name().to_string(),
912 self.0.forward_source().to_string(),
913 "entry_point".to_string(),
914 Arc::clone(&self.0) as Arc<dyn RuntimeOp>,
915 ))
916 }
917
918 fn lower_backward_source(&self) -> String {
919 self.0.backward_source().to_string()
920 }
921}
922
923pub struct FlashAttn2PsaRuntimeOp<D: Float + Send + Sync + 'static> {
931 fwd: FlashAttention2Forward<D>,
932 bwd: PsaFa2Backward<D>,
933 v_section: usize,
934}
935
936impl<D: Float + Send + Sync + 'static> FlashAttn2PsaRuntimeOp<D> {
937 pub fn new_lo(key_dim: i32) -> Self {
939 Self { fwd: FlashAttention2Forward::<D>::new(key_dim), bwd: PsaFa2Backward::<D>::new(key_dim), v_section: 2 }
940 }
941
942 pub fn new_hi(key_dim: i32) -> Self {
944 Self { fwd: FlashAttention2Forward::<D>::new(key_dim), bwd: PsaFa2Backward::<D>::new(key_dim), v_section: 3 }
945 }
946
947 pub fn kernel_name(&self) -> &str { self.fwd.name }
949 pub fn forward_source(&self) -> &str { &self.fwd.source }
951 pub fn backward_source(&self) -> &str { &self.bwd.source }
953}
954
955impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for FlashAttn2PsaRuntimeOp<D> {
956 fn n_activation_inputs(&self) -> usize { 1 }
957
958 fn param_shapes(&self, input_shapes: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
959 let bh = input_shapes[0][0] * input_shapes[0][2]; let n = input_shapes[0][3];
962 vec![vec![bh * n]] }
964
965 fn pack_args(
966 &self,
967 inputs: &[(teeny_core::model::RawPtr, &[usize])],
968 params: &[teeny_core::model::RawPtr],
969 output: teeny_core::model::RawPtr,
970 _output_shape: &[usize],
971 _output_row_stride: i32,
972 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
973 ) {
974 let b = inputs[0].1[0];
976 let nh = inputs[0].1[2];
977 let n = inputs[0].1[3];
978 let kd = inputs[0].1[4];
979 let bh = b * nh; let section_elems = bh * n * kd;
981
982 let base = inputs[0].0 as *mut D;
983 let q_ptr = base as *mut c_void;
984 let k_ptr = unsafe { base.add(section_elems) } as *mut c_void;
985 let v_ptr = unsafe { base.add(self.v_section * section_elems) } as *mut c_void;
986 let softmax_scale = 1.0_f32 / (kd as f32).sqrt();
987
988 visitor.visit_ptr(q_ptr);
989 visitor.visit_ptr(k_ptr);
990 visitor.visit_ptr(v_ptr);
991 visitor.visit_ptr(output);
992 visitor.visit_ptr(params[0]);
993 visitor.visit_i32(n as i32);
994 visitor.visit_i32(n as i32);
995 visitor.visit_f32(softmax_scale);
996 visitor.visit_f32(f32::NEG_INFINITY);
997 }
998
999 fn block(&self) -> [u32; 3] { [1, 1, 1] }
1000
1001 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
1002 [output_shape[2] as u32, (output_shape[0] * output_shape[1]) as u32, 1]
1004 }
1005
1006 #[cfg(feature = "training")]
1007 fn has_backward(&self) -> bool { true }
1008
1009 #[cfg(feature = "training")]
1011 fn pack_backward_args(
1012 &self,
1013 inputs: &[(teeny_core::model::RawPtr, &[usize])],
1014 params: &[teeny_core::model::RawPtr],
1015 output: teeny_core::model::RawPtr,
1016 _output_shape: &[usize],
1017 grad_output: teeny_core::model::RawPtr,
1018 _grad_output_row_stride: i32,
1019 grad_inputs: &[teeny_core::model::RawPtr],
1020 _grad_params: &[teeny_core::model::RawPtr],
1021 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1022 ) {
1023 let b = inputs[0].1[0];
1025 let nh = inputs[0].1[2];
1026 let n = inputs[0].1[3];
1027 let kd = inputs[0].1[4];
1028 let bh = b * nh;
1029 let section_elems = bh * n * kd;
1030 let softmax_scale = 1.0_f32 / (kd as f32).sqrt();
1031
1032 let fwd_base = inputs[0].0 as *mut D;
1034 let q_ptr = fwd_base as *mut c_void;
1035 let k_ptr = unsafe { fwd_base.add(section_elems) } as *mut c_void;
1036 let v_ptr = unsafe { fwd_base.add(self.v_section * section_elems) } as *mut c_void;
1037
1038 let d_base = grad_inputs[0] as *mut D;
1040 let dq_ptr = d_base as *mut c_void;
1041 let dk_ptr = unsafe { d_base.add(section_elems) } as *mut c_void;
1042 let dv_ptr = unsafe { d_base.add(self.v_section * section_elems) } as *mut c_void;
1043
1044 visitor.visit_ptr(q_ptr);
1045 visitor.visit_ptr(k_ptr);
1046 visitor.visit_ptr(v_ptr);
1047 visitor.visit_ptr(output); visitor.visit_ptr(grad_output); visitor.visit_ptr(params[0]); visitor.visit_ptr(dq_ptr);
1051 visitor.visit_ptr(dk_ptr);
1052 visitor.visit_ptr(dv_ptr);
1053 visitor.visit_i32(n as i32); visitor.visit_f32(softmax_scale);
1055 }
1056
1057 #[cfg(feature = "training")]
1058 fn backward_block(&self) -> [u32; 3] { [1, 1, 1] }
1059
1060 #[cfg(feature = "training")]
1062 fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
1063 [output_shape[2] as u32, (output_shape[0] * output_shape[1]) as u32, 1]
1065 }
1066}