1use std::collections::HashMap;
52
53pub const SPIRV_VERSION_1_3: u32 = 0x0001_0300;
56pub const SPIRV_MAGIC: u32 = 0x0723_0203;
58
59const TWO_OVER_PI_BITS: &[u32] = &[
65 0x0000_0000,
66 0xa2f9_836e,
67 0x4e44_1529,
68 0xfc27_57d1,
69 0xf534_ddc0,
70 0xdb62_9599,
71 0x3c43_9041,
72 0xfe51_63ab,
73 0xdebb_c561,
74 0xb724_6e3a,
75 0x424d_d2e0,
76 0x0649_2eea,
77];
78
79const SINCOS_FAST_RANGE: f32 = 8192.0;
82
83pub const MAX_RANK: usize = 6;
85pub const MAX_ELEMENTWISE_INPUTS: usize = 3;
87
88const OP_EXT_INST_IMPORT: u16 = 11;
90const OP_EXT_INST: u16 = 12;
91const OP_EXTENSION: u16 = 10;
92const OP_MEMORY_MODEL: u16 = 14;
93const OP_ENTRY_POINT: u16 = 15;
94const OP_EXECUTION_MODE: u16 = 16;
95const OP_CAPABILITY: u16 = 17;
96const OP_TYPE_VOID: u16 = 19;
97const OP_TYPE_BOOL: u16 = 20;
98const OP_TYPE_INT: u16 = 21;
99const OP_TYPE_FLOAT: u16 = 22;
100const OP_TYPE_VECTOR: u16 = 23;
101const OP_TYPE_ARRAY: u16 = 28;
102const OP_TYPE_RUNTIME_ARRAY: u16 = 29;
103const OP_TYPE_STRUCT: u16 = 30;
104const OP_TYPE_POINTER: u16 = 32;
105const OP_TYPE_FUNCTION: u16 = 33;
106const OP_CONSTANT_FALSE: u16 = 42;
107const OP_CONSTANT: u16 = 43;
108const OP_CONSTANT_COMPOSITE: u16 = 44;
109const OP_SPEC_CONSTANT: u16 = 50;
110const OP_FUNCTION: u16 = 54;
111const OP_FUNCTION_END: u16 = 56;
112const OP_VARIABLE: u16 = 59;
113const OP_LOAD: u16 = 61;
114const OP_STORE: u16 = 62;
115const OP_ACCESS_CHAIN: u16 = 65;
116const OP_DECORATE: u16 = 71;
117const OP_MEMBER_DECORATE: u16 = 72;
118const OP_CONVERT_F_TO_U: u16 = 109;
119const OP_CONVERT_U_TO_F: u16 = 112;
120const OP_F_CONVERT: u16 = 115;
121const OP_BITCAST: u16 = 124;
122const OP_F_NEGATE: u16 = 127;
123const OP_I_ADD: u16 = 128;
124const OP_F_ADD: u16 = 129;
125const OP_I_SUB: u16 = 130;
126const OP_F_SUB: u16 = 131;
127const OP_I_MUL: u16 = 132;
128const OP_F_MUL: u16 = 133;
129const OP_U_DIV: u16 = 134;
130const OP_F_DIV: u16 = 136;
131const OP_U_MOD: u16 = 137;
132const OP_IS_NAN: u16 = 156;
133const OP_LOGICAL_NOT_EQUAL: u16 = 165;
134const OP_LOGICAL_OR: u16 = 166;
135const OP_LOGICAL_AND: u16 = 167;
136const OP_LOGICAL_NOT: u16 = 168;
137const OP_SELECT: u16 = 169;
138const OP_I_EQUAL: u16 = 170;
139const OP_I_NOT_EQUAL: u16 = 171;
140const OP_U_GREATER_THAN_EQUAL: u16 = 174;
141const OP_U_LESS_THAN: u16 = 176;
142const OP_F_ORD_EQUAL: u16 = 180;
143const OP_F_ORD_LESS_THAN: u16 = 184;
144const OP_F_ORD_GREATER_THAN: u16 = 186;
145const OP_F_ORD_GREATER_THAN_EQUAL: u16 = 190;
146const OP_SHIFT_RIGHT_LOGICAL: u16 = 194;
147const OP_SHIFT_LEFT_LOGICAL: u16 = 196;
148const OP_BITWISE_OR: u16 = 197;
149const OP_BITWISE_XOR: u16 = 198;
150const OP_BITWISE_AND: u16 = 199;
151const OP_NOT: u16 = 200;
152const OP_CONTROL_BARRIER: u16 = 224;
153const OP_ATOMIC_AND: u16 = 240;
154const OP_ATOMIC_OR: u16 = 241;
155const OP_LOOP_MERGE: u16 = 246;
156const OP_SELECTION_MERGE: u16 = 247;
157const OP_LABEL: u16 = 248;
158const OP_BRANCH: u16 = 249;
159const OP_BRANCH_CONDITIONAL: u16 = 250;
160const OP_RETURN: u16 = 253;
161const OP_GROUP_NON_UNIFORM_F_ADD: u16 = 350;
162const OP_TYPE_COOPERATIVE_MATRIX_KHR: u16 = 4456;
163const OP_COOPERATIVE_MATRIX_LOAD_KHR: u16 = 4457;
164const OP_COOPERATIVE_MATRIX_STORE_KHR: u16 = 4458;
165const OP_COOPERATIVE_MATRIX_MUL_ADD_KHR: u16 = 4459;
166
167const CAPABILITY_SHADER: u32 = 1;
169const CAPABILITY_FLOAT16: u32 = 9;
170const CAPABILITY_GROUP_NON_UNIFORM_ARITHMETIC: u32 = 63;
171const CAPABILITY_VULKAN_MEMORY_MODEL: u32 = 5345;
172const CAPABILITY_COOPERATIVE_MATRIX_KHR: u32 = 6022;
173const ADDRESSING_MODEL_LOGICAL: u32 = 0;
174const MEMORY_MODEL_GLSL450: u32 = 1;
175const EXECUTION_MODEL_GL_COMPUTE: u32 = 5;
176const EXECUTION_MODE_LOCAL_SIZE: u32 = 17;
177const STORAGE_CLASS_INPUT: u32 = 1;
178const STORAGE_CLASS_WORKGROUP: u32 = 4;
179const STORAGE_CLASS_PRIVATE: u32 = 6;
180const STORAGE_CLASS_FUNCTION: u32 = 7;
181const STORAGE_CLASS_STORAGE_BUFFER: u32 = 12;
182const DECORATION_SPEC_ID: u32 = 1;
183const DECORATION_BLOCK: u32 = 2;
184const DECORATION_ARRAY_STRIDE: u32 = 6;
185const DECORATION_BUILT_IN: u32 = 11;
186const DECORATION_BINDING: u32 = 33;
187const DECORATION_DESCRIPTOR_SET: u32 = 34;
188const DECORATION_OFFSET: u32 = 35;
189const DECORATION_NO_CONTRACTION: u32 = 42;
190const BUILT_IN_NUM_WORKGROUPS: u32 = 24;
191const BUILT_IN_WORKGROUP_ID: u32 = 26;
192const BUILT_IN_LOCAL_INVOCATION_ID: u32 = 27;
193const BUILT_IN_GLOBAL_INVOCATION_ID: u32 = 28;
194const FUNCTION_CONTROL_NONE: u32 = 0;
195const SELECTION_CONTROL_NONE: u32 = 0;
196const LOOP_CONTROL_NONE: u32 = 0;
197const SCOPE_DEVICE: u32 = 1;
198const SCOPE_WORKGROUP: u32 = 2;
199const MEMORY_SEMANTICS_RELAXED: u32 = 0;
200const MEMORY_SEMANTICS_ACQUIRE_RELEASE_WORKGROUP: u32 = 0x8 | 0x100;
201
202const GLSL_ROUND_EVEN: u32 = 2;
204const GLSL_FABS: u32 = 4;
205const GLSL_FLOOR: u32 = 8;
206const GLSL_CEIL: u32 = 9;
207const GLSL_POW: u32 = 26;
208const GLSL_EXP: u32 = 27;
209const GLSL_LOG: u32 = 28;
210const GLSL_INVERSE_SQRT: u32 = 32;
211const GLSL_UMIN: u32 = 38;
212const GLSL_FMA: u32 = 50;
213const GLSL_FIND_U_MSB: u32 = 75;
214
215pub type Id = u32;
217
218#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
222pub enum Fp8Format {
223 E4M3,
224 E5M2,
225}
226
227#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
229pub enum Storage {
230 Word,
232 Byte,
234 Quarter(Fp8Format),
239 Half,
243}
244
245impl Storage {
246 pub const fn lanes(self) -> u32 {
248 match self {
249 Self::Word => 1,
250 Self::Half => 2,
251 Self::Byte | Self::Quarter(_) => 4,
252 }
253 }
254}
255
256#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
258pub enum NanMode {
259 Propagate,
261 Ignore,
263}
264
265#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
267pub enum ElementwiseOp {
268 Abs,
269 Ceil,
270 Cos,
271 Erf,
272 Exp,
273 Floor,
274 Log,
275 Negate,
276 Reciprocal,
277 Rsqrt,
278 Sin,
279 Sigmoid,
280 Tanh,
281 Clamp(NanMode),
284 Add,
285 Sub,
286 Mul,
287 Pow,
288 Maximum(NanMode),
289 Minimum(NanMode),
290 Equal,
291 Greater,
292 GreaterEqual,
293 LogicalAnd,
294 LogicalOr,
295 LogicalXor,
296 LogicalNot,
297 Select,
298 CopyBytes,
300}
301
302impl ElementwiseOp {
303 pub const fn inputs(self) -> &'static [Storage] {
307 match self {
308 Self::Abs
309 | Self::Ceil
310 | Self::Cos
311 | Self::Erf
312 | Self::Exp
313 | Self::Floor
314 | Self::Log
315 | Self::Negate
316 | Self::Reciprocal
317 | Self::Rsqrt
318 | Self::Sin
319 | Self::Sigmoid
320 | Self::Tanh
321 | Self::Clamp(_) => &[Storage::Word],
322 Self::Add
323 | Self::Sub
324 | Self::Mul
325 | Self::Pow
326 | Self::Maximum(_)
327 | Self::Minimum(_)
328 | Self::Equal
329 | Self::Greater
330 | Self::GreaterEqual => &[Storage::Word, Storage::Word],
331 Self::LogicalAnd | Self::LogicalOr | Self::LogicalXor => {
332 &[Storage::Byte, Storage::Byte]
333 }
334 Self::LogicalNot | Self::CopyBytes => &[Storage::Byte],
335 Self::Select => &[Storage::Byte, Storage::Word, Storage::Word],
336 }
337 }
338
339 pub const fn output(self) -> Storage {
341 match self {
342 Self::Equal
343 | Self::Greater
344 | Self::GreaterEqual
345 | Self::LogicalAnd
346 | Self::LogicalOr
347 | Self::LogicalXor
348 | Self::LogicalNot
349 | Self::CopyBytes => Storage::Byte,
350 _ => Storage::Word,
351 }
352 }
353
354 pub const fn extra_spec_constants(self) -> u32 {
356 match self {
357 Self::Clamp(_) => 2,
358 _ => 0,
359 }
360 }
361}
362
363#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
365pub enum ReduceOp {
366 Sum,
367 Product,
368 Max(NanMode),
369 Min(NanMode),
370 ArgMax(NanMode),
372}
373
374#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
377pub enum KernelKey {
378 Nvfp4Matmul {
380 buffers: u32,
381 cooperative: bool,
382 subgroup: bool,
383 },
384 Elementwise {
390 op: ElementwiseOp,
391 float: Storage,
392 broadcast: bool,
393 workgroup: u32,
394 buffers: u32,
395 },
396 Reduce {
400 op: ReduceOp,
401 float: Storage,
402 workgroup: u32,
403 buffers: u32,
404 },
405 Matmul {
411 input: Storage,
412 output: Storage,
413 tile: u32,
414 buffers: u32,
415 },
416 MatmulStream {
422 rhs: Storage,
423 output: Storage,
424 buffers: u32,
425 },
426 MaxPool {
428 nan_mode: NanMode,
429 float: Storage,
430 workgroup: u32,
431 buffers: u32,
432 },
433 Cast {
436 input: Storage,
437 output: Storage,
438 workgroup: u32,
439 buffers: u32,
440 },
441 Move {
444 storage: Storage,
445 contiguous: bool,
446 workgroup: u32,
447 buffers: u32,
448 },
449}
450
451impl KernelKey {
452 pub fn assemble(self) -> Vec<u32> {
454 match self {
455 Self::Nvfp4Matmul {
456 buffers,
457 cooperative: false,
458 subgroup: false,
459 } => assemble_nvfp4_matmul(buffers),
460 Self::Nvfp4Matmul {
461 buffers,
462 cooperative: true,
463 ..
464 } => assemble_nvfp4_matmul_cooperative(buffers),
465 Self::Nvfp4Matmul {
466 buffers,
467 cooperative: false,
468 subgroup: true,
469 } => assemble_nvfp4_matmul_subgroup(buffers),
470 Self::Elementwise {
471 op,
472 float,
473 broadcast,
474 workgroup,
475 buffers,
476 } => assemble_elementwise(op, float, broadcast, workgroup, buffers),
477 Self::Reduce {
478 op,
479 float,
480 workgroup,
481 buffers,
482 } => assemble_reduce(op, float, workgroup, buffers),
483 Self::Matmul {
484 input,
485 output,
486 tile,
487 buffers,
488 } => assemble_matmul(input, output, MatmulGeometry::wide(tile), buffers),
489 Self::MatmulStream {
490 rhs,
491 output,
492 buffers,
493 } => assemble_matmul_stream(rhs, output, buffers),
494 Self::MaxPool {
495 nan_mode,
496 float,
497 workgroup,
498 buffers,
499 } => assemble_max_pool(nan_mode, float, workgroup, buffers),
500 Self::Move {
501 storage,
502 contiguous,
503 workgroup,
504 buffers,
505 } => assemble_move(storage, contiguous, workgroup, buffers),
506 Self::Cast {
507 input,
508 output,
509 workgroup,
510 buffers,
511 } => assemble_cast(input, output, workgroup, buffers),
512 }
513 }
514
515 pub fn every_variant() -> Vec<KernelKey> {
519 let mut keys = Vec::new();
520 keys.push(KernelKey::Nvfp4Matmul {
521 buffers: 17,
522 cooperative: false,
523 subgroup: false,
524 });
525 keys.push(KernelKey::Nvfp4Matmul {
526 buffers: 17,
527 cooperative: true,
528 subgroup: false,
529 });
530 keys.push(KernelKey::Nvfp4Matmul {
531 buffers: 17,
532 cooperative: false,
533 subgroup: true,
534 });
535 let ops = [
536 ElementwiseOp::Abs,
537 ElementwiseOp::Ceil,
538 ElementwiseOp::Cos,
539 ElementwiseOp::Erf,
540 ElementwiseOp::Exp,
541 ElementwiseOp::Floor,
542 ElementwiseOp::Log,
543 ElementwiseOp::Negate,
544 ElementwiseOp::Reciprocal,
545 ElementwiseOp::Rsqrt,
546 ElementwiseOp::Sin,
547 ElementwiseOp::Sigmoid,
548 ElementwiseOp::Tanh,
549 ElementwiseOp::Clamp(NanMode::Propagate),
550 ElementwiseOp::Clamp(NanMode::Ignore),
551 ElementwiseOp::Add,
552 ElementwiseOp::Sub,
553 ElementwiseOp::Mul,
554 ElementwiseOp::Pow,
555 ElementwiseOp::Maximum(NanMode::Propagate),
556 ElementwiseOp::Minimum(NanMode::Ignore),
557 ElementwiseOp::Equal,
558 ElementwiseOp::Greater,
559 ElementwiseOp::GreaterEqual,
560 ElementwiseOp::LogicalAnd,
561 ElementwiseOp::LogicalOr,
562 ElementwiseOp::LogicalXor,
563 ElementwiseOp::LogicalNot,
564 ElementwiseOp::Select,
565 ElementwiseOp::CopyBytes,
566 ];
567 for op in ops {
568 for float in [Storage::Word, Storage::Half] {
569 for broadcast in [false, true] {
570 keys.push(KernelKey::Elementwise {
571 op,
572 float,
573 broadcast,
574 workgroup: 64,
575 buffers: 17,
576 });
577 }
578 }
579 }
580 for op in [
581 ReduceOp::Sum,
582 ReduceOp::Product,
583 ReduceOp::Max(NanMode::Propagate),
584 ReduceOp::Min(NanMode::Ignore),
585 ReduceOp::ArgMax(NanMode::Propagate),
586 ReduceOp::ArgMax(NanMode::Ignore),
587 ] {
588 for float in [Storage::Word, Storage::Half] {
589 keys.push(KernelKey::Reduce {
590 op,
591 float,
592 workgroup: 64,
593 buffers: 17,
594 });
595 }
596 }
597 for float in [Storage::Word, Storage::Half] {
598 keys.push(KernelKey::Matmul {
599 input: float,
600 output: float,
601 tile: 16,
602 buffers: 17,
603 });
604 keys.push(KernelKey::Matmul {
605 input: float,
606 output: float,
607 tile: 8,
608 buffers: 5,
609 });
610 }
611 for float in [Storage::Word, Storage::Half] {
612 keys.push(KernelKey::MatmulStream {
613 rhs: float,
614 output: float,
615 buffers: 17,
616 });
617 }
618 keys.push(KernelKey::Cast {
620 input: Storage::Half,
621 output: Storage::Word,
622 workgroup: 64,
623 buffers: 17,
624 });
625 keys.push(KernelKey::Cast {
626 input: Storage::Word,
627 output: Storage::Half,
628 workgroup: 64,
629 buffers: 17,
630 });
631 for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
633 keys.push(KernelKey::MatmulStream {
634 rhs: Storage::Quarter(format),
635 output: Storage::Half,
636 buffers: 17,
637 });
638 keys.push(KernelKey::Matmul {
639 input: Storage::Quarter(format),
640 output: Storage::Half,
641 tile: 16,
642 buffers: 17,
643 });
644 keys.push(KernelKey::Matmul {
645 input: Storage::Quarter(format),
646 output: Storage::Half,
647 tile: 8,
648 buffers: 5,
649 });
650 }
651 for nan_mode in [NanMode::Propagate, NanMode::Ignore] {
652 for float in [Storage::Word, Storage::Half] {
653 keys.push(KernelKey::MaxPool {
654 nan_mode,
655 float,
656 workgroup: 64,
657 buffers: 17,
658 });
659 }
660 }
661 for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
664 for nan_mode in [NanMode::Propagate, NanMode::Ignore] {
665 keys.push(KernelKey::MaxPool {
666 nan_mode,
667 float: Storage::Quarter(format),
668 workgroup: 64,
669 buffers: 17,
670 });
671 keys.push(KernelKey::Reduce {
672 op: ReduceOp::ArgMax(nan_mode),
673 float: Storage::Quarter(format),
674 workgroup: 64,
675 buffers: 17,
676 });
677 }
678 }
679 for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
681 for wide in [Storage::Word, Storage::Half] {
682 keys.push(KernelKey::Cast {
683 input: Storage::Quarter(format),
684 output: wide,
685 workgroup: 64,
686 buffers: 17,
687 });
688 keys.push(KernelKey::Cast {
689 input: wide,
690 output: Storage::Quarter(format),
691 workgroup: 64,
692 buffers: 17,
693 });
694 }
695 }
696 for storage in [
697 Storage::Word,
698 Storage::Byte,
699 Storage::Half,
700 Storage::Quarter(Fp8Format::E4M3),
701 Storage::Quarter(Fp8Format::E5M2),
702 ] {
703 for contiguous in [false, true] {
704 keys.push(KernelKey::Move {
705 storage,
706 contiguous,
707 workgroup: 64,
708 buffers: 17,
709 });
710 }
711 }
712 keys
713 }
714
715 pub const fn spec_constant_count(self) -> u32 {
717 match self {
718 Self::Elementwise { op, broadcast, .. } => {
719 let inputs = op.inputs().len() as u32;
720 let base = 1 + 2 * inputs + 2;
721 let shape = if broadcast {
722 MAX_RANK as u32 * (1 + inputs)
723 } else {
724 0
725 };
726 base + shape + op.extra_spec_constants()
727 }
728 Self::Reduce { .. } => 2 + 2 + 3,
729 Self::Matmul { .. } | Self::MatmulStream { .. } => 3 * 2 + 4,
730 Self::Nvfp4Matmul { .. } => 15 * 2 + 5,
733 Self::MaxPool { .. } => 2 + 2 + 12,
734 Self::Move { contiguous, .. } => {
735 if contiguous {
736 2 + 2 + 1
737 } else {
738 2 + 2 + 1 + MAX_RANK as u32 * 3 + 2
739 }
740 }
741 Self::Cast { .. } => 2 + 2 + 1,
742 }
743 }
744
745 pub const fn local_size(self) -> [u32; 3] {
747 match self {
748 Self::Elementwise { workgroup, .. }
749 | Self::Reduce { workgroup, .. }
750 | Self::MaxPool { workgroup, .. }
751 | Self::Move { workgroup, .. }
752 | Self::Cast { workgroup, .. } => [workgroup, 1, 1],
753 Self::Matmul { tile, .. } => MatmulGeometry::wide(tile).local_size(),
754 Self::MatmulStream { .. } => [STREAM_WORKGROUP, 1, 1],
755 Self::Nvfp4Matmul {
756 cooperative: false,
757 subgroup: false,
758 ..
759 } => [STREAM_WORKGROUP, 1, 1],
760 Self::Nvfp4Matmul {
761 cooperative: true, ..
762 } => [32, 1, 1],
763 Self::Nvfp4Matmul {
764 cooperative: false,
765 subgroup: true,
766 ..
767 } => [32, 1, 1],
768 }
769 }
770}
771
772#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
774pub struct Operand {
775 pub buffer: u32,
776 pub base: u32,
777}
778
779impl Operand {
780 fn push(self, words: &mut Vec<u32>) {
781 words.push(self.buffer);
782 words.push(self.base);
783 }
784}
785
786#[derive(Clone, Debug, PartialEq, Eq)]
790pub struct ElementwiseSpec<'a> {
791 pub count: u32,
792 pub inputs: &'a [Operand],
793 pub output: Operand,
794 pub dims: [u32; MAX_RANK],
796 pub strides: &'a [[u32; MAX_RANK]],
798 pub clamp: Option<[u32; 2]>,
800}
801
802impl ElementwiseSpec<'_> {
803 pub fn words(&self, broadcast: bool) -> Vec<u32> {
804 let mut words = vec![self.count];
805 for input in self.inputs {
806 input.push(&mut words);
807 }
808 self.output.push(&mut words);
809 if broadcast {
810 words.extend_from_slice(&self.dims);
811 for strides in self.strides {
812 words.extend_from_slice(strides);
813 }
814 }
815 if let Some(bounds) = self.clamp {
816 words.extend_from_slice(&bounds);
817 }
818 words
819 }
820}
821
822pub fn reduce_spec(input: Operand, output: Operand, outer: u32, axis: u32, inner: u32) -> Vec<u32> {
824 let mut words = Vec::with_capacity(7);
825 input.push(&mut words);
826 output.push(&mut words);
827 words.extend_from_slice(&[outer, axis, inner]);
828 words
829}
830
831pub fn matmul_spec(
833 lhs: Operand,
834 rhs: Operand,
835 output: Operand,
836 m: u32,
837 n: u32,
838 k: u32,
839 batch: u32,
840) -> Vec<u32> {
841 let mut words = Vec::with_capacity(10);
842 lhs.push(&mut words);
843 rhs.push(&mut words);
844 output.push(&mut words);
845 words.extend_from_slice(&[m, n, k, batch]);
846 words
847}
848
849pub struct Nvfp4MatmulSpec<'a> {
852 pub activation: Operand,
853 pub packed: &'a [Operand],
854 pub block_scales: &'a [Operand],
855 pub tensor_scale: Operand,
856 pub output: Operand,
857 pub m: u32,
858 pub n: u32,
859 pub k: u32,
860 pub epilogue: u32,
861 pub weight_mode: u32,
862}
863
864impl Nvfp4MatmulSpec<'_> {
865 pub fn words(&self) -> Vec<u32> {
866 assert!(
867 !self.packed.is_empty()
868 && self.packed.len() <= 6
869 && self.packed.len() == self.block_scales.len()
870 );
871 let mut words = Vec::with_capacity(35);
872 self.activation.push(&mut words);
873 for index in 0..6 {
874 self.packed[index.min(self.packed.len() - 1)].push(&mut words);
875 }
876 for index in 0..6 {
877 self.block_scales[index.min(self.block_scales.len() - 1)].push(&mut words);
878 }
879 self.tensor_scale.push(&mut words);
880 self.output.push(&mut words);
881 words.extend_from_slice(&[self.m, self.n, self.k, self.epilogue, self.weight_mode]);
882 words
883 }
884}
885
886#[derive(Clone, Copy, Debug, PartialEq, Eq)]
889pub struct PoolGeometry {
890 pub batch: u32,
891 pub height: u32,
892 pub width: u32,
893 pub channels: u32,
894 pub out_height: u32,
895 pub out_width: u32,
896 pub kernel: [u32; 2],
897 pub stride: [u32; 2],
898 pub pad_top: u32,
899 pub pad_left: u32,
900}
901
902pub fn max_pool_spec(input: Operand, output: Operand, geometry: PoolGeometry) -> Vec<u32> {
904 let mut words = Vec::with_capacity(16);
905 input.push(&mut words);
906 output.push(&mut words);
907 words.extend_from_slice(&[
908 geometry.batch,
909 geometry.height,
910 geometry.width,
911 geometry.channels,
912 geometry.out_height,
913 geometry.out_width,
914 geometry.kernel[0],
915 geometry.kernel[1],
916 geometry.stride[0],
917 geometry.stride[1],
918 geometry.pad_top,
919 geometry.pad_left,
920 ]);
921 words
922}
923
924#[derive(Clone, Copy, Debug, PartialEq, Eq)]
927pub struct MoveGeometry {
928 pub count: u32,
929 pub dims: [u32; MAX_RANK],
930 pub in_strides: [u32; MAX_RANK],
931 pub in_offset: u32,
932 pub out_strides: [u32; MAX_RANK],
933 pub out_offset: u32,
934}
935
936pub fn move_spec(
938 input: Operand,
939 output: Operand,
940 geometry: MoveGeometry,
941 contiguous: bool,
942) -> Vec<u32> {
943 let mut words = Vec::with_capacity(25);
944 input.push(&mut words);
945 output.push(&mut words);
946 words.push(geometry.count);
947 if !contiguous {
948 words.extend_from_slice(&geometry.dims);
949 words.extend_from_slice(&geometry.in_strides);
950 words.push(geometry.in_offset);
951 words.extend_from_slice(&geometry.out_strides);
952 words.push(geometry.out_offset);
953 }
954 words
955}
956
957pub const fn linear_workgroups(count: u32, workgroup: u32, limit: u32) -> u32 {
961 let needed = count.div_ceil(workgroup);
962 if needed == 0 {
963 1
964 } else if needed > limit {
965 limit
966 } else {
967 needed
968 }
969}
970
971pub const MATMUL_MICRO: u32 = 4;
975
976pub const STREAM_ROWS: u32 = 8;
979
980pub const STREAM_COLUMNS: u32 = 16;
982
983pub const STREAM_WORKGROUP: u32 = 64;
986
987pub const fn stream_matmul_shared_bytes() -> u32 {
990 STREAM_WORKGROUP * STREAM_ROWS * 4 * 4
991}
992
993pub const fn stream_matmul_workgroups(n: u32, batch: u32) -> [u32; 3] {
995 [n.div_ceil(STREAM_COLUMNS), 1, batch]
996}
997
998pub const fn matmul_block(tile: u32) -> u32 {
1000 tile * MATMUL_MICRO
1001}
1002
1003pub const fn matmul_shared_bytes(tile: u32) -> u32 {
1006 let wide = MatmulGeometry::wide(tile).shared_bytes();
1007 let stream = stream_matmul_shared_bytes();
1008 if wide > stream { wide } else { stream }
1009}
1010
1011pub const fn matmul_workgroups(m: u32, n: u32, batch: u32, tile: u32) -> [u32; 3] {
1013 let block = matmul_block(tile);
1014 [n.div_ceil(block), m.div_ceil(block), batch]
1015}
1016
1017#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1021pub struct MatmulGeometry {
1022 pub tile_x: u32,
1023 pub tile_y: u32,
1024 pub micro_m: u32,
1025 pub micro_n: u32,
1026 pub depth: u32,
1027}
1028
1029impl MatmulGeometry {
1030 pub const fn wide(tile: u32) -> Self {
1033 Self {
1034 tile_x: tile,
1035 tile_y: tile,
1036 micro_m: MATMUL_MICRO,
1037 micro_n: MATMUL_MICRO,
1038 depth: tile,
1039 }
1040 }
1041
1042 pub const fn invocations(self) -> u32 {
1043 self.tile_x * self.tile_y
1044 }
1045
1046 pub const fn block_m(self) -> u32 {
1047 self.tile_y * self.micro_m
1048 }
1049
1050 pub const fn block_n(self) -> u32 {
1051 self.tile_x * self.micro_n
1052 }
1053
1054 pub const fn local_size(self) -> [u32; 3] {
1055 [self.tile_x, self.tile_y, 1]
1056 }
1057
1058 pub const fn shared_bytes(self) -> u32 {
1060 (self.block_m() + self.block_n()) * self.depth * 4
1061 }
1062}
1063
1064pub fn f32_to_fp8_bits(format: Fp8Format, value: f32) -> u8 {
1085 let (mantissa_bits, bias, nan_out, max_finite): (u32, u32, u8, u32) = match format {
1086 Fp8Format::E4M3 => (3, 7, 0x7f, 0x7e),
1088 Fp8Format::E5M2 => (2, 15, 0x7e, 0x7b),
1090 };
1091 let bits = value.to_bits();
1092 let sign = ((bits >> 24) as u8) & 0x80;
1093 let magnitude = bits & 0x7fff_ffff;
1094 if magnitude > 0x7f80_0000 {
1095 return sign | nan_out;
1096 }
1097 let shift = 23 - mantissa_bits;
1098 let normal_floor = (128 - bias) << 23;
1099 let body = if magnitude >= normal_floor {
1100 let adjusted = magnitude - ((127 - bias) << 23);
1103 let guard = (adjusted >> shift) & 1;
1104 let half_ulp = (1 << (shift - 1)) - 1;
1105 (adjusted + half_ulp + guard) >> shift
1106 } else {
1107 let scale = f32::from_bits((127 + bias - 1 + mantissa_bits) << 23);
1111 (f32::from_bits(magnitude) * scale).round_ties_even() as u32
1112 };
1113 if body > max_finite {
1116 let overflow = match format {
1117 Fp8Format::E4M3 => nan_out,
1118 Fp8Format::E5M2 => 0x7c,
1119 };
1120 return sign | overflow;
1121 }
1122 sign | body as u8
1123}
1124
1125pub fn f16_to_f32(bits: u16) -> f32 {
1126 let bits = u32::from(bits);
1127 let sign = (bits & 0x8000) << 16;
1128 let exponent = (bits >> 10) & 0x1f;
1129 let mantissa = bits & 0x3ff;
1130 let converted = if exponent == 0 {
1131 (mantissa as f32) * (1.0 / 16_777_216.0)
1133 } else if exponent == 31 {
1134 f32::from_bits(0x7f80_0000 | (mantissa << 13))
1135 } else {
1136 f32::from_bits(((exponent + 112) << 23) | (mantissa << 13))
1137 };
1138 f32::from_bits(converted.to_bits() | sign)
1139}
1140
1141pub fn f32_to_f16_bits(value: f32) -> u16 {
1145 let bits = value.to_bits();
1146 let sign = ((bits >> 16) & 0x8000) as u16;
1147 let magnitude = bits & 0x7fff_ffff;
1148 if magnitude > 0x7f80_0000 {
1149 return sign | 0x7e00;
1150 }
1151 if magnitude >= 0x477f_f000 {
1152 return sign | 0x7c00;
1155 }
1156 if magnitude >= 0x3880_0000 {
1157 let adjusted = magnitude - 0x3800_0000;
1160 let rounded = adjusted + 0x0000_0fff + ((adjusted >> 13) & 1);
1161 return sign | (rounded >> 13) as u16;
1162 }
1163 let scaled = f32::from_bits(magnitude) * 16_777_216.0;
1166 let truncated = scaled as u16;
1167 let fraction = scaled - f32::from(truncated);
1168 let rounded = if fraction > 0.5 || (fraction == 0.5 && truncated & 1 != 0) {
1169 truncated + 1
1170 } else {
1171 truncated
1172 };
1173 sign | rounded
1174}
1175
1176#[derive(Clone, PartialEq, Eq, Hash)]
1181enum TypeKey {
1182 Void,
1183 Bool,
1184 U32,
1185 F32,
1186 F16,
1187 Vector(Id, u32),
1188 Pointer(u32, Id),
1189 RuntimeArray(Id),
1190 Array(Id, Id),
1191 Function(Id),
1192}
1193
1194#[derive(Clone, Copy, PartialEq, Eq, Hash)]
1195enum ConstKey {
1196 U32(u32),
1197 F32(u32),
1198 False,
1199}
1200
1201struct Builder {
1204 next_id: Id,
1205 capabilities: Vec<u32>,
1206 extensions: Vec<u32>,
1207 imports: Vec<u32>,
1208 memory_model: Vec<u32>,
1209 entry_point: Vec<u32>,
1210 execution_modes: Vec<u32>,
1211 annotations: Vec<u32>,
1212 declarations: Vec<u32>,
1213 functions: Vec<u32>,
1214 types: HashMap<TypeKey, Id>,
1215 constants: HashMap<ConstKey, Id>,
1216 glsl: Id,
1217 spec_next: u32,
1218 local_variable_cursor: usize,
1220 interface: Vec<Id>,
1221 main: Id,
1222}
1223
1224fn instruction(target: &mut Vec<u32>, opcode: u16, operands: &[u32]) {
1225 let word_count = u32::try_from(operands.len() + 1).expect("instruction fits");
1226 target.push((word_count << 16) | u32::from(opcode));
1227 target.extend_from_slice(operands);
1228}
1229
1230fn literal_string(text: &str) -> Vec<u32> {
1232 let mut bytes = text.as_bytes().to_vec();
1233 bytes.push(0);
1234 while bytes.len() % 4 != 0 {
1235 bytes.push(0);
1236 }
1237 bytes
1238 .chunks_exact(4)
1239 .map(|chunk| u32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
1240 .collect()
1241}
1242
1243impl Builder {
1244 fn new() -> Self {
1245 let mut builder = Self {
1246 next_id: 1,
1247 capabilities: Vec::new(),
1248 extensions: Vec::new(),
1249 imports: Vec::new(),
1250 memory_model: Vec::new(),
1251 entry_point: Vec::new(),
1252 execution_modes: Vec::new(),
1253 annotations: Vec::new(),
1254 declarations: Vec::new(),
1255 functions: Vec::new(),
1256 types: HashMap::new(),
1257 constants: HashMap::new(),
1258 glsl: 0,
1259 spec_next: 0,
1260 local_variable_cursor: 0,
1261 interface: Vec::new(),
1262 main: 0,
1263 };
1264 instruction(
1265 &mut builder.capabilities,
1266 OP_CAPABILITY,
1267 &[CAPABILITY_SHADER],
1268 );
1269 builder.glsl = builder.id();
1270 let mut import = vec![builder.glsl];
1271 import.extend(literal_string("GLSL.std.450"));
1272 instruction(&mut builder.imports, OP_EXT_INST_IMPORT, &import);
1273 instruction(
1274 &mut builder.memory_model,
1275 OP_MEMORY_MODEL,
1276 &[ADDRESSING_MODEL_LOGICAL, MEMORY_MODEL_GLSL450],
1277 );
1278 builder.main = builder.id();
1279 builder
1280 }
1281
1282 fn id(&mut self) -> Id {
1283 let id = self.next_id;
1284 self.next_id += 1;
1285 id
1286 }
1287
1288 fn finish(mut self, local_size: [u32; 3]) -> Vec<u32> {
1289 let mut entry = vec![EXECUTION_MODEL_GL_COMPUTE, self.main];
1290 entry.extend(literal_string("main"));
1291 entry.extend_from_slice(&self.interface);
1292 instruction(&mut self.entry_point, OP_ENTRY_POINT, &entry);
1293 instruction(
1294 &mut self.execution_modes,
1295 OP_EXECUTION_MODE,
1296 &[
1297 self.main,
1298 EXECUTION_MODE_LOCAL_SIZE,
1299 local_size[0],
1300 local_size[1],
1301 local_size[2],
1302 ],
1303 );
1304 let mut words = vec![SPIRV_MAGIC, SPIRV_VERSION_1_3, 0, self.next_id, 0];
1305 for section in [
1306 &self.capabilities,
1307 &self.extensions,
1308 &self.imports,
1309 &self.memory_model,
1310 &self.entry_point,
1311 &self.execution_modes,
1312 &self.annotations,
1313 &self.declarations,
1314 &self.functions,
1315 ] {
1316 words.extend_from_slice(section);
1317 }
1318 words
1319 }
1320
1321 fn ty(&mut self, key: TypeKey) -> Id {
1324 if let Some(id) = self.types.get(&key) {
1325 return *id;
1326 }
1327 let id = self.id();
1328 match &key {
1329 TypeKey::Void => instruction(&mut self.declarations, OP_TYPE_VOID, &[id]),
1330 TypeKey::Bool => instruction(&mut self.declarations, OP_TYPE_BOOL, &[id]),
1331 TypeKey::U32 => instruction(&mut self.declarations, OP_TYPE_INT, &[id, 32, 0]),
1332 TypeKey::F32 => instruction(&mut self.declarations, OP_TYPE_FLOAT, &[id, 32]),
1333 TypeKey::F16 => instruction(&mut self.declarations, OP_TYPE_FLOAT, &[id, 16]),
1334 TypeKey::Vector(element, count) => {
1335 instruction(
1336 &mut self.declarations,
1337 OP_TYPE_VECTOR,
1338 &[id, *element, *count],
1339 );
1340 }
1341 TypeKey::Pointer(class, pointee) => {
1342 instruction(
1343 &mut self.declarations,
1344 OP_TYPE_POINTER,
1345 &[id, *class, *pointee],
1346 );
1347 }
1348 TypeKey::RuntimeArray(element) => {
1349 instruction(
1350 &mut self.declarations,
1351 OP_TYPE_RUNTIME_ARRAY,
1352 &[id, *element],
1353 );
1354 }
1355 TypeKey::Array(element, length) => {
1356 instruction(
1357 &mut self.declarations,
1358 OP_TYPE_ARRAY,
1359 &[id, *element, *length],
1360 );
1361 }
1362 TypeKey::Function(ret) => {
1363 instruction(&mut self.declarations, OP_TYPE_FUNCTION, &[id, *ret]);
1364 }
1365 }
1366 self.types.insert(key, id);
1367 id
1368 }
1369
1370 fn void(&mut self) -> Id {
1371 self.ty(TypeKey::Void)
1372 }
1373 fn bool_ty(&mut self) -> Id {
1374 self.ty(TypeKey::Bool)
1375 }
1376 fn u32_ty(&mut self) -> Id {
1377 self.ty(TypeKey::U32)
1378 }
1379 fn f32_ty(&mut self) -> Id {
1380 self.ty(TypeKey::F32)
1381 }
1382 fn f16_ty(&mut self) -> Id {
1383 self.ty(TypeKey::F16)
1384 }
1385 fn uvec3(&mut self) -> Id {
1386 let u32_ty = self.u32_ty();
1387 self.ty(TypeKey::Vector(u32_ty, 3))
1388 }
1389 fn pointer(&mut self, class: u32, pointee: Id) -> Id {
1390 self.ty(TypeKey::Pointer(class, pointee))
1391 }
1392
1393 fn constant(&mut self, key: ConstKey) -> Id {
1394 if let Some(id) = self.constants.get(&key) {
1395 return *id;
1396 }
1397 let id = self.id();
1398 match key {
1399 ConstKey::U32(value) => {
1400 let ty = self.u32_ty();
1401 instruction(&mut self.declarations, OP_CONSTANT, &[ty, id, value]);
1402 }
1403 ConstKey::F32(bits) => {
1404 let ty = self.f32_ty();
1405 instruction(&mut self.declarations, OP_CONSTANT, &[ty, id, bits]);
1406 }
1407 ConstKey::False => {
1408 let ty = self.bool_ty();
1409 instruction(&mut self.declarations, OP_CONSTANT_FALSE, &[ty, id]);
1410 }
1411 }
1412 self.constants.insert(key, id);
1413 id
1414 }
1415
1416 fn c_u32(&mut self, value: u32) -> Id {
1417 self.constant(ConstKey::U32(value))
1418 }
1419 fn c_f32(&mut self, value: f32) -> Id {
1420 self.constant(ConstKey::F32(value.to_bits()))
1421 }
1422 fn c_false(&mut self) -> Id {
1423 self.constant(ConstKey::False)
1424 }
1425
1426 fn spec_u32(&mut self, default: u32) -> Id {
1428 let ty = self.u32_ty();
1429 let id = self.id();
1430 instruction(&mut self.declarations, OP_SPEC_CONSTANT, &[ty, id, default]);
1431 instruction(
1432 &mut self.annotations,
1433 OP_DECORATE,
1434 &[id, DECORATION_SPEC_ID, self.spec_next],
1435 );
1436 self.spec_next += 1;
1437 id
1438 }
1439
1440 fn spec_operand(&mut self) -> (Id, Id) {
1441 let buffer = self.spec_u32(0);
1442 let base = self.spec_u32(0);
1443 (buffer, base)
1444 }
1445
1446 fn spec_dims(&mut self) -> [Id; MAX_RANK] {
1447 let mut ids = [0; MAX_RANK];
1448 for id in &mut ids {
1449 *id = self.spec_u32(1);
1450 }
1451 ids
1452 }
1453
1454 fn spec_strides(&mut self) -> [Id; MAX_RANK] {
1455 let mut ids = [0; MAX_RANK];
1456 for id in &mut ids {
1457 *id = self.spec_u32(0);
1458 }
1459 ids
1460 }
1461
1462 fn builtin_uvec3(&mut self, builtin: u32) -> Id {
1465 let uvec3 = self.uvec3();
1466 let pointer = self.pointer(STORAGE_CLASS_INPUT, uvec3);
1467 let id = self.id();
1468 instruction(
1469 &mut self.declarations,
1470 OP_VARIABLE,
1471 &[pointer, id, STORAGE_CLASS_INPUT],
1472 );
1473 instruction(
1474 &mut self.annotations,
1475 OP_DECORATE,
1476 &[id, DECORATION_BUILT_IN, builtin],
1477 );
1478 self.interface.push(id);
1479 id
1480 }
1481
1482 fn buffer_array(&mut self, buffers: u32) -> Id {
1484 let u32_ty = self.u32_ty();
1485 let words = self.ty(TypeKey::RuntimeArray(u32_ty));
1486 instruction(
1487 &mut self.annotations,
1488 OP_DECORATE,
1489 &[words, DECORATION_ARRAY_STRIDE, 4],
1490 );
1491 let block = self.id();
1492 instruction(&mut self.declarations, OP_TYPE_STRUCT, &[block, words]);
1493 instruction(
1494 &mut self.annotations,
1495 OP_DECORATE,
1496 &[block, DECORATION_BLOCK],
1497 );
1498 instruction(
1499 &mut self.annotations,
1500 OP_MEMBER_DECORATE,
1501 &[block, 0, DECORATION_OFFSET, 0],
1502 );
1503 let length = self.c_u32(buffers);
1504 let array = self.ty(TypeKey::Array(block, length));
1505 let pointer = self.pointer(STORAGE_CLASS_STORAGE_BUFFER, array);
1506 let variable = self.id();
1507 instruction(
1508 &mut self.declarations,
1509 OP_VARIABLE,
1510 &[pointer, variable, STORAGE_CLASS_STORAGE_BUFFER],
1511 );
1512 instruction(
1513 &mut self.annotations,
1514 OP_DECORATE,
1515 &[variable, DECORATION_DESCRIPTOR_SET, 0],
1516 );
1517 instruction(
1518 &mut self.annotations,
1519 OP_DECORATE,
1520 &[variable, DECORATION_BINDING, 0],
1521 );
1522 variable
1523 }
1524
1525 fn private_u32_array(&mut self, values: &[u32]) -> Id {
1527 let u32_ty = self.u32_ty();
1528 let length = self.c_u32(values.len() as u32);
1529 let array = self.ty(TypeKey::Array(u32_ty, length));
1530 let elements: Vec<Id> = values.iter().map(|value| self.c_u32(*value)).collect();
1531 let composite = self.id();
1532 let mut operands = vec![array, composite];
1533 operands.extend_from_slice(&elements);
1534 instruction(&mut self.declarations, OP_CONSTANT_COMPOSITE, &operands);
1535 let pointer = self.pointer(STORAGE_CLASS_PRIVATE, array);
1536 let variable = self.id();
1537 instruction(
1538 &mut self.declarations,
1539 OP_VARIABLE,
1540 &[pointer, variable, STORAGE_CLASS_PRIVATE, composite],
1541 );
1542 variable
1543 }
1544
1545 fn shared_f32_array(&mut self, length: u32) -> Id {
1547 let f32_ty = self.f32_ty();
1548 let length = self.c_u32(length);
1549 let array = self.ty(TypeKey::Array(f32_ty, length));
1550 let pointer = self.pointer(STORAGE_CLASS_WORKGROUP, array);
1551 let variable = self.id();
1552 instruction(
1553 &mut self.declarations,
1554 OP_VARIABLE,
1555 &[pointer, variable, STORAGE_CLASS_WORKGROUP],
1556 );
1557 variable
1558 }
1559
1560 fn shared_f16_array(&mut self, length: u32) -> Id {
1562 let f16_ty = self.f16_ty();
1563 let length = self.c_u32(length);
1564 let array = self.ty(TypeKey::Array(f16_ty, length));
1565 let pointer = self.pointer(STORAGE_CLASS_WORKGROUP, array);
1566 let variable = self.id();
1567 instruction(
1568 &mut self.declarations,
1569 OP_VARIABLE,
1570 &[pointer, variable, STORAGE_CLASS_WORKGROUP],
1571 );
1572 variable
1573 }
1574
1575 fn enable_cooperative_matrix(&mut self) {
1576 instruction(&mut self.capabilities, OP_CAPABILITY, &[CAPABILITY_FLOAT16]);
1577 instruction(
1578 &mut self.capabilities,
1579 OP_CAPABILITY,
1580 &[CAPABILITY_COOPERATIVE_MATRIX_KHR],
1581 );
1582 instruction(
1583 &mut self.capabilities,
1584 OP_CAPABILITY,
1585 &[CAPABILITY_VULKAN_MEMORY_MODEL],
1586 );
1587 let memory_extension = literal_string("SPV_KHR_vulkan_memory_model");
1588 instruction(&mut self.extensions, OP_EXTENSION, &memory_extension);
1589 let extension = literal_string("SPV_KHR_cooperative_matrix");
1590 instruction(&mut self.extensions, OP_EXTENSION, &extension);
1591 self.memory_model.clear();
1592 instruction(
1593 &mut self.memory_model,
1594 OP_MEMORY_MODEL,
1595 &[ADDRESSING_MODEL_LOGICAL, 3], );
1597 }
1598
1599 fn enable_subgroup_arithmetic(&mut self) {
1600 instruction(
1601 &mut self.capabilities,
1602 OP_CAPABILITY,
1603 &[CAPABILITY_GROUP_NON_UNIFORM_ARITHMETIC],
1604 );
1605 }
1606
1607 fn cooperative_matrix_ty(&mut self, component: Id, rows: u32, columns: u32, usage: u32) -> Id {
1608 let ty = self.id();
1609 let scope = self.c_u32(3); let rows = self.c_u32(rows);
1611 let columns = self.c_u32(columns);
1612 let usage = self.c_u32(usage);
1613 instruction(
1614 &mut self.declarations,
1615 OP_TYPE_COOPERATIVE_MATRIX_KHR,
1616 &[ty, component, scope, rows, columns, usage],
1617 );
1618 ty
1619 }
1620
1621 fn begin_main(&mut self) {
1624 let void = self.void();
1625 let fn_type = self.ty(TypeKey::Function(void));
1626 instruction(
1627 &mut self.functions,
1628 OP_FUNCTION,
1629 &[void, self.main, FUNCTION_CONTROL_NONE, fn_type],
1630 );
1631 let entry = self.id();
1632 instruction(&mut self.functions, OP_LABEL, &[entry]);
1633 self.local_variable_cursor = self.functions.len();
1634 }
1635
1636 fn end_main(&mut self) {
1637 instruction(&mut self.functions, OP_RETURN, &[]);
1638 instruction(&mut self.functions, OP_FUNCTION_END, &[]);
1639 }
1640
1641 fn local(&mut self, ty: Id) -> Id {
1643 let pointer = self.pointer(STORAGE_CLASS_FUNCTION, ty);
1644 let id = self.id();
1645 let mut declaration = Vec::with_capacity(4);
1646 instruction(
1647 &mut declaration,
1648 OP_VARIABLE,
1649 &[pointer, id, STORAGE_CLASS_FUNCTION],
1650 );
1651 let cursor = self.local_variable_cursor;
1652 self.functions
1653 .splice(cursor..cursor, declaration.iter().copied());
1654 self.local_variable_cursor += declaration.len();
1655 id
1656 }
1657
1658 fn emit(&mut self, opcode: u16, operands: &[u32]) {
1659 instruction(&mut self.functions, opcode, operands);
1660 }
1661
1662 fn value(&mut self, opcode: u16, ty: Id, operands: &[u32]) -> Id {
1663 let id = self.id();
1664 let mut all = Vec::with_capacity(operands.len() + 2);
1665 all.push(ty);
1666 all.push(id);
1667 all.extend_from_slice(operands);
1668 instruction(&mut self.functions, opcode, &all);
1669 id
1670 }
1671
1672 fn float_value(&mut self, opcode: u16, operands: &[u32]) -> Id {
1673 let f32_ty = self.f32_ty();
1674 self.float_typed(f32_ty, opcode, operands)
1675 }
1676
1677 fn float_typed(&mut self, ty: Id, opcode: u16, operands: &[u32]) -> Id {
1680 let id = self.value(opcode, ty, operands);
1681 instruction(
1682 &mut self.annotations,
1683 OP_DECORATE,
1684 &[id, DECORATION_NO_CONTRACTION],
1685 );
1686 id
1687 }
1688
1689 fn load(&mut self, ty: Id, pointer: Id) -> Id {
1690 self.value(OP_LOAD, ty, &[pointer])
1691 }
1692 fn store(&mut self, pointer: Id, value: Id) {
1693 self.emit(OP_STORE, &[pointer, value]);
1694 }
1695 fn access_chain(&mut self, pointer_ty: Id, base: Id, indices: &[Id]) -> Id {
1696 let mut operands = vec![base];
1697 operands.extend_from_slice(indices);
1698 self.value(OP_ACCESS_CHAIN, pointer_ty, &operands)
1699 }
1700
1701 fn builtin_component(&mut self, variable: Id, index: u32) -> Id {
1703 let u32_ty = self.u32_ty();
1704 let pointer = self.pointer(STORAGE_CLASS_INPUT, u32_ty);
1705 let component = self.c_u32(index);
1706 let chain = self.access_chain(pointer, variable, &[component]);
1707 self.load(u32_ty, chain)
1708 }
1709
1710 fn iadd(&mut self, a: Id, b: Id) -> Id {
1711 let ty = self.u32_ty();
1712 self.value(OP_I_ADD, ty, &[a, b])
1713 }
1714 fn isub(&mut self, a: Id, b: Id) -> Id {
1715 let ty = self.u32_ty();
1716 self.value(OP_I_SUB, ty, &[a, b])
1717 }
1718 fn imul(&mut self, a: Id, b: Id) -> Id {
1719 let ty = self.u32_ty();
1720 self.value(OP_I_MUL, ty, &[a, b])
1721 }
1722 fn udiv(&mut self, a: Id, b: Id) -> Id {
1723 let ty = self.u32_ty();
1724 self.value(OP_U_DIV, ty, &[a, b])
1725 }
1726 fn umod(&mut self, a: Id, b: Id) -> Id {
1727 let ty = self.u32_ty();
1728 self.value(OP_U_MOD, ty, &[a, b])
1729 }
1730 fn umin(&mut self, a: Id, b: Id) -> Id {
1731 let ty = self.u32_ty();
1732 let glsl = self.glsl;
1733 self.value(OP_EXT_INST, ty, &[glsl, GLSL_UMIN, a, b])
1734 }
1735 fn shl(&mut self, a: Id, shift: Id) -> Id {
1736 let ty = self.u32_ty();
1737 self.value(OP_SHIFT_LEFT_LOGICAL, ty, &[a, shift])
1738 }
1739 fn shr(&mut self, a: Id, shift: Id) -> Id {
1740 let ty = self.u32_ty();
1741 self.value(OP_SHIFT_RIGHT_LOGICAL, ty, &[a, shift])
1742 }
1743 fn bor(&mut self, a: Id, b: Id) -> Id {
1744 let ty = self.u32_ty();
1745 self.value(OP_BITWISE_OR, ty, &[a, b])
1746 }
1747 fn bxor(&mut self, a: Id, b: Id) -> Id {
1748 let ty = self.u32_ty();
1749 self.value(OP_BITWISE_XOR, ty, &[a, b])
1750 }
1751 fn band(&mut self, a: Id, b: Id) -> Id {
1752 let ty = self.u32_ty();
1753 self.value(OP_BITWISE_AND, ty, &[a, b])
1754 }
1755 fn bnot(&mut self, a: Id) -> Id {
1756 let ty = self.u32_ty();
1757 self.value(OP_NOT, ty, &[a])
1758 }
1759 fn ult(&mut self, a: Id, b: Id) -> Id {
1760 let ty = self.bool_ty();
1761 self.value(OP_U_LESS_THAN, ty, &[a, b])
1762 }
1763 fn uge(&mut self, a: Id, b: Id) -> Id {
1764 let ty = self.bool_ty();
1765 self.value(OP_U_GREATER_THAN_EQUAL, ty, &[a, b])
1766 }
1767 fn ieq(&mut self, a: Id, b: Id) -> Id {
1768 let ty = self.bool_ty();
1769 self.value(OP_I_EQUAL, ty, &[a, b])
1770 }
1771 fn ine(&mut self, a: Id, b: Id) -> Id {
1772 let ty = self.bool_ty();
1773 self.value(OP_I_NOT_EQUAL, ty, &[a, b])
1774 }
1775 fn land(&mut self, a: Id, b: Id) -> Id {
1776 let ty = self.bool_ty();
1777 self.value(OP_LOGICAL_AND, ty, &[a, b])
1778 }
1779 fn lor(&mut self, a: Id, b: Id) -> Id {
1780 let ty = self.bool_ty();
1781 self.value(OP_LOGICAL_OR, ty, &[a, b])
1782 }
1783 fn lxor(&mut self, a: Id, b: Id) -> Id {
1784 let ty = self.bool_ty();
1785 self.value(OP_LOGICAL_NOT_EQUAL, ty, &[a, b])
1786 }
1787 fn lnot(&mut self, a: Id) -> Id {
1788 let ty = self.bool_ty();
1789 self.value(OP_LOGICAL_NOT, ty, &[a])
1790 }
1791 fn select(&mut self, ty: Id, condition: Id, then: Id, otherwise: Id) -> Id {
1792 self.value(OP_SELECT, ty, &[condition, then, otherwise])
1793 }
1794 fn select_u32(&mut self, condition: Id, then: Id, otherwise: Id) -> Id {
1795 let ty = self.u32_ty();
1796 self.select(ty, condition, then, otherwise)
1797 }
1798 fn select_f32(&mut self, condition: Id, then: Id, otherwise: Id) -> Id {
1799 let ty = self.f32_ty();
1800 self.select(ty, condition, then, otherwise)
1801 }
1802 fn bitcast_f32(&mut self, word: Id) -> Id {
1803 let ty = self.f32_ty();
1804 self.value(OP_BITCAST, ty, &[word])
1805 }
1806 fn bitcast_u32(&mut self, float: Id) -> Id {
1807 let ty = self.u32_ty();
1808 self.value(OP_BITCAST, ty, &[float])
1809 }
1810 fn f32_to_f16(&mut self, float: Id) -> Id {
1811 let ty = self.f16_ty();
1812 self.value(OP_F_CONVERT, ty, &[float])
1813 }
1814
1815 fn cooperative_load(&mut self, ty: Id, pointer: Id, stride: Id) -> Id {
1816 let row_major = self.c_u32(0);
1817 self.value(
1818 OP_COOPERATIVE_MATRIX_LOAD_KHR,
1819 ty,
1820 &[pointer, row_major, stride],
1821 )
1822 }
1823
1824 fn cooperative_store(&mut self, pointer: Id, value: Id, stride: Id) {
1825 let row_major = self.c_u32(0);
1826 self.emit(
1827 OP_COOPERATIVE_MATRIX_STORE_KHR,
1828 &[pointer, value, row_major, stride],
1829 );
1830 }
1831
1832 fn cooperative_mul_add(&mut self, ty: Id, a: Id, b: Id, c: Id) -> Id {
1833 self.value(OP_COOPERATIVE_MATRIX_MUL_ADD_KHR, ty, &[a, b, c])
1834 }
1835
1836 fn subgroup_sum_f32(&mut self, value: Id) -> Id {
1837 let ty = self.f32_ty();
1838 let subgroup_scope = self.c_u32(3);
1839 self.value(
1840 OP_GROUP_NON_UNIFORM_F_ADD,
1841 ty,
1842 &[subgroup_scope, 0, value], )
1844 }
1845 fn u_to_f(&mut self, word: Id) -> Id {
1846 let ty = self.f32_ty();
1847 self.value(OP_CONVERT_U_TO_F, ty, &[word])
1848 }
1849 fn f_to_u(&mut self, float: Id) -> Id {
1850 let ty = self.u32_ty();
1851 self.value(OP_CONVERT_F_TO_U, ty, &[float])
1852 }
1853
1854 fn fadd(&mut self, a: Id, b: Id) -> Id {
1855 self.float_value(OP_F_ADD, &[a, b])
1856 }
1857 fn fsub(&mut self, a: Id, b: Id) -> Id {
1858 self.float_value(OP_F_SUB, &[a, b])
1859 }
1860 fn fmul(&mut self, a: Id, b: Id) -> Id {
1861 self.float_value(OP_F_MUL, &[a, b])
1862 }
1863 fn fdiv(&mut self, a: Id, b: Id) -> Id {
1864 self.float_value(OP_F_DIV, &[a, b])
1865 }
1866 fn fneg(&mut self, a: Id) -> Id {
1867 let ty = self.f32_ty();
1868 self.value(OP_F_NEGATE, ty, &[a])
1869 }
1870 fn fma_free(&mut self, a: Id, b: Id, c: Id) -> Id {
1873 let product = self.fmul(a, b);
1874 self.fadd(product, c)
1875 }
1876 fn fma(&mut self, a: Id, b: Id, c: Id) -> Id {
1880 let ty = self.f32_ty();
1881 let glsl = self.glsl;
1882 let id = self.value(OP_EXT_INST, ty, &[glsl, GLSL_FMA, a, b, c]);
1883 instruction(
1884 &mut self.annotations,
1885 OP_DECORATE,
1886 &[id, DECORATION_NO_CONTRACTION],
1887 );
1888 id
1889 }
1890
1891 fn widen_fp8(&mut self, format: Fp8Format, bits: Id) -> Id {
1909 let (mantissa_bits, k, top_exponent) = match format {
1911 Fp8Format::E4M3 => (3_u32, 120_u32, 15_u32),
1912 Fp8Format::E5M2 => (2, 112, 31),
1913 };
1914 let magnitude_mask = self.c_u32(0x7f);
1915 let shift = self.c_u32(23 - mantissa_bits);
1916 let k_bits = self.c_u32(k << 23);
1917 let mantissa_shift = self.c_u32(mantissa_bits);
1918 let zero = self.c_u32(0);
1919 let implicit_one = self.c_f32(f32::from_bits(k << 23));
1920 let sign_mask = self.c_u32(0x80);
1921 let twenty_four = self.c_u32(24);
1922 let magnitude = self.band(bits, magnitude_mask);
1923 let placed = self.shl(magnitude, shift);
1924 let x_bits = self.iadd(placed, k_bits);
1927 let x = self.bitcast_f32(x_bits);
1928 let exponent = self.shr(magnitude, mantissa_shift);
1929 let is_subnormal = self.ieq(exponent, zero);
1930 let difference = self.fsub(x, implicit_one);
1931 let subnormal = self.fadd(difference, difference);
1932 let value = self.select_f32(is_subnormal, subnormal, x);
1933 let value = self.bitcast_u32(value);
1934 let value = match format {
1937 Fp8Format::E4M3 => {
1938 let nan_pattern = self.c_u32(0x7f);
1939 let quiet_nan = self.c_u32(0x7fc0_0000);
1940 let is_nan = self.ieq(magnitude, nan_pattern);
1941 self.select_u32(is_nan, quiet_nan, value)
1942 }
1943 Fp8Format::E5M2 => {
1944 let top = self.c_u32(top_exponent);
1945 let inf_bits = self.c_u32(0x7f80_0000);
1946 let fraction_mask = self.c_u32((1 << mantissa_bits) - 1);
1947 let is_infnan = self.ieq(exponent, top);
1948 let fraction = self.band(magnitude, fraction_mask);
1949 let payload = self.shl(fraction, shift);
1950 let infnan = self.bor(inf_bits, payload);
1951 self.select_u32(is_infnan, infnan, value)
1952 }
1953 };
1954 let sign = self.band(bits, sign_mask);
1955 let sign = self.shl(sign, twenty_four);
1956 let value = self.bor(value, sign);
1957 self.bitcast_f32(value)
1958 }
1959
1960 fn widen_f16(&mut self, bits: Id) -> Id {
1966 let magnitude_mask = self.c_u32(0x7fff);
1967 let thirteen = self.c_u32(13);
1968 let ten = self.c_u32(10);
1969 let sixteen = self.c_u32(16);
1970 let k_bits = self.c_u32(112 << 23);
1971 let implicit_one = self.c_f32(f32::from_bits(112 << 23));
1972 let zero = self.c_u32(0);
1973 let thirty_one = self.c_u32(31);
1974 let inf_bits = self.c_u32(0x7f80_0000);
1975 let mantissa_mask = self.c_u32(0x3ff);
1976 let sign_mask = self.c_u32(0x8000);
1977 let magnitude = self.band(bits, magnitude_mask);
1978 let placed = self.shl(magnitude, thirteen);
1979 let x_bits = self.iadd(placed, k_bits);
1982 let x = self.bitcast_f32(x_bits);
1983 let exponent = self.shr(magnitude, ten);
1984 let is_subnormal = self.ieq(exponent, zero);
1985 let difference = self.fsub(x, implicit_one);
1986 let subnormal = self.fadd(difference, difference);
1987 let value = self.select_f32(is_subnormal, subnormal, x);
1988 let value = self.bitcast_u32(value);
1989 let is_infnan = self.ieq(exponent, thirty_one);
1991 let mantissa = self.band(magnitude, mantissa_mask);
1992 let payload = self.shl(mantissa, thirteen);
1993 let infnan = self.bor(inf_bits, payload);
1994 let value = self.select_u32(is_infnan, infnan, value);
1995 let sign = self.band(bits, sign_mask);
1996 let sign = self.shl(sign, sixteen);
1997 let value = self.bor(value, sign);
1998 self.bitcast_f32(value)
1999 }
2000
2001 fn narrow_fp8(&mut self, format: Fp8Format, value: Id) -> Id {
2010 let (mantissa_bits, bias, nan_out, max_finite, overflow_out) = match format {
2011 Fp8Format::E4M3 => (3_u32, 7_u32, 0x7f_u32, 0x7e_u32, 0x7f_u32),
2012 Fp8Format::E5M2 => (2, 15, 0x7e, 0x7b, 0x7c),
2013 };
2014 let shift = 23 - mantissa_bits;
2015 let twenty_four = self.c_u32(24);
2016 let one = self.c_u32(1);
2017 let sign_mask = self.c_u32(0x80);
2018 let magnitude_mask = self.c_u32(0x7fff_ffff);
2019 let inf_bits = self.c_u32(0x7f80_0000);
2020 let normal_floor = self.c_u32((128 - bias) << 23);
2021 let rebias = self.c_u32((127 - bias) << 23);
2022 let shift_c = self.c_u32(shift);
2023 let half_ulp = self.c_u32((1 << (shift - 1)) - 1);
2024 let scale = self.c_f32(f32::from_bits((127 + bias - 1 + mantissa_bits) << 23));
2025 let max_finite_c = self.c_u32(max_finite);
2026 let nan_c = self.c_u32(nan_out);
2027 let overflow_c = self.c_u32(overflow_out);
2028 let bits = self.bitcast_u32(value);
2029 let sign = self.shr(bits, twenty_four);
2030 let sign = self.band(sign, sign_mask);
2031 let magnitude = self.band(bits, magnitude_mask);
2032 let is_nan = self.ult(inf_bits, magnitude);
2033 let is_normal = self.uge(magnitude, normal_floor);
2034 let adjusted = self.isub(magnitude, rebias);
2037 let guard = self.shr(adjusted, shift_c);
2038 let guard = self.band(guard, one);
2039 let rounding = self.iadd(half_ulp, guard);
2040 let rounded = self.iadd(adjusted, rounding);
2041 let normal_out = self.shr(rounded, shift_c);
2042 let safe_magnitude = self.select_u32(is_normal, normal_floor, magnitude);
2045 let scaled = self.bitcast_f32(safe_magnitude);
2046 let scaled = self.fmul(scaled, scale);
2047 let scaled = self.ext_f32(GLSL_ROUND_EVEN, &[scaled]);
2048 let subnormal_out = self.f_to_u(scaled);
2049 let body = self.select_u32(is_normal, normal_out, subnormal_out);
2050 let overflowed = self.ult(max_finite_c, body);
2052 let body = self.select_u32(overflowed, overflow_c, body);
2053 let body = self.select_u32(is_nan, nan_c, body);
2054 self.bor(sign, body)
2055 }
2056
2057 fn narrow_f16(&mut self, value: Id) -> Id {
2058 let sixteen = self.c_u32(16);
2059 let thirteen = self.c_u32(13);
2060 let one = self.c_u32(1);
2061 let sign_mask = self.c_u32(0x8000);
2062 let magnitude_mask = self.c_u32(0x7fff_ffff);
2063 let inf_bits = self.c_u32(0x7f80_0000);
2064 let overflow_bits = self.c_u32(0x477f_f000);
2065 let normal_floor = self.c_u32(0x3880_0000);
2066 let rebias = self.c_u32(0x3800_0000);
2067 let half_ulp = self.c_u32(0x0000_0fff);
2068 let inf_out = self.c_u32(0x7c00);
2069 let nan_out = self.c_u32(0x7e00);
2070 let scale = self.c_f32(16_777_216.0);
2071 let bits = self.bitcast_u32(value);
2072 let sign = self.shr(bits, sixteen);
2073 let sign = self.band(sign, sign_mask);
2074 let magnitude = self.band(bits, magnitude_mask);
2075 let is_nan = self.ult(inf_bits, magnitude);
2076 let overflow = self.uge(magnitude, overflow_bits);
2079 let is_normal = self.uge(magnitude, normal_floor);
2082 let adjusted = self.isub(magnitude, rebias);
2083 let guard = self.shr(adjusted, thirteen);
2084 let guard = self.band(guard, one);
2085 let rounding = self.iadd(half_ulp, guard);
2086 let rounded = self.iadd(adjusted, rounding);
2087 let normal_out = self.shr(rounded, thirteen);
2088 let safe_magnitude = self.select_u32(is_normal, normal_floor, magnitude);
2093 let scaled = self.bitcast_f32(safe_magnitude);
2094 let scaled = self.fmul(scaled, scale);
2095 let scaled = self.ext_f32(GLSL_ROUND_EVEN, &[scaled]);
2096 let subnormal_out = self.f_to_u(scaled);
2097 let body = self.select_u32(is_normal, normal_out, subnormal_out);
2098 let body = self.select_u32(overflow, inf_out, body);
2099 let body = self.select_u32(is_nan, nan_out, body);
2100 self.bor(sign, body)
2101 }
2102 fn ext_f32(&mut self, op: u32, args: &[Id]) -> Id {
2103 let ty = self.f32_ty();
2104 let glsl = self.glsl;
2105 let mut operands = vec![glsl, op];
2106 operands.extend_from_slice(args);
2107 self.value(OP_EXT_INST, ty, &operands)
2108 }
2109 fn fabs(&mut self, a: Id) -> Id {
2110 self.ext_f32(GLSL_FABS, &[a])
2111 }
2112 fn is_nan(&mut self, a: Id) -> Id {
2113 let ty = self.bool_ty();
2114 self.value(OP_IS_NAN, ty, &[a])
2115 }
2116 fn foeq(&mut self, a: Id, b: Id) -> Id {
2117 let ty = self.bool_ty();
2118 self.value(OP_F_ORD_EQUAL, ty, &[a, b])
2119 }
2120 fn folt(&mut self, a: Id, b: Id) -> Id {
2121 let ty = self.bool_ty();
2122 self.value(OP_F_ORD_LESS_THAN, ty, &[a, b])
2123 }
2124 fn fogt(&mut self, a: Id, b: Id) -> Id {
2125 let ty = self.bool_ty();
2126 self.value(OP_F_ORD_GREATER_THAN, ty, &[a, b])
2127 }
2128 fn foge(&mut self, a: Id, b: Id) -> Id {
2129 let ty = self.bool_ty();
2130 self.value(OP_F_ORD_GREATER_THAN_EQUAL, ty, &[a, b])
2131 }
2132
2133 fn private_element(&mut self, array: Id, index: Id) -> Id {
2135 let u32_ty = self.u32_ty();
2136 let pointer = self.pointer(STORAGE_CLASS_PRIVATE, u32_ty);
2137 let chain = self.access_chain(pointer, array, &[index]);
2138 self.load(u32_ty, chain)
2139 }
2140
2141 fn find_msb(&mut self, value: Id) -> Id {
2143 let ty = self.u32_ty();
2144 let glsl = self.glsl;
2145 self.value(OP_EXT_INST, ty, &[glsl, GLSL_FIND_U_MSB, value])
2146 }
2147
2148 fn mul_wide(&mut self, a: Id, b: Id) -> (Id, Id) {
2151 let sixteen = self.c_u32(16);
2152 let mask = self.c_u32(0xffff);
2153 let a_lo = self.band(a, mask);
2154 let a_hi = self.shr(a, sixteen);
2155 let b_lo = self.band(b, mask);
2156 let b_hi = self.shr(b, sixteen);
2157 let ll = self.imul(a_lo, b_lo);
2158 let lh = self.imul(a_lo, b_hi);
2159 let hl = self.imul(a_hi, b_lo);
2160 let hh = self.imul(a_hi, b_hi);
2161 let ll_hi = self.shr(ll, sixteen);
2162 let lh_lo = self.band(lh, mask);
2163 let hl_lo = self.band(hl, mask);
2164 let middle = self.iadd(ll_hi, lh_lo);
2165 let middle = self.iadd(middle, hl_lo);
2166 let ll_lo = self.band(ll, mask);
2167 let middle_lo = self.band(middle, mask);
2168 let middle_shifted = self.shl(middle_lo, sixteen);
2169 let low = self.bor(ll_lo, middle_shifted);
2170 let lh_hi = self.shr(lh, sixteen);
2171 let hl_hi = self.shr(hl, sixteen);
2172 let middle_hi = self.shr(middle, sixteen);
2173 let high = self.iadd(hh, lh_hi);
2174 let high = self.iadd(high, hl_hi);
2175 let high = self.iadd(high, middle_hi);
2176 (high, low)
2177 }
2178
2179 fn add_carry(&mut self, a: Id, b: Id) -> (Id, Id) {
2181 let sum = self.iadd(a, b);
2182 let carried = self.ult(sum, a);
2183 let one = self.c_u32(1);
2184 let zero = self.c_u32(0);
2185 let carry = self.select_u32(carried, one, zero);
2186 (sum, carry)
2187 }
2188
2189 fn pow2(&mut self, exponent: Id) -> Id {
2191 let bias = self.c_u32(127);
2192 let biased = self.iadd(exponent, bias);
2193 let twenty_three = self.c_u32(23);
2194 let bits = self.shl(biased, twenty_three);
2195 self.bitcast_f32(bits)
2196 }
2197
2198 fn label(&mut self, id: Id) {
2199 self.emit(OP_LABEL, &[id]);
2200 }
2201 fn branch(&mut self, target: Id) {
2202 self.emit(OP_BRANCH, &[target]);
2203 }
2204 fn branch_conditional(&mut self, condition: Id, then: Id, otherwise: Id) {
2205 self.emit(OP_BRANCH_CONDITIONAL, &[condition, then, otherwise]);
2206 }
2207 fn workgroup_barrier(&mut self) {
2208 let scope = self.c_u32(SCOPE_WORKGROUP);
2209 let semantics = self.c_u32(MEMORY_SEMANTICS_ACQUIRE_RELEASE_WORKGROUP);
2210 self.emit(OP_CONTROL_BARRIER, &[scope, scope, semantics]);
2211 }
2212
2213 fn begin_loop(&mut self, counter: Id, limit: Id) -> (LoopScope, Id) {
2218 let scope = LoopScope {
2219 header: self.id(),
2220 body: self.id(),
2221 cont: self.id(),
2222 merge: self.id(),
2223 };
2224 self.branch(scope.header);
2225 self.label(scope.header);
2226 let u32_ty = self.u32_ty();
2227 let current = self.load(u32_ty, counter);
2228 let in_range = self.ult(current, limit);
2229 self.emit(OP_LOOP_MERGE, &[scope.merge, scope.cont, LOOP_CONTROL_NONE]);
2230 self.branch_conditional(in_range, scope.body, scope.merge);
2231 self.label(scope.body);
2232 (scope, current)
2233 }
2234
2235 fn end_loop(&mut self, scope: LoopScope, counter: Id, step: Id) {
2236 self.branch(scope.cont);
2237 self.label(scope.cont);
2238 let u32_ty = self.u32_ty();
2239 let current = self.load(u32_ty, counter);
2240 let next = self.iadd(current, step);
2241 self.store(counter, next);
2242 self.branch(scope.header);
2243 self.label(scope.merge);
2244 }
2245
2246 fn if_then(&mut self, condition: Id, then: impl FnOnce(&mut Self)) {
2248 let then_label = self.id();
2249 let merge = self.id();
2250 self.emit(OP_SELECTION_MERGE, &[merge, SELECTION_CONTROL_NONE]);
2251 self.branch_conditional(condition, then_label, merge);
2252 self.label(then_label);
2253 then(self);
2254 self.branch(merge);
2255 self.label(merge);
2256 }
2257
2258 fn word_pointer(&mut self, buffers: Id, operand: (Id, Id), word_index: Id) -> Id {
2262 let u32_ty = self.u32_ty();
2263 let pointer = self.pointer(STORAGE_CLASS_STORAGE_BUFFER, u32_ty);
2264 let zero = self.c_u32(0);
2265 let index = self.iadd(operand.1, word_index);
2266 self.access_chain(pointer, buffers, &[operand.0, zero, index])
2267 }
2268
2269 fn load_word(&mut self, buffers: Id, operand: (Id, Id), word_index: Id) -> Id {
2270 let pointer = self.word_pointer(buffers, operand, word_index);
2271 let u32_ty = self.u32_ty();
2272 self.load(u32_ty, pointer)
2273 }
2274
2275 fn load_f32(&mut self, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2276 let word = self.load_word(buffers, operand, element);
2277 self.bitcast_f32(word)
2278 }
2279
2280 fn store_word(&mut self, buffers: Id, operand: (Id, Id), word_index: Id, value: Id) {
2281 let pointer = self.word_pointer(buffers, operand, word_index);
2282 self.store(pointer, value);
2283 }
2284
2285 fn store_f32(&mut self, buffers: Id, operand: (Id, Id), element: Id, value: Id) {
2286 let word = self.bitcast_u32(value);
2287 self.store_word(buffers, operand, element, word);
2288 }
2289
2290 fn load_byte_bits(&mut self, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2293 let two = self.c_u32(2);
2294 let three = self.c_u32(3);
2295 let mask = self.c_u32(0xff);
2296 let word_index = self.shr(element, two);
2297 let word = self.load_word(buffers, operand, word_index);
2298 let lane = self.band(element, three);
2299 let eight = self.c_u32(8);
2300 let shift = self.imul(lane, eight);
2301 let shifted = self.shr(word, shift);
2302 self.band(shifted, mask)
2303 }
2304
2305 fn load_bool(&mut self, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2307 let zero = self.c_u32(0);
2308 let byte = self.load_byte_bits(buffers, operand, element);
2309 self.ine(byte, zero)
2310 }
2311
2312 fn store_byte_bits(&mut self, buffers: Id, operand: (Id, Id), element: Id, byte: Id) {
2316 let two = self.c_u32(2);
2317 let three = self.c_u32(3);
2318 let mask = self.c_u32(0xff);
2319 let word_index = self.shr(element, two);
2320 let pointer = self.word_pointer(buffers, operand, word_index);
2321 let lane = self.band(element, three);
2322 let eight = self.c_u32(8);
2323 let shift = self.imul(lane, eight);
2324 let clear = self.shl(mask, shift);
2325 let clear = self.bnot(clear);
2326 let masked = self.band(byte, mask);
2327 let set = self.shl(masked, shift);
2328 let scope = self.c_u32(SCOPE_DEVICE);
2329 let semantics = self.c_u32(MEMORY_SEMANTICS_RELAXED);
2330 let u32_ty = self.u32_ty();
2331 self.value(OP_ATOMIC_AND, u32_ty, &[pointer, scope, semantics, clear]);
2332 self.value(OP_ATOMIC_OR, u32_ty, &[pointer, scope, semantics, set]);
2333 }
2334
2335 fn store_bool(&mut self, buffers: Id, operand: (Id, Id), element: Id, value: Id) {
2337 let one = self.c_u32(1);
2338 let zero = self.c_u32(0);
2339 let byte = self.select_u32(value, one, zero);
2340 self.store_byte_bits(buffers, operand, element, byte);
2341 }
2342
2343 fn load_half_bits(&mut self, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2347 let one = self.c_u32(1);
2348 let sixteen = self.c_u32(16);
2349 let mask = self.c_u32(0xffff);
2350 let word_index = self.shr(element, one);
2351 let word = self.load_word(buffers, operand, word_index);
2352 let lane = self.band(element, one);
2353 let shift = self.imul(lane, sixteen);
2354 let shifted = self.shr(word, shift);
2355 self.band(shifted, mask)
2356 }
2357
2358 fn store_half_bits(&mut self, buffers: Id, operand: (Id, Id), element: Id, bits: Id) {
2362 let one = self.c_u32(1);
2363 let sixteen = self.c_u32(16);
2364 let mask = self.c_u32(0xffff);
2365 let word_index = self.shr(element, one);
2366 let pointer = self.word_pointer(buffers, operand, word_index);
2367 let lane = self.band(element, one);
2368 let shift = self.imul(lane, sixteen);
2369 let clear = self.shl(mask, shift);
2370 let clear = self.bnot(clear);
2371 let bits = self.band(bits, mask);
2372 let set = self.shl(bits, shift);
2373 let scope = self.c_u32(SCOPE_DEVICE);
2374 let semantics = self.c_u32(MEMORY_SEMANTICS_RELAXED);
2375 let u32_ty = self.u32_ty();
2376 self.value(OP_ATOMIC_AND, u32_ty, &[pointer, scope, semantics, clear]);
2377 self.value(OP_ATOMIC_OR, u32_ty, &[pointer, scope, semantics, set]);
2378 }
2379
2380 fn store_float(
2382 &mut self,
2383 storage: Storage,
2384 buffers: Id,
2385 operand: (Id, Id),
2386 element: Id,
2387 value: Id,
2388 ) {
2389 match storage {
2390 Storage::Word => self.store_f32(buffers, operand, element, value),
2391 Storage::Half => {
2392 let bits = self.narrow_f16(value);
2393 self.store_half_bits(buffers, operand, element, bits);
2394 }
2395 Storage::Quarter(format) => {
2396 let bits = self.narrow_fp8(format, value);
2397 self.store_byte_bits(buffers, operand, element, bits);
2398 }
2399 Storage::Byte => unreachable!("BOOL is not a float storage"),
2400 }
2401 }
2402
2403 fn load_float(&mut self, storage: Storage, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2406 match storage {
2407 Storage::Word => self.load_f32(buffers, operand, element),
2408 Storage::Quarter(format) => {
2409 let bits = self.load_byte_bits(buffers, operand, element);
2410 self.widen_fp8(format, bits)
2411 }
2412 Storage::Half => {
2413 let bits = self.load_half_bits(buffers, operand, element);
2414 self.widen_f16(bits)
2415 }
2416 Storage::Byte => unreachable!("byte storage is not a float lane"),
2417 }
2418 }
2419
2420 fn load_lane_group(
2429 &mut self,
2430 storage: Storage,
2431 buffers: Id,
2432 operand: (Id, Id),
2433 base: Id,
2434 lanes: u32,
2435 ) -> Vec<Id> {
2436 let source_lanes = storage.lanes();
2437 let lane_bits = 32 / source_lanes;
2438 let first = match source_lanes.trailing_zeros() {
2439 0 => base,
2440 shift => {
2441 let shift = self.c_u32(shift);
2442 self.shr(base, shift)
2443 }
2444 };
2445 let words: Vec<Id> = (0..lanes.div_ceil(source_lanes))
2446 .map(|word| {
2447 let word = self.c_u32(word);
2448 let index = self.iadd(first, word);
2449 self.load_word(buffers, operand, index)
2450 })
2451 .collect();
2452 if lane_bits == 32 {
2453 return words;
2454 }
2455 let mask = self.c_u32((1 << lane_bits) - 1);
2456 let offset = (source_lanes > lanes).then(|| {
2457 let modulus = self.c_u32(source_lanes - 1);
2458 self.band(base, modulus)
2459 });
2460 (0..lanes)
2461 .map(|lane| {
2462 let word = words[(lane / source_lanes) as usize];
2463 let shift = match offset {
2464 None => self.c_u32((lane % source_lanes) * lane_bits),
2465 Some(offset) => {
2466 let lane = self.c_u32(lane % source_lanes);
2467 let position = self.iadd(offset, lane);
2468 let bits = self.c_u32(lane_bits);
2469 self.imul(position, bits)
2470 }
2471 };
2472 let shifted = self.shr(word, shift);
2473 self.band(shifted, mask)
2474 })
2475 .collect()
2476 }
2477
2478 fn pack_lanes(&mut self, values: &[Id]) -> Id {
2481 let lane_bits = 32 / values.len() as u32;
2482 let mut word = values[0];
2483 for (lane, value) in values.iter().enumerate().skip(1) {
2484 let shift = self.c_u32(lane as u32 * lane_bits);
2485 let shifted = self.shl(*value, shift);
2486 word = self.bor(word, shifted);
2487 }
2488 word
2489 }
2490
2491 fn store_lane_bits(
2493 &mut self,
2494 storage: Storage,
2495 buffers: Id,
2496 operand: (Id, Id),
2497 element: Id,
2498 bits: Id,
2499 ) {
2500 match storage {
2501 Storage::Word => self.store_word(buffers, operand, element, bits),
2502 Storage::Half => self.store_half_bits(buffers, operand, element, bits),
2503 Storage::Byte | Storage::Quarter(_) => {
2504 self.store_byte_bits(buffers, operand, element, bits)
2505 }
2506 }
2507 }
2508
2509 fn widen_bits(&mut self, storage: Storage, bits: Id) -> Id {
2511 match storage {
2512 Storage::Word => self.bitcast_f32(bits),
2513 Storage::Half => self.widen_f16(bits),
2514 Storage::Quarter(format) => self.widen_fp8(format, bits),
2515 Storage::Byte => unreachable!("byte storage is not a float lane"),
2516 }
2517 }
2518
2519 fn narrow_bits(&mut self, storage: Storage, value: Id) -> Id {
2521 match storage {
2522 Storage::Word => self.bitcast_u32(value),
2523 Storage::Half => self.narrow_f16(value),
2524 Storage::Quarter(format) => self.narrow_fp8(format, value),
2525 Storage::Byte => unreachable!("BOOL is not a float storage"),
2526 }
2527 }
2528
2529 fn strided_indices(
2532 &mut self,
2533 i: Id,
2534 dims: &[Id; MAX_RANK],
2535 strides: &[[Id; MAX_RANK]],
2536 ) -> Vec<Id> {
2537 let zero = self.c_u32(0);
2538 let mut indices = vec![zero; strides.len()];
2539 let mut remainder = i;
2540 for d in (0..MAX_RANK).rev() {
2541 let coordinate = self.umod(remainder, dims[d]);
2542 remainder = self.udiv(remainder, dims[d]);
2543 for (operand, stride) in strides.iter().enumerate() {
2544 let term = self.imul(coordinate, stride[d]);
2545 indices[operand] = self.iadd(indices[operand], term);
2546 }
2547 }
2548 indices
2549 }
2550
2551 fn grid_stride(&mut self, workgroup: u32) -> (Id, Id) {
2554 let gid = self.builtin_uvec3(BUILT_IN_GLOBAL_INVOCATION_ID);
2555 let num_workgroups = self.builtin_uvec3(BUILT_IN_NUM_WORKGROUPS);
2556 self.begin_main();
2557 let u32_ty = self.u32_ty();
2558 let counter = self.local(u32_ty);
2559 let start = self.builtin_component(gid, 0);
2560 self.store(counter, start);
2561 let groups = self.builtin_component(num_workgroups, 0);
2562 let size = self.c_u32(workgroup);
2563 let stride = self.imul(groups, size);
2564 (counter, stride)
2565 }
2566
2567 fn apply_max(&mut self, a: Id, b: Id, nan_mode: NanMode) -> Id {
2571 let ordered = self.foge(a, b);
2572 let picked = self.select_f32(ordered, a, b);
2573 self.apply_nan_mode(a, b, picked, nan_mode)
2574 }
2575
2576 fn apply_min(&mut self, a: Id, b: Id, nan_mode: NanMode) -> Id {
2578 let ordered = self.folt(a, b);
2579 let picked = self.select_f32(ordered, a, b);
2580 self.apply_nan_mode(a, b, picked, nan_mode)
2581 }
2582
2583 fn apply_nan_mode(&mut self, a: Id, b: Id, picked: Id, nan_mode: NanMode) -> Id {
2584 let a_nan = self.is_nan(a);
2585 let b_nan = self.is_nan(b);
2586 match nan_mode {
2587 NanMode::Propagate => {
2588 let nan = self.c_f32(f32::NAN);
2589 let any_nan = self.lor(a_nan, b_nan);
2590 self.select_f32(any_nan, nan, picked)
2591 }
2592 NanMode::Ignore => {
2593 let without_b = self.select_f32(b_nan, a, picked);
2594 self.select_f32(a_nan, b, without_b)
2595 }
2596 }
2597 }
2598
2599 fn horner(&mut self, x: Id, coefficients: &[f32]) -> Id {
2601 let mut acc = self.c_f32(coefficients[0]);
2602 for coefficient in &coefficients[1..] {
2603 let c = self.c_f32(*coefficient);
2604 acc = self.fma_free(acc, x, c);
2605 }
2606 acc
2607 }
2608
2609 fn cephes_reduce(&mut self, magnitude: Id) -> (Id, Id) {
2613 let four_over_pi = self.c_f32(1.273_239_5);
2614 let scaled = self.fmul(magnitude, four_over_pi);
2615 let octant = self.f_to_u(scaled);
2616 let one = self.c_u32(1);
2617 let zero = self.c_u32(0);
2618 let odd = self.band(octant, one);
2619 let is_odd = self.ine(odd, zero);
2620 let bumped = self.iadd(octant, one);
2621 let octant = self.select_u32(is_odd, bumped, octant);
2622 let y = self.u_to_f(octant);
2623 let dp1 = self.c_f32(0.785_156_25);
2625 let dp2 = self.c_f32(2.418_756_5e-4);
2626 let dp3 = self.c_f32(3.774_895e-8);
2627 let t1 = self.fmul(y, dp1);
2628 let r = self.fsub(magnitude, t1);
2629 let t2 = self.fmul(y, dp2);
2630 let r = self.fsub(r, t2);
2631 let t3 = self.fmul(y, dp3);
2632 let z = self.fsub(r, t3);
2633 let seven = self.c_u32(7);
2634 let octant = self.band(octant, seven);
2635 (octant, z)
2636 }
2637
2638 fn payne_hanek_reduce(&mut self, magnitude: Id) -> (Id, Id) {
2650 let table = self.private_u32_array(TWO_OVER_PI_BITS);
2651 let bits = self.bitcast_u32(magnitude);
2652 let twenty_three = self.c_u32(23);
2653 let exponent_mask = self.c_u32(0xff);
2654 let exponent = self.shr(bits, twenty_three);
2655 let exponent = self.band(exponent, exponent_mask);
2656 let significand_mask = self.c_u32(0x007f_ffff);
2657 let implicit = self.c_u32(0x0080_0000);
2658 let significand = self.band(bits, significand_mask);
2659 let m = self.bor(significand, implicit);
2660
2661 let bias = self.c_u32(120);
2664 let offset = self.isub(exponent, bias);
2665 let five = self.c_u32(5);
2666 let thirty_one = self.c_u32(31);
2667 let base = self.shr(offset, five);
2668 let shift = self.band(offset, thirty_one);
2669 let thirty_two = self.c_u32(32);
2670 let complement = self.isub(thirty_two, shift);
2671 let complement = self.band(complement, thirty_one);
2672 let zero = self.c_u32(0);
2673 let aligned = self.ieq(shift, zero);
2674 let mut window = [0; 4];
2675 for (index, slot) in window.iter_mut().enumerate() {
2676 let step = self.c_u32(index as u32);
2677 let first = self.iadd(base, step);
2678 let one = self.c_u32(1);
2679 let second = self.iadd(first, one);
2680 let high = self.private_element(table, first);
2681 let low = self.private_element(table, second);
2682 let high = self.shl(high, shift);
2683 let low = self.shr(low, complement);
2684 let low = self.select_u32(aligned, zero, low);
2685 *slot = self.bor(high, low);
2686 }
2687
2688 let mut product = [0; 5];
2690 let mut carry = zero;
2691 for (index, word) in window.iter().rev().enumerate() {
2692 let (high, low) = self.mul_wide(m, *word);
2693 let (sum, overflow) = self.add_carry(low, carry);
2694 product[index] = sum;
2695 carry = self.iadd(high, overflow);
2696 }
2697 product[4] = carry;
2698
2699 let twenty_nine = self.c_u32(29);
2702 let three = self.c_u32(3);
2703 let seven = self.c_u32(7);
2704 let integer_low = self.shr(product[3], twenty_nine);
2705 let integer_high = self.shl(product[4], three);
2706 let integer = self.bor(integer_low, integer_high);
2707 let integer = self.band(integer, seven);
2708 let fraction_high = self.shl(product[3], three);
2709 let carry_in = self.shr(product[2], twenty_nine);
2710 let fraction_high = self.bor(fraction_high, carry_in);
2711 let fraction_low = self.shl(product[2], three);
2712 let carry_in = self.shr(product[1], twenty_nine);
2713 let fraction_low = self.bor(fraction_low, carry_in);
2714
2715 let one = self.c_u32(1);
2717 let odd = self.band(integer, one);
2718 let is_odd = self.ine(odd, zero);
2719 let bumped = self.iadd(integer, one);
2720 let octant = self.select_u32(is_odd, bumped, integer);
2721 let octant = self.band(octant, seven);
2722 let negated_low = self.bnot(fraction_low);
2724 let (negated_low, overflow) = self.add_carry(negated_low, one);
2725 let negated_high = self.bnot(fraction_high);
2726 let negated_high = self.iadd(negated_high, overflow);
2727 let magnitude_high = self.select_u32(is_odd, negated_high, fraction_high);
2728 let magnitude_low = self.select_u32(is_odd, negated_low, fraction_low);
2729
2730 let nonzero_high = self.ine(magnitude_high, zero);
2732 let msb_high = self.find_msb(magnitude_high);
2733 let msb_low = self.find_msb(magnitude_low);
2734 let shift_high = self.isub(thirty_one, msb_high);
2735 let shift_low = self.isub(thirty_one, msb_low);
2736 let shift_low_total = self.iadd(shift_low, thirty_two);
2737 let leading = self.select_u32(nonzero_high, shift_high, shift_low_total);
2738 let complement = self.isub(thirty_two, leading);
2741 let complement_masked = self.band(complement, thirty_one);
2742 let aligned = self.ieq(leading, zero);
2743 let top_high = self.shl(magnitude_high, leading);
2744 let carried = self.shr(magnitude_low, complement_masked);
2745 let carried = self.select_u32(aligned, zero, carried);
2746 let top_high = self.bor(top_high, carried);
2747 let top_low = self.shl(magnitude_low, shift_low);
2748 let top = self.select_u32(nonzero_high, top_high, top_low);
2749 let eight = self.c_u32(8);
2750 let significand = self.shr(top, eight);
2751 let value = self.u_to_f(significand);
2752 let twenty_four = self.c_u32(24);
2754 let exponent = self.iadd(twenty_four, leading);
2755 let exponent = self.isub(zero, exponent);
2756 let scale = self.pow2(exponent);
2757 let fraction = self.fmul(value, scale);
2758 let empty_low = self.ieq(magnitude_low, zero);
2760 let empty_high = self.ieq(magnitude_high, zero);
2761 let empty = self.land(empty_low, empty_high);
2762 let zero_f = self.c_f32(0.0);
2763 let fraction = self.select_f32(empty, zero_f, fraction);
2764 let negated = self.fneg(fraction);
2765 let fraction = self.select_f32(is_odd, negated, fraction);
2766 let pi_over_four_high = self.c_f32(0.785_156_25);
2768 let pi_over_four_low = self.c_f32(2.419_134e-4);
2769 let high = self.fmul(fraction, pi_over_four_high);
2770 let low = self.fmul(fraction, pi_over_four_low);
2771 let z = self.fadd(high, low);
2772 (octant, z)
2773 }
2774
2775 fn sincos(&mut self, x: Id, cosine: bool) -> Id {
2781 let magnitude = self.fabs(x);
2782 let (fast_octant, fast_z) = self.cephes_reduce(magnitude);
2783 let (exact_octant, exact_z) = self.payne_hanek_reduce(magnitude);
2784 let threshold = self.c_f32(SINCOS_FAST_RANGE);
2785 let fast = self.folt(magnitude, threshold);
2786 let octant = self.select_u32(fast, fast_octant, exact_octant);
2787 let z = self.select_f32(fast, fast_z, exact_z);
2788
2789 let three = self.c_u32(3);
2790 let reflect = self.ult(three, octant);
2791 let four = self.c_u32(4);
2792 let reduced = self.isub(octant, four);
2793 let octant = self.select_u32(reflect, reduced, octant);
2794 let zz = self.fmul(z, z);
2795 let cos_poly = self.horner(zz, &[2.443_315_7e-5, -1.388_731_6e-3, 4.166_664_6e-2]);
2797 let zz2 = self.fmul(zz, zz);
2798 let cos_tail = self.fmul(cos_poly, zz2);
2799 let half = self.c_f32(0.5);
2800 let half_zz = self.fmul(half, zz);
2801 let cos_value = self.fsub(cos_tail, half_zz);
2802 let one_f = self.c_f32(1.0);
2803 let cos_value = self.fadd(cos_value, one_f);
2804 let sin_poly = self.horner(zz, &[-1.951_529_6e-4, 8.332_161e-3, -1.666_665_5e-1]);
2806 let sin_tail = self.fmul(sin_poly, zz);
2807 let sin_tail = self.fmul(sin_tail, z);
2808 let sin_value = self.fadd(sin_tail, z);
2809 let one = self.c_u32(1);
2810 let two = self.c_u32(2);
2811 let octant_is_1 = self.ieq(octant, one);
2812 let octant_is_2 = self.ieq(octant, two);
2813 let middle = self.lor(octant_is_1, octant_is_2);
2814 let value = if cosine {
2816 self.select_f32(middle, sin_value, cos_value)
2817 } else {
2818 self.select_f32(middle, cos_value, sin_value)
2819 };
2820 let mut negate = reflect;
2821 if cosine {
2822 let upper_half = self.ult(one, octant);
2823 negate = self.lxor(negate, upper_half);
2824 } else {
2825 let negative_input = self.signbit(x);
2827 negate = self.lxor(negate, negative_input);
2828 }
2829 let negated = self.fneg(value);
2830 let value = self.select_f32(negate, negated, value);
2831 let infinity = self.c_f32(f32::INFINITY);
2833 let finite = self.folt(magnitude, infinity);
2834 let nan = self.c_f32(f32::NAN);
2835 self.select_f32(finite, value, nan)
2836 }
2837
2838 fn tanh(&mut self, x: Id) -> Id {
2840 let magnitude = self.fabs(x);
2841 let square = self.fmul(x, x);
2842 let poly = self.horner(
2843 square,
2844 &[
2845 -5.704_988_7e-3,
2846 2.063_909e-2,
2847 -5.373_971_6e-2,
2848 1.333_144_2e-1,
2849 -3.333_328e-1,
2850 ],
2851 );
2852 let small = self.fmul(poly, square);
2853 let small = self.fma_free(small, x, x);
2854 let two = self.c_f32(2.0);
2855 let doubled = self.fmul(magnitude, two);
2856 let exp = self.ext_f32(GLSL_EXP, &[doubled]);
2857 let one = self.c_f32(1.0);
2858 let denominator = self.fadd(exp, one);
2859 let ratio = self.fdiv(two, denominator);
2860 let large = self.fsub(one, ratio);
2861 let negative = self.signbit(x);
2862 let negated = self.fneg(large);
2863 let large = self.select_f32(negative, negated, large);
2864 let threshold = self.c_f32(0.625);
2865 let use_large = self.foge(magnitude, threshold);
2866 let value = self.select_f32(use_large, large, small);
2867 let zero = self.c_f32(0.0);
2869 let is_zero = self.foeq(x, zero);
2870 self.select_f32(is_zero, x, value)
2871 }
2872
2873 fn erf(&mut self, x: Id) -> Id {
2877 let magnitude = self.fabs(x);
2878 let square = self.fmul(x, x);
2879 let series = self.horner(
2880 square,
2881 &[
2882 1.0 / 76_204_800.0,
2883 -1.0 / 6_894_720.0,
2884 1.0 / 685_440.0,
2885 -1.0 / 75_600.0,
2886 1.0 / 9_360.0,
2887 -1.0 / 1_320.0,
2888 1.0 / 216.0,
2889 -1.0 / 42.0,
2890 1.0 / 10.0,
2891 -1.0 / 3.0,
2892 1.0,
2893 ],
2894 );
2895 let two_over_sqrt_pi = self.c_f32(core::f32::consts::FRAC_2_SQRT_PI);
2896 let series = self.fmul(series, two_over_sqrt_pi);
2897 let series = self.fmul(series, x);
2898 let half = self.c_f32(0.5);
2899 let one = self.c_f32(1.0);
2900 let half_magnitude = self.fmul(magnitude, half);
2901 let denominator = self.fadd(one, half_magnitude);
2902 let t = self.fdiv(one, denominator);
2903 let poly = self.horner(
2904 t,
2905 &[
2906 0.170_872_77,
2907 -0.822_152_23,
2908 1.488_515_9,
2909 -1.135_204,
2910 0.278_868_07,
2911 -0.186_288_06,
2912 0.096_784_18,
2913 0.374_091_96,
2914 1.000_023_7,
2915 -1.265_512_2,
2916 ],
2917 );
2918 let exponent = self.fsub(poly, square);
2919 let exp = self.ext_f32(GLSL_EXP, &[exponent]);
2920 let erfc = self.fmul(t, exp);
2921 let tail = self.fsub(one, erfc);
2922 let zero = self.c_f32(0.0);
2923 let negative = self.folt(x, zero);
2924 let negated = self.fneg(tail);
2925 let tail = self.select_f32(negative, negated, tail);
2926 let use_series = self.folt(magnitude, one);
2927 let value = self.select_f32(use_series, series, tail);
2928 let is_zero = self.foeq(x, zero);
2930 self.select_f32(is_zero, x, value)
2931 }
2932
2933 fn pow(&mut self, x: Id, y: Id) -> Id {
2937 let magnitude = self.fabs(x);
2938 let raw = self.ext_f32(GLSL_POW, &[magnitude, y]);
2939 let zero = self.c_f32(0.0);
2940 let one = self.c_f32(1.0);
2941 let half = self.c_f32(0.5);
2942 let floor_y = self.ext_f32(GLSL_FLOOR, &[y]);
2943 let y_integral = self.foeq(floor_y, y);
2944 let half_y = self.fmul(y, half);
2945 let floor_half = self.ext_f32(GLSL_FLOOR, &[half_y]);
2946 let y_even = self.foeq(floor_half, half_y);
2947 let y_odd = self.lnot(y_even);
2948 let y_odd = self.land(y_integral, y_odd);
2949 let negated = self.fneg(raw);
2950 let signed = self.select_f32(y_odd, negated, raw);
2951 let nan = self.c_f32(f32::NAN);
2952 let negative_base = self.select_f32(y_integral, signed, nan);
2953 let x_negative = self.folt(x, zero);
2954 let x_zero = self.foeq(x, zero);
2956 let y_positive = self.fogt(y, zero);
2957 let inf = self.c_f32(f32::INFINITY);
2958 let zero_base = self.select_f32(y_positive, zero, inf);
2959 let zero_base_neg = self.fneg(zero_base);
2960 let x_sign = self.bitcast_f32_sign(x);
2961 let x_sign_negative = self.folt(x_sign, zero);
2962 let zero_base_signed = self.land(x_sign_negative, y_odd);
2963 let zero_base = self.select_f32(zero_base_signed, zero_base_neg, zero_base);
2964 let value = self.select_f32(x_negative, negative_base, raw);
2965 let value = self.select_f32(x_zero, zero_base, value);
2966 let y_zero = self.foeq(y, zero);
2967 let value = self.select_f32(y_zero, one, value);
2968 let x_one = self.foeq(x, one);
2969 self.select_f32(x_one, one, value)
2970 }
2971
2972 fn signbit(&mut self, x: Id) -> Id {
2974 let bits = self.bitcast_u32(x);
2975 let sign_bit = self.c_u32(0x8000_0000);
2976 let masked = self.band(bits, sign_bit);
2977 let zero = self.c_u32(0);
2978 self.ine(masked, zero)
2979 }
2980
2981 fn bitcast_f32_sign(&mut self, x: Id) -> Id {
2983 let bits = self.bitcast_u32(x);
2984 let sign_bit = self.c_u32(0x8000_0000);
2985 let masked = self.band(bits, sign_bit);
2986 let zero = self.c_u32(0);
2987 let negative = self.ine(masked, zero);
2988 let minus_one = self.c_f32(-1.0);
2989 let plus_one = self.c_f32(1.0);
2990 self.select_f32(negative, minus_one, plus_one)
2991 }
2992
2993 fn elementwise_lane(
2998 &mut self,
2999 op: ElementwiseOp,
3000 inputs: &[Id],
3001 clamp: Option<(Id, Id)>,
3002 ) -> Id {
3003 let x = inputs[0];
3004 match op {
3005 ElementwiseOp::Abs => self.fabs(x),
3006 ElementwiseOp::Ceil => self.ext_f32(GLSL_CEIL, &[x]),
3007 ElementwiseOp::Floor => self.ext_f32(GLSL_FLOOR, &[x]),
3008 ElementwiseOp::Cos => self.sincos(x, true),
3009 ElementwiseOp::Sin => self.sincos(x, false),
3010 ElementwiseOp::Erf => self.erf(x),
3011 ElementwiseOp::Exp => self.ext_f32(GLSL_EXP, &[x]),
3012 ElementwiseOp::Log => self.ext_f32(GLSL_LOG, &[x]),
3013 ElementwiseOp::Negate => self.fneg(x),
3014 ElementwiseOp::Reciprocal => {
3015 let one = self.c_f32(1.0);
3016 self.fdiv(one, x)
3017 }
3018 ElementwiseOp::Rsqrt => self.ext_f32(GLSL_INVERSE_SQRT, &[x]),
3019 ElementwiseOp::Sigmoid => {
3020 let negated = self.fneg(x);
3021 let exp = self.ext_f32(GLSL_EXP, &[negated]);
3022 let one = self.c_f32(1.0);
3023 let denominator = self.fadd(one, exp);
3024 self.fdiv(one, denominator)
3025 }
3026 ElementwiseOp::Tanh => self.tanh(x),
3027 ElementwiseOp::Clamp(nan_mode) => {
3028 let (lo_bits, hi_bits) = clamp.expect("clamp bounds");
3031 let lo = self.bitcast_f32(lo_bits);
3032 let hi = self.bitcast_f32(hi_bits);
3033 let floored = self.apply_max(x, lo, nan_mode);
3034 self.apply_min(floored, hi, nan_mode)
3035 }
3036 ElementwiseOp::Add => self.fadd(x, inputs[1]),
3037 ElementwiseOp::Sub => self.fsub(x, inputs[1]),
3038 ElementwiseOp::Mul => self.fmul(x, inputs[1]),
3039 ElementwiseOp::Pow => self.pow(x, inputs[1]),
3040 ElementwiseOp::Maximum(nan_mode) => self.apply_max(x, inputs[1], nan_mode),
3041 ElementwiseOp::Minimum(nan_mode) => self.apply_min(x, inputs[1], nan_mode),
3042 ElementwiseOp::Equal => self.foeq(x, inputs[1]),
3043 ElementwiseOp::Greater => self.fogt(x, inputs[1]),
3044 ElementwiseOp::GreaterEqual => self.foge(x, inputs[1]),
3045 ElementwiseOp::LogicalAnd => self.land(x, inputs[1]),
3046 ElementwiseOp::LogicalOr => self.lor(x, inputs[1]),
3047 ElementwiseOp::LogicalXor => self.lxor(x, inputs[1]),
3048 ElementwiseOp::LogicalNot => self.lnot(x),
3049 ElementwiseOp::Select => self.select_f32(x, inputs[1], inputs[2]),
3050 ElementwiseOp::CopyBytes => x,
3051 }
3052 }
3053}
3054
3055#[derive(Clone, Copy)]
3056struct LoopScope {
3057 header: Id,
3058 body: Id,
3059 cont: Id,
3060 merge: Id,
3061}
3062
3063fn assemble_elementwise(
3075 op: ElementwiseOp,
3076 float: Storage,
3077 broadcast: bool,
3078 workgroup: u32,
3079 buffers: u32,
3080) -> Vec<u32> {
3081 let mut b = Builder::new();
3082 let array = b.buffer_array(buffers);
3083 let count = b.spec_u32(1);
3084 let input_storage = op.inputs();
3085 let lane = |storage: Storage| match storage {
3088 Storage::Word => float,
3089 other => other,
3090 };
3091 let inputs: Vec<(Id, Id)> = input_storage.iter().map(|_| b.spec_operand()).collect();
3092 let output = b.spec_operand();
3093 let shape = broadcast.then(|| {
3094 let dims = b.spec_dims();
3095 let strides: Vec<[Id; MAX_RANK]> = input_storage.iter().map(|_| b.spec_strides()).collect();
3096 (dims, strides)
3097 });
3098 let clamp = matches!(op, ElementwiseOp::Clamp(_)).then(|| {
3099 let lo = b.spec_u32(0);
3100 let hi = b.spec_u32(0);
3101 (lo, hi)
3102 });
3103
3104 let (counter, stride) = b.grid_stride(workgroup);
3105 let (scope, i) = b.begin_loop(counter, count);
3106 let indices = match &shape {
3107 Some((dims, strides)) => b.strided_indices(i, dims, strides),
3108 None => vec![i; inputs.len()],
3109 };
3110 let bits_lane =
3115 float == Storage::Half && matches!(op, ElementwiseOp::Negate | ElementwiseOp::Abs);
3116 let mut values = Vec::with_capacity(inputs.len());
3117 for (k, operand) in inputs.iter().enumerate() {
3118 let value = match lane(input_storage[k]) {
3119 Storage::Word => b.load_f32(array, *operand, indices[k]),
3120 Storage::Half if bits_lane => b.load_half_bits(array, *operand, indices[k]),
3121 Storage::Half => {
3122 let bits = b.load_half_bits(array, *operand, indices[k]);
3123 b.widen_f16(bits)
3124 }
3125 Storage::Byte => b.load_bool(array, *operand, indices[k]),
3126 Storage::Quarter(_) => unreachable!("elementwise FP8 operand"),
3128 };
3129 values.push(value);
3130 }
3131 let result = if bits_lane {
3132 let bits = values[0];
3133 match op {
3134 ElementwiseOp::Negate => {
3135 let sign = b.c_u32(0x8000);
3136 b.bxor(bits, sign)
3137 }
3138 ElementwiseOp::Abs => {
3139 let magnitude = b.c_u32(0x7fff);
3140 b.band(bits, magnitude)
3141 }
3142 _ => unreachable!("bits lane is negate/abs only"),
3143 }
3144 } else {
3145 b.elementwise_lane(op, &values, clamp)
3146 };
3147 match lane(op.output()) {
3148 Storage::Word => b.store_f32(array, output, i, result),
3149 Storage::Half if bits_lane => b.store_half_bits(array, output, i, result),
3150 Storage::Half => {
3151 let bits = b.narrow_f16(result);
3152 b.store_half_bits(array, output, i, bits);
3153 }
3154 Storage::Byte => b.store_bool(array, output, i, result),
3155 Storage::Quarter(_) => unreachable!("elementwise FP8 result"),
3156 }
3157 b.end_loop(scope, counter, stride);
3158 b.end_main();
3159 b.finish([workgroup, 1, 1])
3160}
3161
3162fn assemble_reduce(op: ReduceOp, float: Storage, workgroup: u32, buffers: u32) -> Vec<u32> {
3170 let mut b = Builder::new();
3171 let array = b.buffer_array(buffers);
3172 let input = b.spec_operand();
3173 let output = b.spec_operand();
3174 let outer = b.spec_u32(1);
3175 let axis = b.spec_u32(1);
3176 let inner = b.spec_u32(1);
3177
3178 let (counter, stride) = b.grid_stride(workgroup);
3179 let u32_ty = b.u32_ty();
3180 let f32_ty = b.f32_ty();
3181 let acc_var = b.local(f32_ty);
3182 let index_var = b.local(u32_ty);
3183 let done_var = {
3184 let bool_ty = b.bool_ty();
3185 b.local(bool_ty)
3186 };
3187 let a_var = b.local(u32_ty);
3188 let count = b.imul(outer, inner);
3189 let (scope, o) = b.begin_loop(counter, count);
3190 let outer_index = b.udiv(o, inner);
3191 let inner_index = b.umod(o, inner);
3192 let row = b.imul(outer_index, axis);
3193 let row = b.imul(row, inner);
3194 let base = b.iadd(row, inner_index);
3195 let init = match op {
3196 ReduceOp::Sum => b.c_f32(0.0),
3197 ReduceOp::Product => b.c_f32(1.0),
3198 ReduceOp::Max(_) | ReduceOp::ArgMax(_) => b.c_f32(f32::NEG_INFINITY),
3199 ReduceOp::Min(_) => b.c_f32(f32::INFINITY),
3200 };
3201 b.store(acc_var, init);
3202 let zero = b.c_u32(0);
3203 let one = b.c_u32(1);
3204 b.store(index_var, zero);
3205 let false_id = b.c_false();
3206 b.store(done_var, false_id);
3207 b.store(a_var, zero);
3208 let (inner_scope, a) = b.begin_loop(a_var, axis);
3209 let offset = b.imul(a, inner);
3210 let element = b.iadd(base, offset);
3211 let value = b.load_float(float, array, input, element);
3212 let acc = b.load(f32_ty, acc_var);
3213 match op {
3214 ReduceOp::Sum => {
3215 let next = b.fadd(acc, value);
3216 b.store(acc_var, next);
3217 }
3218 ReduceOp::Product => {
3219 let next = b.fmul(acc, value);
3220 b.store(acc_var, next);
3221 }
3222 ReduceOp::Max(nan_mode) => {
3223 let next = b.apply_max(acc, value, nan_mode);
3224 b.store(acc_var, next);
3225 }
3226 ReduceOp::Min(nan_mode) => {
3227 let next = b.apply_min(acc, value, nan_mode);
3228 b.store(acc_var, next);
3229 }
3230 ReduceOp::ArgMax(nan_mode) => {
3231 let bool_ty = b.bool_ty();
3232 let done = b.load(bool_ty, done_var);
3233 let index = b.load(u32_ty, index_var);
3234 let greater = b.fogt(value, acc);
3235 let value_nan = b.is_nan(value);
3236 let take = match nan_mode {
3237 NanMode::Propagate => {
3238 let candidate = b.lor(greater, value_nan);
3240 let not_done = b.lnot(done);
3241 let take = b.land(candidate, not_done);
3242 let next_done = b.lor(done, value_nan);
3243 b.store(done_var, next_done);
3244 take
3245 }
3246 NanMode::Ignore => {
3247 let ordered = b.lnot(value_nan);
3248 b.land(greater, ordered)
3249 }
3250 };
3251 let next_acc = b.select_f32(take, value, acc);
3252 let next_index = b.select_u32(take, a, index);
3253 b.store(acc_var, next_acc);
3254 b.store(index_var, next_index);
3255 }
3256 }
3257 b.end_loop(inner_scope, a_var, one);
3258 match op {
3259 ReduceOp::ArgMax(_) => {
3260 let index = b.load(u32_ty, index_var);
3261 b.store_word(array, output, o, index);
3262 }
3263 _ => {
3264 let acc = b.load(f32_ty, acc_var);
3265 b.store_float(float, array, output, o, acc);
3266 }
3267 }
3268 b.end_loop(scope, counter, stride);
3269 b.end_main();
3270 b.finish([workgroup, 1, 1])
3271}
3272
3273struct Slab {
3278 operand: (Id, Id),
3279 rows: u32,
3280 cols: u32,
3281 row0: Id,
3282 col0: Id,
3283 stride: Id,
3284 base: Id,
3285 row_limit: Id,
3286 col_limit: Id,
3287 transposed: bool,
3288 tile: Id,
3289}
3290
3291fn stage_slab(b: &mut Builder, input: Storage, array: Id, invocations: u32, lid: Id, slab: &Slab) {
3304 let f32_ty = b.f32_ty();
3305 let workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
3306 let zero = b.c_u32(0);
3307 let zero_f = b.c_f32(0.0);
3308 let rows_c = b.c_u32(slab.rows);
3309 let cols_c = b.c_u32(slab.cols);
3310 let elements = slab.rows * slab.cols;
3311 let slot = |b: &mut Builder, r: Id, c: Id| {
3312 if slab.transposed {
3313 let slot = b.imul(c, rows_c);
3314 b.iadd(slot, r)
3315 } else {
3316 let slot = b.imul(r, cols_c);
3317 b.iadd(slot, c)
3318 }
3319 };
3320 let locate = |b: &mut Builder, r: Id, c: Id| {
3322 let global_row = b.iadd(slab.row0, r);
3323 let global_col = b.iadd(slab.col0, c);
3324 let row_ok = b.ult(global_row, slab.row_limit);
3325 let col_ok = b.ult(global_col, slab.col_limit);
3326 let ok = b.land(row_ok, col_ok);
3327 let element = b.imul(global_row, slab.stride);
3328 let element = b.iadd(element, slab.base);
3329 let element = b.iadd(element, global_col);
3330 let element = b.select_u32(ok, element, zero);
3331 (element, ok)
3332 };
3333 let each = |b: &mut Builder, count: u32, body: &dyn Fn(&mut Builder, Id)| {
3335 for q in 0..count.div_ceil(invocations) {
3336 let offset = b.c_u32(q * invocations);
3337 let index = b.iadd(offset, lid);
3338 if count % invocations == 0 {
3339 body(b, index);
3340 } else {
3341 let limit = b.c_u32(count);
3342 let in_range = b.ult(index, limit);
3343 b.if_then(in_range, |b| body(b, index));
3344 }
3345 }
3346 };
3347 let scalar = |b: &mut Builder| {
3348 each(b, elements, &|b, index| {
3349 let r = b.udiv(index, cols_c);
3350 let c = b.umod(index, cols_c);
3351 let (element, ok) = locate(b, r, c);
3352 let loaded = b.load_float(input, array, slab.operand, element);
3353 let value = b.select_f32(ok, loaded, zero_f);
3354 let slot = slot(b, r, c);
3355 let pointer = b.access_chain(workgroup_ptr, slab.tile, &[slot]);
3356 b.store(pointer, value);
3357 });
3358 };
3359 let lanes = input.lanes();
3360 if lanes == 1 || slab.cols % lanes != 0 {
3361 scalar(b);
3362 return;
3363 }
3364 let lanes_c = b.c_u32(lanes);
3365 let words_per_row = slab.cols / lanes;
3366 let words_per_row_c = b.c_u32(words_per_row);
3367 let remainder = b.umod(slab.stride, lanes_c);
3368 let aligned = b.ieq(remainder, zero);
3369 b.if_then(aligned, |b| {
3370 each(b, elements / lanes, &|b, word| {
3371 let r = b.udiv(word, words_per_row_c);
3372 let word_in_row = b.umod(word, words_per_row_c);
3373 let c = b.imul(word_in_row, lanes_c);
3374 let (element, ok) = locate(b, r, c);
3378 let last = b.c_u32(lanes - 1);
3379 let last_col = b.iadd(c, last);
3380 let (_, last_ok) = locate(b, r, last_col);
3381 let ok = b.land(ok, last_ok);
3382 let element = b.select_u32(ok, element, zero);
3383 let bits = b.load_lane_group(input, array, slab.operand, element, lanes);
3384 for (lane, bits) in bits.into_iter().enumerate() {
3385 let widened = b.widen_bits(input, bits);
3386 let value = b.select_f32(ok, widened, zero_f);
3387 let lane_c = b.c_u32(lane as u32);
3388 let col = b.iadd(c, lane_c);
3389 let slot = slot(b, r, col);
3390 let pointer = b.access_chain(workgroup_ptr, slab.tile, &[slot]);
3391 b.store(pointer, value);
3392 }
3393 });
3394 });
3395 let unaligned = b.lnot(aligned);
3396 b.if_then(unaligned, scalar);
3397}
3398
3399fn assemble_matmul(
3428 input: Storage,
3429 output_storage: Storage,
3430 geometry: MatmulGeometry,
3431 buffers: u32,
3432) -> Vec<u32> {
3433 let MatmulGeometry {
3434 tile_x,
3435 tile_y,
3436 micro_m,
3437 micro_n,
3438 depth,
3439 } = geometry;
3440 let invocations = geometry.invocations();
3441 let block_m = geometry.block_m();
3442 let block_n = geometry.block_n();
3443 debug_assert_eq!((block_m * depth) % invocations, 0, "lhs slab stages evenly");
3444 debug_assert_eq!((block_n * depth) % invocations, 0, "rhs slab stages evenly");
3445 let mut b = Builder::new();
3446 let array = b.buffer_array(buffers);
3447 let lhs = b.spec_operand();
3448 let rhs = b.spec_operand();
3449 let output = b.spec_operand();
3450 let m = b.spec_u32(1);
3451 let n = b.spec_u32(1);
3452 let k = b.spec_u32(1);
3453 let _batch = b.spec_u32(1);
3454 let lhs_tile = b.shared_f32_array(block_m * depth);
3457 let rhs_tile = b.shared_f32_array(block_n * depth);
3458 let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
3459 let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
3460
3461 b.begin_main();
3462 let u32_ty = b.u32_ty();
3463 let f32_ty = b.f32_ty();
3464 let accumulators: Vec<Id> = (0..micro_m * micro_n).map(|_| b.local(f32_ty)).collect();
3465 let t_var = b.local(u32_ty);
3466 let kk_var = b.local(u32_ty);
3467 let tx = b.builtin_component(local_id, 0);
3468 let ty = b.builtin_component(local_id, 1);
3469 let gx = b.builtin_component(group_id, 0);
3470 let gy = b.builtin_component(group_id, 1);
3471 let z = b.builtin_component(group_id, 2);
3472 let tile_x_c = b.c_u32(tile_x);
3473 let depth_c = b.c_u32(depth);
3474 let block_m_c = b.c_u32(block_m);
3475 let block_n_c = b.c_u32(block_n);
3476 let zero = b.c_u32(0);
3477 let one = b.c_u32(1);
3478 let zero_f = b.c_f32(0.0);
3479 for accumulator in &accumulators {
3480 b.store(*accumulator, zero_f);
3481 }
3482 b.store(t_var, zero);
3483 let row0 = b.imul(gy, block_m_c);
3484 let col0 = b.imul(gx, block_n_c);
3485 let lid = b.imul(ty, tile_x_c);
3486 let lid = b.iadd(lid, tx);
3487 let lhs_batch = b.imul(z, m);
3489 let lhs_batch = b.imul(lhs_batch, k);
3490 let rhs_batch = b.imul(z, k);
3491 let rhs_batch = b.imul(rhs_batch, n);
3492 let workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
3493 let steps = {
3494 let k_plus = b.iadd(k, depth_c);
3495 let k_plus = b.isub(k_plus, one);
3496 b.udiv(k_plus, depth_c)
3497 };
3498 let (outer, t) = b.begin_loop(t_var, steps);
3499 let k0 = b.imul(t, depth_c);
3500 stage_slab(
3503 &mut b,
3504 input,
3505 array,
3506 invocations,
3507 lid,
3508 &Slab {
3509 operand: lhs,
3510 rows: block_m,
3511 cols: depth,
3512 row0,
3513 col0: k0,
3514 stride: k,
3515 base: lhs_batch,
3516 row_limit: m,
3517 col_limit: k,
3518 transposed: true,
3519 tile: lhs_tile,
3520 },
3521 );
3522 stage_slab(
3523 &mut b,
3524 input,
3525 array,
3526 invocations,
3527 lid,
3528 &Slab {
3529 operand: rhs,
3530 rows: depth,
3531 cols: block_n,
3532 row0: k0,
3533 col0,
3534 stride: n,
3535 base: rhs_batch,
3536 row_limit: k,
3537 col_limit: n,
3538 transposed: false,
3539 tile: rhs_tile,
3540 },
3541 );
3542 b.workgroup_barrier();
3543 let remaining = b.isub(k, k0);
3544 let k_max = b.umin(remaining, depth_c);
3545 let step = |b: &mut Builder, kk: Id| {
3547 let lhs_row = b.imul(kk, block_m_c);
3548 let rhs_row = b.imul(kk, block_n_c);
3549 let a: Vec<Id> = (0..micro_m)
3550 .map(|i| {
3551 let offset = b.c_u32(i * tile_y);
3552 let local_row = b.iadd(offset, ty);
3553 let slot = b.iadd(lhs_row, local_row);
3554 let pointer = b.access_chain(workgroup_ptr, lhs_tile, &[slot]);
3555 b.load(f32_ty, pointer)
3556 })
3557 .collect();
3558 let bv: Vec<Id> = (0..micro_n)
3559 .map(|j| {
3560 let offset = b.c_u32(j * tile_x);
3561 let local_col = b.iadd(offset, tx);
3562 let slot = b.iadd(rhs_row, local_col);
3563 let pointer = b.access_chain(workgroup_ptr, rhs_tile, &[slot]);
3564 b.load(f32_ty, pointer)
3565 })
3566 .collect();
3567 for i in 0..micro_m {
3568 for j in 0..micro_n {
3569 let accumulator = accumulators[(i * micro_n + j) as usize];
3570 let acc = b.load(f32_ty, accumulator);
3571 let next = b.fma(a[i as usize], bv[j as usize], acc);
3572 b.store(accumulator, next);
3573 }
3574 }
3575 };
3576 b.store(kk_var, zero);
3577 let (inner, kk) = b.begin_loop(kk_var, k_max);
3578 step(&mut b, kk);
3579 b.end_loop(inner, kk_var, one);
3580 b.workgroup_barrier();
3581 b.end_loop(outer, t_var, one);
3582 let out_batch = b.imul(z, m);
3583 let out_batch = b.imul(out_batch, n);
3584 for i in 0..micro_m {
3585 let offset = b.c_u32(i * tile_y);
3586 let row = b.iadd(row0, offset);
3587 let row = b.iadd(row, ty);
3588 let row_ok = b.ult(row, m);
3589 let out_row = b.imul(row, n);
3590 let out_row = b.iadd(out_batch, out_row);
3591 for j in 0..micro_n {
3592 let offset = b.c_u32(j * tile_x);
3593 let col = b.iadd(col0, offset);
3594 let col = b.iadd(col, tx);
3595 let col_ok = b.ult(col, n);
3596 let in_range = b.land(row_ok, col_ok);
3597 let accumulator = accumulators[(i * micro_n + j) as usize];
3598 b.if_then(in_range, |b| {
3599 let out_index = b.iadd(out_row, col);
3600 let acc = b.load(f32_ty, accumulator);
3601 match output_storage {
3602 Storage::Half => {
3603 let bits = b.narrow_f16(acc);
3604 b.store_half_bits(array, output, out_index, bits);
3605 }
3606 Storage::Byte | Storage::Quarter(_) => unreachable!("MATMUL result storage"),
3609 Storage::Word => b.store_f32(array, output, out_index, acc),
3610 }
3611 });
3612 }
3613 }
3614 b.end_main();
3615 b.finish(geometry.local_size())
3616}
3617
3618fn assemble_matmul_stream(rhs_storage: Storage, output_storage: Storage, buffers: u32) -> Vec<u32> {
3645 let lanes = rhs_storage.lanes();
3646 let words = STREAM_COLUMNS / lanes;
3647 let splits = STREAM_WORKGROUP / words;
3648 let rows = STREAM_ROWS;
3649 let mut b = Builder::new();
3650 let array = b.buffer_array(buffers);
3651 let lhs = b.spec_operand();
3652 let rhs = b.spec_operand();
3653 let output = b.spec_operand();
3654 let m = b.spec_u32(1);
3655 let n = b.spec_u32(1);
3656 let k = b.spec_u32(1);
3657 let _batch = b.spec_u32(1);
3658 let partials = b.shared_f32_array(STREAM_WORKGROUP * rows * lanes);
3659 let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
3660 let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
3661
3662 b.begin_main();
3663 let u32_ty = b.u32_ty();
3664 let f32_ty = b.f32_ty();
3665 let accumulators: Vec<Id> = (0..rows * lanes).map(|_| b.local(f32_ty)).collect();
3666 let kk_var = b.local(u32_ty);
3667 let lid = b.builtin_component(local_id, 0);
3668 let gx = b.builtin_component(group_id, 0);
3669 let z = b.builtin_component(group_id, 2);
3670 let zero = b.c_u32(0);
3671 let one = b.c_u32(1);
3672 let zero_f = b.c_f32(0.0);
3673 let words_c = b.c_u32(words);
3674 let lanes_c = b.c_u32(lanes);
3675 let splits_c = b.c_u32(splits);
3676 let columns_c = b.c_u32(STREAM_COLUMNS);
3677 for accumulator in &accumulators {
3678 b.store(*accumulator, zero_f);
3679 }
3680 let w = b.umod(lid, words_c);
3681 let slice = b.udiv(lid, words_c);
3682 let col0 = b.imul(gx, columns_c);
3683 let word_col = b.imul(w, lanes_c);
3684 let col_base = b.iadd(col0, word_col);
3685 let splits_less_one = b.c_u32(splits - 1);
3687 let k_plus = b.iadd(k, splits_less_one);
3688 let chunk = b.udiv(k_plus, splits_c);
3689 let k_begin = b.imul(slice, chunk);
3690 let k_end = b.iadd(k_begin, chunk);
3691 let k_end = b.umin(k_end, k);
3692 let lhs_batch = b.imul(z, m);
3693 let lhs_batch = b.imul(lhs_batch, k);
3694 let rhs_batch = b.imul(z, k);
3695 let rhs_batch = b.imul(rhs_batch, n);
3696 let rhs_col = b.iadd(rhs_batch, col_base);
3697 let row_ok: Vec<Id> = (0..rows)
3698 .map(|i| {
3699 let row = b.c_u32(i);
3700 b.ult(row, m)
3701 })
3702 .collect();
3703 let lhs_rows: Vec<Id> = (0..rows)
3704 .map(|i| {
3705 let row = b.c_u32(i);
3706 let offset = b.imul(row, k);
3707 b.iadd(lhs_batch, offset)
3708 })
3709 .collect();
3710 let rhs_word_element = |b: &mut Builder, kk: Id| {
3712 let col_ok = b.ult(col_base, n);
3713 let element = b.imul(kk, n);
3714 let element = b.iadd(element, rhs_col);
3715 b.select_u32(col_ok, element, zero)
3716 };
3717 let accumulate = |b: &mut Builder, load_rhs: &dyn Fn(&mut Builder, Id) -> Vec<Id>| {
3719 b.store(kk_var, k_begin);
3720 let (scope, kk) = b.begin_loop(kk_var, k_end);
3721 let weights = load_rhs(b, kk);
3722 for i in 0..rows as usize {
3723 b.if_then(row_ok[i], |b| {
3724 let element = b.iadd(lhs_rows[i], kk);
3725 let a = b.load_f32(array, lhs, element);
3726 for (l, weight) in weights.iter().enumerate() {
3727 let accumulator = accumulators[i * lanes as usize + l];
3728 let acc = b.load(f32_ty, accumulator);
3729 let next = b.fma(a, *weight, acc);
3730 b.store(accumulator, next);
3731 }
3732 });
3733 }
3734 b.end_loop(scope, kk_var, one);
3735 };
3736 if lanes == 1 {
3737 accumulate(&mut b, &|b, kk| {
3738 let element = rhs_word_element(b, kk);
3739 vec![b.load_f32(array, rhs, element)]
3740 });
3741 } else {
3742 let remainder = b.umod(n, lanes_c);
3743 let aligned = b.ieq(remainder, zero);
3744 b.if_then(aligned, |b| {
3745 accumulate(b, &|b, kk| {
3746 let element = rhs_word_element(b, kk);
3747 let bits = b.load_lane_group(rhs_storage, array, rhs, element, lanes);
3748 bits.into_iter()
3749 .map(|bits| b.widen_bits(rhs_storage, bits))
3750 .collect()
3751 });
3752 });
3753 let unaligned = b.lnot(aligned);
3754 b.if_then(unaligned, |b| {
3755 accumulate(b, &|b, kk| {
3756 (0..lanes)
3757 .map(|l| {
3758 let l = b.c_u32(l);
3759 let col = b.iadd(col_base, l);
3760 let col_ok = b.ult(col, n);
3761 let element = b.imul(kk, n);
3762 let element = b.iadd(element, rhs_batch);
3763 let element = b.iadd(element, col);
3764 let element = b.select_u32(col_ok, element, zero);
3765 b.load_float(rhs_storage, array, rhs, element)
3766 })
3767 .collect()
3768 });
3769 });
3770 }
3771 let workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
3773 let block = b.c_u32(rows * lanes);
3774 let base = b.imul(lid, block);
3775 for (index, accumulator) in accumulators.iter().enumerate() {
3776 let offset = b.c_u32(index as u32);
3777 let slot = b.iadd(base, offset);
3778 let value = b.load(f32_ty, *accumulator);
3779 let pointer = b.access_chain(workgroup_ptr, partials, &[slot]);
3780 b.store(pointer, value);
3781 }
3782 b.workgroup_barrier();
3783 let per_invocation = rows * STREAM_COLUMNS / STREAM_WORKGROUP;
3786 let per_invocation_c = b.c_u32(per_invocation);
3787 let first = b.imul(lid, per_invocation_c);
3788 let out_batch = b.imul(z, m);
3789 let out_batch = b.imul(out_batch, n);
3790 for j in 0..per_invocation {
3791 let j = b.c_u32(j);
3792 let o = b.iadd(first, j);
3793 let row = b.udiv(o, columns_c);
3794 let col = b.umod(o, columns_c);
3795 let word = b.udiv(col, lanes_c);
3796 let lane = b.umod(col, lanes_c);
3797 let within = b.imul(row, lanes_c);
3798 let within = b.iadd(within, lane);
3799 let mut sum = None;
3800 for split in 0..splits {
3801 let offset = b.c_u32(split * words);
3802 let owner = b.iadd(offset, word);
3803 let slot = b.imul(owner, block);
3804 let slot = b.iadd(slot, within);
3805 let pointer = b.access_chain(workgroup_ptr, partials, &[slot]);
3806 let partial = b.load(f32_ty, pointer);
3807 sum = Some(match sum {
3808 None => partial,
3809 Some(sum) => b.fadd(sum, partial),
3810 });
3811 }
3812 let sum = sum.expect("at least one slice");
3813 let global_col = b.iadd(col0, col);
3814 let row_in = b.ult(row, m);
3815 let col_in = b.ult(global_col, n);
3816 let in_range = b.land(row_in, col_in);
3817 b.if_then(in_range, |b| {
3818 let out = b.imul(row, n);
3819 let out = b.iadd(out, out_batch);
3820 let out = b.iadd(out, global_col);
3821 b.store_float(output_storage, array, output, out, sum);
3822 });
3823 }
3824 b.end_main();
3825 b.finish([STREAM_WORKGROUP, 1, 1])
3826}
3827
3828fn assemble_nvfp4_matmul_cooperative(buffers: u32) -> Vec<u32> {
3833 let mut b = Builder::new();
3834 b.enable_cooperative_matrix();
3835 let array = b.buffer_array(buffers);
3836 let activation = b.spec_operand();
3837 let packed = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
3838 let block_scales = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
3839 let tensor_scale = b.spec_operand();
3840 let output = b.spec_operand();
3841 let m = b.spec_u32(1);
3842 let n = b.spec_u32(1);
3843 let k = b.spec_u32(16);
3844 let epilogue = b.spec_u32(0);
3845 let weight_mode = b.spec_u32(0);
3846
3847 let a_tile = b.shared_f16_array(8 * 16);
3848 let b_tile = b.shared_f16_array(16 * 16);
3849 let c_tile = b.shared_f32_array(8 * 16);
3850 let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
3851 let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
3852
3853 b.begin_main();
3854 let u32_ty = b.u32_ty();
3855 let f32_ty = b.f32_ty();
3856 let f16_ty = b.f16_ty();
3857 let matrix_a_ty = b.cooperative_matrix_ty(f16_ty, 8, 16, 0);
3858 let matrix_b_ty = b.cooperative_matrix_ty(f16_ty, 16, 16, 1);
3859 let matrix_c_ty = b.cooperative_matrix_ty(f32_ty, 8, 16, 2);
3860 let accumulator = b.local(matrix_c_ty);
3861 let block_var = b.local(u32_ty);
3862 let fp4 = b.private_u32_array(&[
3863 0.0f32.to_bits(),
3864 0.5f32.to_bits(),
3865 1.0f32.to_bits(),
3866 1.5f32.to_bits(),
3867 2.0f32.to_bits(),
3868 3.0f32.to_bits(),
3869 4.0f32.to_bits(),
3870 6.0f32.to_bits(),
3871 (-0.0f32).to_bits(),
3872 (-0.5f32).to_bits(),
3873 (-1.0f32).to_bits(),
3874 (-1.5f32).to_bits(),
3875 (-2.0f32).to_bits(),
3876 (-3.0f32).to_bits(),
3877 (-4.0f32).to_bits(),
3878 (-6.0f32).to_bits(),
3879 ]);
3880 let lid = b.builtin_component(local_id, 0);
3881 let output_tile = b.builtin_component(group_id, 0);
3882 let token_tile = b.builtin_component(group_id, 1);
3883 let zero = b.c_u32(0);
3884 let one = b.c_u32(1);
3885 let two = b.c_u32(2);
3886 let four = b.c_u32(4);
3887 let eight = b.c_u32(8);
3888 let sixteen = b.c_u32(16);
3889 let zero_f = b.c_f32(0.0);
3890 let zero_h = b.f32_to_f16(zero_f);
3891 let blocks = b.udiv(k, sixteen);
3892 let output_base = b.imul(output_tile, sixteen);
3893 let token_base = b.imul(token_tile, eight);
3894 let shared_mode = b.ieq(weight_mode, zero);
3897 let row_bytes = b.udiv(k, two);
3898
3899 let f32_workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
3900 let f16_workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f16_ty);
3901 for wave in 0..4 {
3903 let offset = b.c_u32(wave * 32);
3904 let element = b.iadd(lid, offset);
3905 let pointer = b.access_chain(f32_workgroup_ptr, c_tile, &[element]);
3906 b.store(pointer, zero_f);
3907 }
3908 b.workgroup_barrier();
3909 let c_pointer = b.access_chain(f32_workgroup_ptr, c_tile, &[zero]);
3910 let initial = b.cooperative_load(matrix_c_ty, c_pointer, sixteen);
3911 b.store(accumulator, initial);
3912
3913 b.store(block_var, zero);
3914 let (scope, block) = b.begin_loop(block_var, blocks);
3915 let activation_block = b.imul(block, sixteen);
3916 for wave in 0..4 {
3919 let offset = b.c_u32(wave * 32);
3920 let element = b.iadd(lid, offset);
3921 let tile_row = b.udiv(element, sixteen);
3922 let tile_column = b.umod(element, sixteen);
3923 let token = b.iadd(token_base, tile_row);
3924 let token_in_range = b.ult(token, m);
3925 let activation_row = b.imul(token, k);
3926 let source = b.iadd(activation_row, activation_block);
3927 let source = b.iadd(source, tile_column);
3928 let source = b.select_u32(token_in_range, source, zero);
3929 let value = b.load_f32(array, activation, source);
3930 let value = b.select_f32(token_in_range, value, zero_f);
3931 let value = b.f32_to_f16(value);
3932 let pointer = b.access_chain(f16_workgroup_ptr, a_tile, &[element]);
3933 b.store(pointer, value);
3934 }
3935
3936 for wave in 0..8 {
3939 let offset = b.c_u32(wave * 32);
3940 let element = b.iadd(lid, offset);
3941 let k_lane = b.udiv(element, sixteen);
3942 let column = b.umod(element, sixteen);
3943 let output_row = b.iadd(output_base, column);
3944 let in_range = b.ult(output_row, n);
3945 let pointer = b.access_chain(f16_workgroup_ptr, b_tile, &[element]);
3946 b.store(pointer, zero_h);
3947 b.if_then(in_range, |b| {
3948 let weight_row = output_row;
3949 let packed_row = b.imul(weight_row, row_bytes);
3950 let packed_block = b.imul(block, eight);
3951 let packed_base = b.iadd(packed_row, packed_block);
3952 let packed_lane = b.udiv(k_lane, two);
3953 let packed_index = b.iadd(packed_base, packed_lane);
3954 let codes = b.load_byte_bits(array, packed[0], packed_index);
3955 let nibble_mask = b.c_u32(15);
3956 let low = b.band(codes, nibble_mask);
3957 let high = b.shr(codes, four);
3958 let parity = b.band(k_lane, one);
3959 let odd = b.ine(parity, zero);
3960 let code = b.select_u32(odd, high, low);
3961 let private_u32 = b.pointer(STORAGE_CLASS_PRIVATE, u32_ty);
3962 let fp4_pointer = b.access_chain(private_u32, fp4, &[code]);
3963 let weight_bits = b.load(u32_ty, fp4_pointer);
3964 let weight = b.bitcast_f32(weight_bits);
3965 let scale_row = b.imul(weight_row, blocks);
3966 let scale_index = b.iadd(scale_row, block);
3967 let scale_bits = b.load_byte_bits(array, block_scales[0], scale_index);
3968 let scale = b.widen_fp8(Fp8Format::E4M3, scale_bits);
3969 let weight = b.fmul(weight, scale);
3970 let weight = b.f32_to_f16(weight);
3971 b.store(pointer, weight);
3972 });
3973 }
3974 b.workgroup_barrier();
3975 let a_pointer = b.access_chain(f16_workgroup_ptr, a_tile, &[zero]);
3976 let b_pointer = b.access_chain(f16_workgroup_ptr, b_tile, &[zero]);
3977 let matrix_a = b.cooperative_load(matrix_a_ty, a_pointer, sixteen);
3978 let matrix_b = b.cooperative_load(matrix_b_ty, b_pointer, sixteen);
3979 let matrix_c = b.load(matrix_c_ty, accumulator);
3980 let matrix_c = b.cooperative_mul_add(matrix_c_ty, matrix_a, matrix_b, matrix_c);
3981 b.store(accumulator, matrix_c);
3982 b.workgroup_barrier();
3983 b.end_loop(scope, block_var, one);
3984
3985 let matrix_c = b.load(matrix_c_ty, accumulator);
3986 b.cooperative_store(c_pointer, matrix_c, sixteen);
3987 b.workgroup_barrier();
3988 for wave in 0..4 {
3989 let offset = b.c_u32(wave * 32);
3990 let element = b.iadd(lid, offset);
3991 let tile_row = b.udiv(element, sixteen);
3992 let tile_column = b.umod(element, sixteen);
3993 let token = b.iadd(token_base, tile_row);
3994 let output_row = b.iadd(output_base, tile_column);
3995 let token_in_range = b.ult(token, m);
3996 let row_in_tensor = b.ult(output_row, n);
3997 let write = b.land(token_in_range, row_in_tensor);
3998 let write = b.land(write, shared_mode);
3999 b.if_then(write, |b| {
4000 let pointer = b.access_chain(f32_workgroup_ptr, c_tile, &[element]);
4001 let sum = b.load(f32_ty, pointer);
4002 let scale = b.load_f32(array, tensor_scale, zero);
4003 let sum = b.fmul(sum, scale);
4004 let negated = b.fneg(sum);
4005 let exp = b.ext_f32(GLSL_EXP, &[negated]);
4006 let one_f = b.c_f32(1.0);
4007 let denominator = b.fadd(one_f, exp);
4008 let sigmoid = b.fdiv(one_f, denominator);
4009 let silu = b.fmul(sum, sigmoid);
4010 let is_silu = b.ieq(epilogue, one);
4011 let is_sigmoid = b.ieq(epilogue, two);
4012 let three = b.c_u32(3);
4013 let zero_f = b.c_f32(0.0);
4014 let squared_input = sum;
4015 let positive = b.fogt(squared_input, zero_f);
4016 let rectified = b.select_f32(positive, squared_input, zero_f);
4017 let squared = b.fmul(rectified, rectified);
4018 let is_squared = b.ieq(epilogue, three);
4019 let sum = b.select_f32(is_silu, silu, sum);
4020 let sum = b.select_f32(is_sigmoid, sigmoid, sum);
4021 let sum = b.select_f32(is_squared, squared, sum);
4022 let destination = b.imul(token, n);
4023 let destination = b.iadd(destination, output_row);
4024 b.store_f32(array, output, destination, sum);
4025 });
4026 }
4027 b.end_main();
4028 b.finish([32, 1, 1])
4029}
4030
4031fn assemble_nvfp4_matmul_subgroup(buffers: u32) -> Vec<u32> {
4035 let mut b = Builder::new();
4036 b.enable_subgroup_arithmetic();
4037 let array = b.buffer_array(buffers);
4038 let activation = b.spec_operand();
4039 let packed = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
4040 let block_scales = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
4041 let tensor_scale = b.spec_operand();
4042 let output = b.spec_operand();
4043 let _m = b.spec_u32(1);
4044 let n = b.spec_u32(1);
4045 let k = b.spec_u32(16);
4046 let epilogue = b.spec_u32(0);
4047 let weight_mode = b.spec_u32(0);
4048 let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
4049 let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
4050
4051 b.begin_main();
4052 let u32_ty = b.u32_ty();
4053 let f32_ty = b.f32_ty();
4054 let block_var = b.local(u32_ty);
4055 let accumulators = core::array::from_fn::<_, 4, _>(|_| b.local(f32_ty));
4056 let fp4 = b.private_u32_array(&[
4057 0.0f32.to_bits(),
4058 0.5f32.to_bits(),
4059 1.0f32.to_bits(),
4060 1.5f32.to_bits(),
4061 2.0f32.to_bits(),
4062 3.0f32.to_bits(),
4063 4.0f32.to_bits(),
4064 6.0f32.to_bits(),
4065 (-0.0f32).to_bits(),
4066 (-0.5f32).to_bits(),
4067 (-1.0f32).to_bits(),
4068 (-1.5f32).to_bits(),
4069 (-2.0f32).to_bits(),
4070 (-3.0f32).to_bits(),
4071 (-4.0f32).to_bits(),
4072 (-6.0f32).to_bits(),
4073 ]);
4074 let lid = b.builtin_component(local_id, 0);
4075 let output_group = b.builtin_component(group_id, 0);
4076 let token = b.builtin_component(group_id, 1);
4077 let zero = b.c_u32(0);
4078 let zero_f = b.c_f32(0.0);
4079 let four = b.c_u32(4);
4080 let eight = b.c_u32(8);
4081 let sixteen = b.c_u32(16);
4082 let thirty_two = b.c_u32(32);
4083 let nibble_mask = b.c_u32(15);
4084 let blocks = b.udiv(k, sixteen);
4085 let output_base = b.imul(output_group, four);
4086 let mut packed_operand = packed[0];
4087 let mut scale_operand = block_scales[0];
4088 for index in 1..6 {
4089 let index_id = b.c_u32(index as u32);
4090 let selected = b.ieq(token, index_id);
4091 packed_operand.0 = b.select_u32(selected, packed[index].0, packed_operand.0);
4092 packed_operand.1 = b.select_u32(selected, packed[index].1, packed_operand.1);
4093 scale_operand.0 = b.select_u32(selected, block_scales[index].0, scale_operand.0);
4094 scale_operand.1 = b.select_u32(selected, block_scales[index].1, scale_operand.1);
4095 }
4096 let contiguous_mode = b.c_u32(1);
4097 let contiguous = b.ieq(weight_mode, contiguous_mode);
4098 let weight_batch = b.select_u32(contiguous, token, zero);
4099 let batch_rows = b.imul(weight_batch, n);
4100 let two = b.c_u32(2);
4101 let row_bytes = b.udiv(k, two);
4102 let activation_row = b.imul(token, k);
4103 let output_rows = core::array::from_fn::<_, 4, _>(|index| {
4104 let offset = b.c_u32(index as u32);
4105 b.iadd(output_base, offset)
4106 });
4107 let row_valid = output_rows.map(|row| b.ult(row, n));
4108 let weight_rows = core::array::from_fn::<_, 4, _>(|index| {
4109 let safe_row = b.select_u32(row_valid[index], output_rows[index], zero);
4110 b.iadd(batch_rows, safe_row)
4111 });
4112 let row_blocks = weight_rows.map(|row| b.imul(row, blocks));
4113 let packed_rows = weight_rows.map(|row| b.imul(row, row_bytes));
4114 for accumulator in accumulators {
4115 b.store(accumulator, zero_f);
4116 }
4117 b.store(block_var, lid);
4118 let (scope, block) = b.begin_loop(block_var, blocks);
4119 let scales = core::array::from_fn::<_, 4, _>(|index| {
4120 let scale_element = b.iadd(row_blocks[index], block);
4121 let scale_bits = b.load_byte_bits(array, scale_operand, scale_element);
4122 b.widen_fp8(Fp8Format::E4M3, scale_bits)
4123 });
4124 let packed_block = b.imul(block, eight);
4125 let packed_bases = packed_rows.map(|row| b.iadd(row, packed_block));
4126 let activation_block = b.imul(block, sixteen);
4127 let activation_base = b.iadd(activation_row, activation_block);
4128 for byte in 0..8 {
4129 let byte_offset = b.c_u32(byte);
4130 let codes = core::array::from_fn::<_, 4, _>(|index| {
4131 let byte_index = b.iadd(packed_bases[index], byte_offset);
4132 b.load_byte_bits(array, packed_operand, byte_index)
4133 });
4134 for lane in 0..2 {
4135 let inner = b.c_u32(byte * 2 + lane as u32);
4136 let element = b.iadd(activation_base, inner);
4137 let value = b.load_f32(array, activation, element);
4138 for index in 0..4 {
4139 let code = if lane == 0 {
4140 b.band(codes[index], nibble_mask)
4141 } else {
4142 b.shr(codes[index], four)
4143 };
4144 let pointer_ty = b.pointer(STORAGE_CLASS_PRIVATE, u32_ty);
4145 let pointer = b.access_chain(pointer_ty, fp4, &[code]);
4146 let weight_bits = b.load(u32_ty, pointer);
4147 let weight = b.bitcast_f32(weight_bits);
4148 let weight = b.fmul(weight, scales[index]);
4149 let acc = b.load(f32_ty, accumulators[index]);
4150 let next = b.fma(value, weight, acc);
4151 b.store(accumulators[index], next);
4152 }
4153 }
4154 }
4155 b.end_loop(scope, block_var, thirty_two);
4156
4157 let sums = accumulators.map(|accumulator| {
4158 let partial = b.load(f32_ty, accumulator);
4159 b.subgroup_sum_f32(partial)
4160 });
4161
4162 let first = b.ieq(lid, zero);
4163 b.if_then(first, |b| {
4164 let shared_mode = b.c_u32(0);
4165 let batched = b.ine(weight_mode, shared_mode);
4166 let tensor_scale_index = b.select_u32(batched, token, zero);
4167 let tensor_scale = b.load_f32(array, tensor_scale, tensor_scale_index);
4168 for index in 0..4 {
4169 b.if_then(row_valid[index], |b| {
4170 let sum = b.fmul(sums[index], tensor_scale);
4171 let negated = b.fneg(sum);
4172 let exp = b.ext_f32(GLSL_EXP, &[negated]);
4173 let one_f = b.c_f32(1.0);
4174 let denominator = b.fadd(one_f, exp);
4175 let sigmoid = b.fdiv(one_f, denominator);
4176 let silu = b.fmul(sum, sigmoid);
4177 let silu_mode = b.c_u32(1);
4178 let sigmoid_mode = b.c_u32(2);
4179 let squared_mode = b.c_u32(3);
4180 let is_silu = b.ieq(epilogue, silu_mode);
4181 let is_sigmoid = b.ieq(epilogue, sigmoid_mode);
4182 let zero_f = b.c_f32(0.0);
4183 let squared_input = sum;
4184 let positive = b.fogt(squared_input, zero_f);
4185 let rectified = b.select_f32(positive, squared_input, zero_f);
4186 let squared = b.fmul(rectified, rectified);
4187 let is_squared = b.ieq(epilogue, squared_mode);
4188 let sum = b.select_f32(is_silu, silu, sum);
4189 let sum = b.select_f32(is_sigmoid, sigmoid, sum);
4190 let sum = b.select_f32(is_squared, squared, sum);
4191 let element = b.imul(token, n);
4192 let element = b.iadd(element, output_rows[index]);
4193 b.store_f32(array, output, element, sum);
4194 });
4195 }
4196 });
4197 b.end_main();
4198 b.finish([32, 1, 1])
4199}
4200
4201fn assemble_nvfp4_matmul(buffers: u32) -> Vec<u32> {
4207 let mut b = Builder::new();
4208 let array = b.buffer_array(buffers);
4209 let activation = b.spec_operand();
4210 let packed = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
4211 let block_scales = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
4212 let tensor_scale = b.spec_operand();
4213 let output = b.spec_operand();
4214 let _m = b.spec_u32(1);
4215 let n = b.spec_u32(1);
4216 let k = b.spec_u32(16);
4217 let epilogue = b.spec_u32(0);
4218 let weight_mode = b.spec_u32(0);
4219 let partials = b.shared_f32_array(STREAM_WORKGROUP);
4220 let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
4221 let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
4222
4223 b.begin_main();
4224 let u32_ty = b.u32_ty();
4225 let f32_ty = b.f32_ty();
4226 let block_var = b.local(u32_ty);
4227 let accumulator = b.local(f32_ty);
4228 let fp4 = b.private_u32_array(&[
4229 0.0f32.to_bits(),
4230 0.5f32.to_bits(),
4231 1.0f32.to_bits(),
4232 1.5f32.to_bits(),
4233 2.0f32.to_bits(),
4234 3.0f32.to_bits(),
4235 4.0f32.to_bits(),
4236 6.0f32.to_bits(),
4237 (-0.0f32).to_bits(),
4238 (-0.5f32).to_bits(),
4239 (-1.0f32).to_bits(),
4240 (-1.5f32).to_bits(),
4241 (-2.0f32).to_bits(),
4242 (-3.0f32).to_bits(),
4243 (-4.0f32).to_bits(),
4244 (-6.0f32).to_bits(),
4245 ]);
4246 let lid = b.builtin_component(local_id, 0);
4247 let out_row = b.builtin_component(group_id, 0);
4248 let token = b.builtin_component(group_id, 1);
4249 let zero = b.c_u32(0);
4250 let zero_f = b.c_f32(0.0);
4251 let one = b.c_u32(1);
4252 let four = b.c_u32(4);
4253 let eight = b.c_u32(8);
4254 let sixteen = b.c_u32(16);
4255 let sixty_four = b.c_u32(STREAM_WORKGROUP);
4256 let nibble_mask = b.c_u32(15);
4257 let blocks = b.udiv(k, sixteen);
4258 let mut packed_operand = packed[0];
4259 let mut scale_operand = block_scales[0];
4260 for index in 1..6 {
4261 let index_id = b.c_u32(index as u32);
4262 let selected = b.ieq(token, index_id);
4263 packed_operand.0 = b.select_u32(selected, packed[index].0, packed_operand.0);
4264 packed_operand.1 = b.select_u32(selected, packed[index].1, packed_operand.1);
4265 scale_operand.0 = b.select_u32(selected, block_scales[index].0, scale_operand.0);
4266 scale_operand.1 = b.select_u32(selected, block_scales[index].1, scale_operand.1);
4267 }
4268 let contiguous_mode = b.c_u32(1);
4269 let contiguous = b.ieq(weight_mode, contiguous_mode);
4270 let weight_batch = b.select_u32(contiguous, token, zero);
4271 let batch_rows = b.imul(weight_batch, n);
4272 let weight_row = b.iadd(batch_rows, out_row);
4273 let row_blocks = b.imul(weight_row, blocks);
4274 let two = b.c_u32(2);
4275 let row_bytes = b.udiv(k, two);
4276 let packed_row = b.imul(weight_row, row_bytes);
4277 let activation_row = b.imul(token, k);
4278 b.store(accumulator, zero_f);
4279 b.store(block_var, lid);
4280 let (scope, block) = b.begin_loop(block_var, blocks);
4281 let scale_element = b.iadd(row_blocks, block);
4282 let scale_bits = b.load_byte_bits(array, scale_operand, scale_element);
4283 let scale = b.widen_fp8(Fp8Format::E4M3, scale_bits);
4284 let packed_block = b.imul(block, eight);
4285 let packed_base = b.iadd(packed_row, packed_block);
4286 let activation_block = b.imul(block, sixteen);
4287 let activation_base = b.iadd(activation_row, activation_block);
4288 for byte in 0..8 {
4289 let byte_offset = b.c_u32(byte);
4290 let byte_index = b.iadd(packed_base, byte_offset);
4291 let codes = b.load_byte_bits(array, packed_operand, byte_index);
4292 let low = b.band(codes, nibble_mask);
4293 let high = b.shr(codes, four);
4294 for (lane, code) in [low, high].into_iter().enumerate() {
4295 let pointer_ty = b.pointer(STORAGE_CLASS_PRIVATE, u32_ty);
4296 let pointer = b.access_chain(pointer_ty, fp4, &[code]);
4297 let weight_bits = b.load(u32_ty, pointer);
4298 let weight = b.bitcast_f32(weight_bits);
4299 let weight = b.fmul(weight, scale);
4300 let inner = b.c_u32(byte * 2 + lane as u32);
4301 let element = b.iadd(activation_base, inner);
4302 let value = b.load_f32(array, activation, element);
4303 let acc = b.load(f32_ty, accumulator);
4304 let next = b.fma(value, weight, acc);
4305 b.store(accumulator, next);
4306 }
4307 }
4308 b.end_loop(scope, block_var, sixty_four);
4309
4310 let workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
4311 let partial_ptr = b.access_chain(workgroup_ptr, partials, &[lid]);
4312 let partial = b.load(f32_ty, accumulator);
4313 b.store(partial_ptr, partial);
4314 b.workgroup_barrier();
4315
4316 let first = b.ieq(lid, zero);
4317 b.if_then(first, |b| {
4318 let sum_var = b.local(f32_ty);
4319 let index_var = b.local(u32_ty);
4320 b.store(sum_var, zero_f);
4321 b.store(index_var, zero);
4322 let (sum_scope, index) = b.begin_loop(index_var, sixty_four);
4323 let pointer = b.access_chain(workgroup_ptr, partials, &[index]);
4324 let value = b.load(f32_ty, pointer);
4325 let sum = b.load(f32_ty, sum_var);
4326 let next = b.fadd(sum, value);
4327 b.store(sum_var, next);
4328 b.end_loop(sum_scope, index_var, one);
4329 let sum = b.load(f32_ty, sum_var);
4330 let shared_mode = b.c_u32(0);
4331 let batched = b.ine(weight_mode, shared_mode);
4332 let tensor_scale_index = b.select_u32(batched, token, zero);
4333 let tensor_scale = b.load_f32(array, tensor_scale, tensor_scale_index);
4334 let sum = b.fmul(sum, tensor_scale);
4335 let negated = b.fneg(sum);
4336 let exp = b.ext_f32(GLSL_EXP, &[negated]);
4337 let one_f = b.c_f32(1.0);
4338 let denominator = b.fadd(one_f, exp);
4339 let sigmoid = b.fdiv(one_f, denominator);
4340 let silu = b.fmul(sum, sigmoid);
4341 let silu_mode = b.c_u32(1);
4342 let sigmoid_mode = b.c_u32(2);
4343 let squared_mode = b.c_u32(3);
4344 let is_silu = b.ieq(epilogue, silu_mode);
4345 let is_sigmoid = b.ieq(epilogue, sigmoid_mode);
4346 let zero_f = b.c_f32(0.0);
4347 let squared_input = sum;
4348 let positive = b.fogt(squared_input, zero_f);
4349 let rectified = b.select_f32(positive, squared_input, zero_f);
4350 let squared = b.fmul(rectified, rectified);
4351 let is_squared = b.ieq(epilogue, squared_mode);
4352 let sum = b.select_f32(is_silu, silu, sum);
4353 let sum = b.select_f32(is_sigmoid, sigmoid, sum);
4354 let sum = b.select_f32(is_squared, squared, sum);
4355 let element = b.imul(token, n);
4356 let element = b.iadd(element, out_row);
4357 b.store_f32(array, output, element, sum);
4358 });
4359 b.end_main();
4360 b.finish([STREAM_WORKGROUP, 1, 1])
4361}
4362
4363fn assemble_max_pool(nan_mode: NanMode, float: Storage, workgroup: u32, buffers: u32) -> Vec<u32> {
4371 let mut b = Builder::new();
4372 let array = b.buffer_array(buffers);
4373 let input = b.spec_operand();
4374 let output = b.spec_operand();
4375 let batch = b.spec_u32(1);
4376 let height = b.spec_u32(1);
4377 let width = b.spec_u32(1);
4378 let channels = b.spec_u32(1);
4379 let out_height = b.spec_u32(1);
4380 let out_width = b.spec_u32(1);
4381 let kernel_h = b.spec_u32(1);
4382 let kernel_w = b.spec_u32(1);
4383 let stride_h = b.spec_u32(1);
4384 let stride_w = b.spec_u32(1);
4385 let pad_top = b.spec_u32(0);
4386 let pad_left = b.spec_u32(0);
4387
4388 let (counter, stride) = b.grid_stride(workgroup);
4389 let u32_ty = b.u32_ty();
4390 let f32_ty = b.f32_ty();
4391 let acc_var = b.local(f32_ty);
4392 let kh_var = b.local(u32_ty);
4393 let kw_var = b.local(u32_ty);
4394 let count = b.imul(batch, out_height);
4395 let count = b.imul(count, out_width);
4396 let count = b.imul(count, channels);
4397 let zero = b.c_u32(0);
4398 let one = b.c_u32(1);
4399 let (scope, o) = b.begin_loop(counter, count);
4400 let c = b.umod(o, channels);
4401 let t = b.udiv(o, channels);
4402 let ow = b.umod(t, out_width);
4403 let t = b.udiv(t, out_width);
4404 let oh = b.umod(t, out_height);
4405 let nb = b.udiv(t, out_height);
4406 let neg_inf = b.c_f32(f32::NEG_INFINITY);
4407 b.store(acc_var, neg_inf);
4408 b.store(kh_var, zero);
4409 let row_origin = b.imul(oh, stride_h);
4410 let col_origin = b.imul(ow, stride_w);
4411 let (rows, kh) = b.begin_loop(kh_var, kernel_h);
4412 let padded_row = b.iadd(row_origin, kh);
4413 let row_in_low = b.uge(padded_row, pad_top);
4414 let ih = b.isub(padded_row, pad_top);
4415 let row_in_high = b.ult(ih, height);
4416 let row_ok = b.land(row_in_low, row_in_high);
4417 b.store(kw_var, zero);
4418 let (cols, kw) = b.begin_loop(kw_var, kernel_w);
4419 let padded_col = b.iadd(col_origin, kw);
4420 let col_in_low = b.uge(padded_col, pad_left);
4421 let iw = b.isub(padded_col, pad_left);
4422 let col_in_high = b.ult(iw, width);
4423 let col_ok = b.land(col_in_low, col_in_high);
4424 let ok = b.land(row_ok, col_ok);
4425 let index = b.imul(nb, height);
4427 let index = b.iadd(index, ih);
4428 let index = b.imul(index, width);
4429 let index = b.iadd(index, iw);
4430 let index = b.imul(index, channels);
4431 let index = b.iadd(index, c);
4432 let index = b.select_u32(ok, index, zero);
4433 let value = b.load_float(float, array, input, index);
4434 let acc = b.load(f32_ty, acc_var);
4435 let folded = b.apply_max(acc, value, nan_mode);
4436 let next = b.select_f32(ok, folded, acc);
4437 b.store(acc_var, next);
4438 b.end_loop(cols, kw_var, one);
4439 b.end_loop(rows, kh_var, one);
4440 let acc = b.load(f32_ty, acc_var);
4441 b.store_float(float, array, output, o, acc);
4442 b.end_loop(scope, counter, stride);
4443 b.end_main();
4444 b.finish([workgroup, 1, 1])
4445}
4446
4447fn assemble_cast(
4454 input: Storage,
4455 output_storage: Storage,
4456 workgroup: u32,
4457 buffers: u32,
4458) -> Vec<u32> {
4459 assemble_contiguous_lanes(input, output_storage, workgroup, buffers, |b, bits| {
4460 let value = b.widen_bits(input, bits);
4461 b.narrow_bits(output_storage, value)
4462 })
4463}
4464
4465fn assemble_contiguous_lanes(
4480 input: Storage,
4481 output: Storage,
4482 workgroup: u32,
4483 buffers: u32,
4484 convert: impl Fn(&mut Builder, Id) -> Id,
4485) -> Vec<u32> {
4486 let mut b = Builder::new();
4487 let array = b.buffer_array(buffers);
4488 let source = b.spec_operand();
4489 let destination = b.spec_operand();
4490 let count = b.spec_u32(1);
4491 let (counter, stride) = b.grid_stride(workgroup);
4492 let lanes = output.lanes();
4493 if lanes == 1 {
4494 let (scope, i) = b.begin_loop(counter, count);
4495 let bits = b.load_lane_group(input, array, source, i, 1)[0];
4496 let word = convert(&mut b, bits);
4497 b.store_word(array, destination, i, word);
4498 b.end_loop(scope, counter, stride);
4499 } else {
4500 let lanes_c = b.c_u32(lanes);
4501 let lanes_less_one = b.c_u32(lanes - 1);
4502 let words = b.iadd(count, lanes_less_one);
4503 let words = b.udiv(words, lanes_c);
4504 let (scope, w) = b.begin_loop(counter, words);
4505 let base = b.imul(w, lanes_c);
4506 let sources = b.load_lane_group(input, array, source, base, lanes);
4507 let bits: Vec<Id> = sources
4508 .iter()
4509 .map(|source_bits| convert(&mut b, *source_bits))
4510 .collect();
4511 let packed = b.pack_lanes(&bits);
4512 let end = b.iadd(base, lanes_c);
4513 let full = b.uge(count, end);
4514 b.if_then(full, |b| b.store_word(array, destination, w, packed));
4515 let partial = b.lnot(full);
4516 b.if_then(partial, |b| {
4517 for (lane, bits) in bits.iter().enumerate() {
4518 let lane = b.c_u32(lane as u32);
4519 let element = b.iadd(base, lane);
4520 let in_range = b.ult(element, count);
4521 b.if_then(in_range, |b| {
4522 b.store_lane_bits(output, array, destination, element, *bits);
4523 });
4524 }
4525 });
4526 b.end_loop(scope, counter, stride);
4527 }
4528 b.end_main();
4529 b.finish([workgroup, 1, 1])
4530}
4531
4532fn assemble_move(storage: Storage, contiguous: bool, workgroup: u32, buffers: u32) -> Vec<u32> {
4540 if contiguous {
4541 return assemble_contiguous_lanes(storage, storage, workgroup, buffers, |b, bits| {
4542 match storage {
4543 Storage::Byte => {
4545 let zero = b.c_u32(0);
4546 let one = b.c_u32(1);
4547 let set = b.ine(bits, zero);
4548 b.select_u32(set, one, zero)
4549 }
4550 Storage::Word | Storage::Half | Storage::Quarter(_) => bits,
4553 }
4554 });
4555 }
4556 let mut b = Builder::new();
4557 let array = b.buffer_array(buffers);
4558 let input = b.spec_operand();
4559 let output = b.spec_operand();
4560 let count = b.spec_u32(1);
4561 let dims = b.spec_dims();
4562 let in_strides = b.spec_strides();
4563 let in_offset = b.spec_u32(0);
4564 let out_strides = b.spec_strides();
4565 let out_offset = b.spec_u32(0);
4566
4567 let (counter, stride) = b.grid_stride(workgroup);
4568 let (scope, i) = b.begin_loop(counter, count);
4569 let indices = b.strided_indices(i, &dims, &[in_strides, out_strides]);
4570 let source = b.iadd(indices[0], in_offset);
4571 let destination = b.iadd(indices[1], out_offset);
4572 match storage {
4573 Storage::Word => {
4574 let word = b.load_word(array, input, source);
4575 b.store_word(array, output, destination, word);
4576 }
4577 Storage::Byte => {
4578 let value = b.load_bool(array, input, source);
4579 b.store_bool(array, output, destination, value);
4580 }
4581 Storage::Quarter(_) => {
4584 let bits = b.load_byte_bits(array, input, source);
4585 b.store_byte_bits(array, output, destination, bits);
4586 }
4587 Storage::Half => {
4590 let bits = b.load_half_bits(array, input, source);
4591 b.store_half_bits(array, output, destination, bits);
4592 }
4593 }
4594 b.end_loop(scope, counter, stride);
4595 b.end_main();
4596 b.finish([workgroup, 1, 1])
4597}
4598
4599#[cfg(test)]
4600mod tests {
4601 use super::*;
4602
4603 fn all_keys() -> Vec<KernelKey> {
4605 KernelKey::every_variant()
4606 }
4607
4608 fn check_well_formed(key: KernelKey, words: &[u32]) {
4611 assert_eq!(words[0], SPIRV_MAGIC, "{key:?}");
4612 assert_eq!(words[1], SPIRV_VERSION_1_3, "{key:?}");
4613 let bound = words[3];
4614 let mut cursor = 5;
4615 let mut opcodes = Vec::new();
4616 let mut spec_ids = Vec::new();
4617 let mut defined = std::collections::HashSet::new();
4618 while cursor < words.len() {
4619 let word_count = (words[cursor] >> 16) as usize;
4620 assert!(
4621 word_count >= 1,
4622 "{key:?}: zero-length instruction at {cursor}"
4623 );
4624 assert!(
4625 cursor + word_count <= words.len(),
4626 "{key:?}: instruction overruns"
4627 );
4628 let opcode = (words[cursor] & 0xffff) as u16;
4629 opcodes.push(opcode);
4630 if opcode == OP_DECORATE && words[cursor + 2] == DECORATION_SPEC_ID {
4631 spec_ids.push(words[cursor + 3]);
4632 }
4633 let result_id = match opcode {
4635 OP_TYPE_VOID
4636 | OP_TYPE_BOOL
4637 | OP_TYPE_INT
4638 | OP_TYPE_FLOAT
4639 | OP_TYPE_VECTOR
4640 | OP_TYPE_ARRAY
4641 | OP_TYPE_RUNTIME_ARRAY
4642 | OP_TYPE_STRUCT
4643 | OP_TYPE_POINTER
4644 | OP_TYPE_FUNCTION
4645 | OP_LABEL
4646 | OP_EXT_INST_IMPORT => Some(words[cursor + 1]),
4647 OP_CONSTANT | OP_SPEC_CONSTANT | OP_CONSTANT_FALSE | OP_VARIABLE | OP_LOAD
4648 | OP_ACCESS_CHAIN | OP_FUNCTION | OP_EXT_INST | OP_SELECT | OP_BITCAST => {
4649 Some(words[cursor + 2])
4650 }
4651 _ => None,
4652 };
4653 if let Some(id) = result_id {
4654 assert!(id < bound, "{key:?}: id {id} exceeds bound {bound}");
4655 assert!(defined.insert(id), "{key:?}: id {id} defined twice");
4656 }
4657 cursor += word_count;
4658 }
4659 assert_eq!(cursor, words.len(), "{key:?}");
4660 assert_eq!(opcodes[0], OP_CAPABILITY, "{key:?}");
4661 let import = opcodes
4662 .iter()
4663 .position(|op| *op == OP_EXT_INST_IMPORT)
4664 .unwrap();
4665 let memory = opcodes
4666 .iter()
4667 .position(|op| *op == OP_MEMORY_MODEL)
4668 .unwrap();
4669 let entry = opcodes.iter().position(|op| *op == OP_ENTRY_POINT).unwrap();
4670 let mode = opcodes
4671 .iter()
4672 .position(|op| *op == OP_EXECUTION_MODE)
4673 .unwrap();
4674 assert!(import < memory && memory < entry && entry < mode, "{key:?}");
4675 assert_eq!(*opcodes.last().unwrap(), OP_FUNCTION_END, "{key:?}");
4676 assert_eq!(
4677 opcodes.iter().filter(|op| **op == OP_FUNCTION).count(),
4678 1,
4679 "{key:?}"
4680 );
4681 let last_decorate = opcodes
4683 .iter()
4684 .rposition(|op| *op == OP_DECORATE || *op == OP_MEMBER_DECORATE)
4685 .unwrap();
4686 let first_type = opcodes.iter().position(|op| *op == OP_TYPE_VOID).unwrap();
4687 assert!(
4688 last_decorate < first_type,
4689 "{key:?}: annotation after types"
4690 );
4691 spec_ids.sort_unstable();
4693 let expected: Vec<u32> = (0..key.spec_constant_count()).collect();
4694 assert_eq!(spec_ids, expected, "{key:?}: specialization ids");
4695 let function_at = opcodes.iter().position(|op| *op == OP_FUNCTION).unwrap();
4697 let body = &opcodes[function_at + 2..];
4698 let first_non_variable = body.iter().position(|op| *op != OP_VARIABLE).unwrap();
4699 assert!(
4700 !body[first_non_variable..].contains(&OP_VARIABLE),
4701 "{key:?}: OpVariable after the entry block prologue"
4702 );
4703 }
4704
4705 #[test]
4706 fn every_kernel_variant_is_well_formed() {
4707 for key in all_keys() {
4708 let words = key.assemble();
4709 check_well_formed(key, &words);
4710 }
4711 }
4712
4713 #[test]
4714 fn assembly_is_deterministic() {
4715 for key in all_keys() {
4716 assert_eq!(key.assemble(), key.assemble(), "{key:?}");
4717 }
4718 }
4719
4720 #[test]
4721 fn spec_encoders_match_declared_counts() {
4722 let operand = Operand { buffer: 1, base: 2 };
4723 let strides = [[1; MAX_RANK]; 3];
4724 for key in all_keys() {
4725 let words = match key {
4726 KernelKey::Elementwise { op, broadcast, .. } => ElementwiseSpec {
4727 count: 8,
4728 inputs: &vec![operand; op.inputs().len()],
4729 output: operand,
4730 dims: [1; MAX_RANK],
4731 strides: &strides[..op.inputs().len()],
4732 clamp: matches!(op, ElementwiseOp::Clamp(_)).then_some([0, 0x3f80_0000]),
4733 }
4734 .words(broadcast),
4735 KernelKey::Reduce { .. } => reduce_spec(operand, operand, 2, 3, 1),
4736 KernelKey::Matmul { .. } | KernelKey::MatmulStream { .. } => {
4737 matmul_spec(operand, operand, operand, 1, 2, 3, 1)
4738 }
4739 KernelKey::Nvfp4Matmul { .. } => Nvfp4MatmulSpec {
4740 activation: operand,
4741 packed: &[operand],
4742 block_scales: &[operand],
4743 tensor_scale: operand,
4744 output: operand,
4745 m: 1,
4746 n: 2,
4747 k: 16,
4748 epilogue: 0,
4749 weight_mode: 0,
4750 }
4751 .words(),
4752 KernelKey::MaxPool { .. } => max_pool_spec(
4753 operand,
4754 operand,
4755 PoolGeometry {
4756 batch: 1,
4757 height: 4,
4758 width: 4,
4759 channels: 2,
4760 out_height: 2,
4761 out_width: 2,
4762 kernel: [2, 2],
4763 stride: [2, 2],
4764 pad_top: 0,
4765 pad_left: 0,
4766 },
4767 ),
4768 KernelKey::Cast { .. } => move_spec(
4771 operand,
4772 operand,
4773 MoveGeometry {
4774 count: 8,
4775 dims: [1; MAX_RANK],
4776 in_strides: [1; MAX_RANK],
4777 in_offset: 0,
4778 out_strides: [1; MAX_RANK],
4779 out_offset: 0,
4780 },
4781 true,
4782 ),
4783 KernelKey::Move { contiguous, .. } => move_spec(
4784 operand,
4785 operand,
4786 MoveGeometry {
4787 count: 6,
4788 dims: [1; MAX_RANK],
4789 in_strides: [1; MAX_RANK],
4790 in_offset: 0,
4791 out_strides: [1; MAX_RANK],
4792 out_offset: 0,
4793 },
4794 contiguous,
4795 ),
4796 };
4797 assert_eq!(
4798 words.len() as u32,
4799 key.spec_constant_count(),
4800 "{key:?}: encoder length"
4801 );
4802 }
4803 }
4804
4805 #[test]
4806 fn literal_strings_are_nul_terminated_and_word_padded() {
4807 assert_eq!(literal_string("main"), vec![0x6e69_616d, 0]);
4808 assert_eq!(literal_string("abc"), vec![0x0063_6261]);
4809 }
4810
4811 #[test]
4812 fn linear_workgroups_cover_and_cap() {
4813 assert_eq!(linear_workgroups(1, 64, 65_535), 1);
4814 assert_eq!(linear_workgroups(64, 64, 65_535), 1);
4815 assert_eq!(linear_workgroups(65, 64, 65_535), 2);
4816 assert_eq!(linear_workgroups(u32::MAX, 64, 65_535), 65_535);
4817 assert_eq!(linear_workgroups(0, 64, 65_535), 1);
4818 }
4819
4820 #[test]
4821 fn matmul_workgroups_cover_all_dimensions() {
4822 assert_eq!(matmul_block(16), 64);
4823 assert_eq!(matmul_block(8), 32);
4824 assert_eq!(MatmulGeometry::wide(16).shared_bytes(), 8192);
4825 assert_eq!(stream_matmul_shared_bytes(), 8192);
4826 assert_eq!(matmul_shared_bytes(8), 8192);
4827 for tile in [8, 16] {
4828 let geometry = MatmulGeometry::wide(tile);
4829 assert_eq!(geometry.invocations(), tile * tile, "{geometry:?}");
4830 assert_eq!(
4831 (geometry.block_m() * geometry.depth) % geometry.invocations(),
4832 0,
4833 "{geometry:?}"
4834 );
4835 assert_eq!(
4836 (geometry.block_n() * geometry.depth) % geometry.invocations(),
4837 0,
4838 "{geometry:?}"
4839 );
4840 }
4841 for lanes in [1, 2, 4] {
4843 let words = STREAM_COLUMNS / lanes;
4844 assert_eq!(STREAM_WORKGROUP % words, 0);
4845 assert_eq!((STREAM_ROWS * STREAM_COLUMNS) % STREAM_WORKGROUP, 0);
4846 }
4847 assert_eq!(stream_matmul_workgroups(4096, 2), [256, 1, 2]);
4848 assert_eq!(stream_matmul_workgroups(17, 1), [2, 1, 1]);
4849 assert_eq!(matmul_workgroups(1, 1, 1, 16), [1, 1, 1]);
4850 assert_eq!(matmul_workgroups(64, 64, 1, 16), [1, 1, 1]);
4851 assert_eq!(matmul_workgroups(65, 64, 1, 16), [1, 2, 1]);
4852 assert_eq!(matmul_workgroups(64, 65, 1, 16), [2, 1, 1]);
4853 assert_eq!(matmul_workgroups(32, 32, 3, 8), [1, 1, 3]);
4854 assert_eq!(matmul_workgroups(33, 8, 3, 8), [1, 2, 3]);
4855 }
4856
4857 #[test]
4858 fn elementwise_storage_tables_are_consistent() {
4859 for key in all_keys() {
4860 if let KernelKey::Elementwise { op, .. } = key {
4861 assert!(!op.inputs().is_empty());
4862 assert!(op.inputs().len() <= MAX_ELEMENTWISE_INPUTS);
4863 }
4864 }
4865 }
4866
4867 fn reference_f32_to_f16(value: f32) -> u16 {
4871 let value = f64::from(value);
4872 if value.is_nan() {
4873 return if value.is_sign_negative() {
4874 0xfe00
4875 } else {
4876 0x7e00
4877 };
4878 }
4879 let sign: u16 = if value.is_sign_negative() { 0x8000 } else { 0 };
4880 let magnitude = value.abs();
4881 if magnitude >= 65520.0 {
4882 return sign | 0x7c00;
4883 }
4884 if magnitude == 0.0 {
4885 return sign;
4886 }
4887 let bits = magnitude.to_bits();
4888 let biased = ((bits >> 52) & 0x7ff) as i32;
4889 let mantissa = (bits & ((1_u64 << 52) - 1)) | (1_u64 << 52);
4891 let exponent = biased - 1023;
4892 if exponent >= -14 {
4894 let kept = mantissa >> 42;
4896 let dropped = mantissa & ((1_u64 << 42) - 1);
4897 let half = 1_u64 << 41;
4898 let rounded = if dropped > half || (dropped == half && kept & 1 == 1) {
4899 kept + 1
4900 } else {
4901 kept
4902 };
4903 let (exponent, mantissa) = if rounded == 2048 {
4904 (exponent + 1, 1024)
4905 } else {
4906 (exponent, rounded)
4907 };
4908 let biased = exponent + 15;
4909 if biased >= 31 {
4910 return sign | 0x7c00;
4911 }
4912 return sign | ((biased as u16) << 10) | (mantissa - 1024) as u16;
4913 }
4914 let drop = (28 - exponent) as u32;
4916 if drop >= 64 {
4917 return sign;
4918 }
4919 let kept = mantissa >> drop;
4920 let dropped = mantissa & ((1_u64 << drop) - 1);
4921 let half = 1_u64 << (drop - 1);
4922 let rounded = if dropped > half || (dropped == half && kept & 1 == 1) {
4923 kept + 1
4924 } else {
4925 kept
4926 };
4927 sign | rounded as u16
4929 }
4930
4931 #[test]
4932 fn binary16_widening_is_exact_for_every_pattern() {
4933 for pattern in 0_u32..=0xffff {
4934 let bits = pattern as u16;
4935 let value = f16_to_f32(bits);
4936 let exponent = (bits >> 10) & 0x1f;
4937 let mantissa = bits & 0x3ff;
4938 match (exponent, mantissa) {
4939 (0, 0) => assert_eq!(value.to_bits(), (u32::from(bits) & 0x8000) << 16),
4940 (31, 0) => assert!(value.is_infinite(), "{bits:#06x}"),
4941 (31, _) => assert!(value.is_nan(), "{bits:#06x}"),
4942 _ => {
4943 assert_eq!(f32_to_f16_bits(value), bits, "{bits:#06x}");
4945 }
4946 }
4947 }
4948 }
4949
4950 #[test]
4951 fn binary16_narrowing_matches_the_reference_everywhere() {
4952 let mut cases: Vec<f32> = (0_u32..=0xffff)
4955 .map(|pattern| f16_to_f32(pattern as u16))
4956 .collect();
4957 for pattern in 0_u32..=0xfffe {
4958 let lo = f16_to_f32(pattern as u16);
4959 let hi = f16_to_f32((pattern + 1) as u16);
4960 if lo.is_finite() && hi.is_finite() {
4961 cases.push(f32::from_bits(
4962 lo.to_bits() + hi.to_bits().abs_diff(lo.to_bits()) / 2,
4963 ));
4964 }
4965 }
4966 let mut state = 0x243f_6a88_85a3_08d3_u64;
4967 for _ in 0..1_000_000 {
4968 state = state
4969 .wrapping_mul(6_364_136_223_846_793_005)
4970 .wrapping_add(1);
4971 cases.push(f32::from_bits((state >> 32) as u32));
4972 }
4973 for case in cases {
4974 assert_eq!(
4975 f32_to_f16_bits(case),
4976 reference_f32_to_f16(case),
4977 "{case:e} ({:#010x})",
4978 case.to_bits()
4979 );
4980 }
4981 }
4982}
4983
4984#[cfg(test)]
4985mod fp8_narrowing_tests {
4986 use super::*;
4987 use virtio_accel_tosa::{fp8e4m3_to_f32, fp8e5m2_to_f32};
4988
4989 fn nearest_by_search(format: Fp8Format, value: f32) -> u8 {
4993 let sign = if value.is_sign_negative() { 0x80 } else { 0x00 };
4994 let magnitude = value.abs();
4995 let mut best: Option<(f32, u8)> = None;
4996 for bits in 0..=0x7f_u8 {
4997 let candidate = fp8_decode(format, bits);
4998 if !candidate.is_finite() {
4999 continue;
5000 }
5001 let distance = (candidate - magnitude).abs();
5002 best = match best {
5003 None => Some((distance, bits)),
5004 Some((best_distance, _)) if distance < best_distance => Some((distance, bits)),
5005 Some((best_distance, best_bits))
5006 if distance == best_distance && bits & 1 == 0 && best_bits & 1 == 1 =>
5007 {
5008 Some((distance, bits))
5009 }
5010 other => other,
5011 };
5012 }
5013 sign | best.expect("a finite encoding exists").1
5014 }
5015
5016 #[test]
5017 fn narrowing_round_trips_every_encoding_exactly() {
5018 for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
5019 for bits in 0..=u8::MAX {
5020 let value = match format {
5021 Fp8Format::E4M3 => fp8e4m3_to_f32(bits),
5022 Fp8Format::E5M2 => fp8e5m2_to_f32(bits),
5023 };
5024 if !value.is_finite() {
5025 continue;
5026 }
5027 assert_eq!(
5028 f32_to_fp8_bits(format, value),
5029 bits,
5030 "{format:?}: {value} did not return to {bits:#04x}"
5031 );
5032 }
5033 }
5034 }
5035
5036 #[test]
5037 fn narrowing_matches_an_independent_nearest_search() {
5038 for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
5039 let max_finite = match format {
5040 Fp8Format::E4M3 => 448.0_f32,
5041 Fp8Format::E5M2 => 57344.0_f32,
5042 };
5043 let mut state = 0x1234_5678_u32;
5044 for index in 0..200_000 {
5045 state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
5047 let value = if index % 3 == 0 {
5048 let a = fp8_decode(format, (index % 256) as u8);
5049 let b = fp8_decode(format, ((index + 1) % 256) as u8);
5050 (a + b) * 0.5
5051 } else {
5052 let scaled = (state >> 8) as f32 / (1_u32 << 24) as f32;
5053 (scaled * 2.0 - 1.0) * max_finite * 1.2
5054 };
5055 if !value.is_finite() {
5056 continue;
5057 }
5058 let actual = f32_to_fp8_bits(format, value);
5059 if value.abs() > max_finite {
5060 continue; }
5062 assert_eq!(
5063 actual,
5064 nearest_by_search(format, value),
5065 "{format:?}: {value}"
5066 );
5067 }
5068 }
5069 }
5070
5071 fn fp8_decode(format: Fp8Format, bits: u8) -> f32 {
5072 match format {
5073 Fp8Format::E4M3 => fp8e4m3_to_f32(bits),
5074 Fp8Format::E5M2 => fp8e5m2_to_f32(bits),
5075 }
5076 }
5077
5078 #[test]
5079 fn overflow_policy_is_nan_for_e4m3_and_infinity_for_e5m2() {
5080 assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, 464.0), 0x7e);
5083 assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, 464.001), 0x7f);
5084 assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, -464.001), 0xff);
5085 assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, f32::INFINITY), 0x7f);
5086 assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, f32::NAN), 0x7f);
5087 assert_eq!(f32_to_fp8_bits(Fp8Format::E5M2, 1.0e9), 0x7c);
5088 assert_eq!(f32_to_fp8_bits(Fp8Format::E5M2, f32::NEG_INFINITY), 0xfc);
5089 assert_eq!(f32_to_fp8_bits(Fp8Format::E5M2, f32::NAN), 0x7e);
5090 assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, -0.0), 0x80);
5092 assert_eq!(f32_to_fp8_bits(Fp8Format::E5M2, 0.0), 0x00);
5093 }
5094}