1use super::{Brick, BrickAssertion, BrickBudget, BrickVerification};
37use std::time::Duration;
38
39#[derive(Debug, Clone, Copy, PartialEq, Eq)]
41pub enum TensorType {
42 F32,
44 F16,
46 I32,
48 U32,
50}
51
52impl TensorType {
53 #[must_use]
55 pub fn to_wgsl(&self) -> &'static str {
56 match self {
57 Self::F32 => "f32",
58 Self::F16 => "f16",
59 Self::I32 => "i32",
60 Self::U32 => "u32",
61 }
62 }
63
64 #[must_use]
66 pub fn to_rust(&self) -> &'static str {
67 match self {
68 Self::F32 => "f32",
69 Self::F16 => "half::f16",
70 Self::I32 => "i32",
71 Self::U32 => "u32",
72 }
73 }
74
75 #[must_use]
77 pub const fn byte_size(&self) -> usize {
78 match self {
79 Self::F32 | Self::I32 | Self::U32 => 4,
80 Self::F16 => 2,
81 }
82 }
83}
84
85#[derive(Debug, Clone)]
87pub struct TensorBinding {
88 pub name: String,
90 pub dtype: TensorType,
92 pub shape: Vec<u32>,
94 pub group: u32,
96 pub binding: u32,
98 pub read_only: bool,
100}
101
102impl TensorBinding {
103 #[must_use]
105 pub fn new(name: impl Into<String>, dtype: TensorType, shape: &[u32]) -> Self {
106 Self {
107 name: name.into(),
108 dtype,
109 shape: shape.to_vec(),
110 group: 0,
111 binding: 0,
112 read_only: true,
113 }
114 }
115
116 #[must_use]
118 pub fn at(mut self, group: u32, binding: u32) -> Self {
119 self.group = group;
120 self.binding = binding;
121 self
122 }
123
124 #[must_use]
126 pub fn writable(mut self) -> Self {
127 self.read_only = false;
128 self
129 }
130
131 #[must_use]
133 pub fn element_count(&self) -> u32 {
134 self.shape.iter().product()
135 }
136
137 #[must_use]
139 pub fn byte_size(&self) -> usize {
140 self.element_count() as usize * self.dtype.byte_size()
141 }
142
143 #[must_use]
145 pub fn to_wgsl_binding(&self) -> String {
146 let access = if self.read_only { "read" } else { "read_write" };
147 format!(
148 "@group({}) @binding({}) var<storage, {}> {}: array<{}>;",
149 self.group,
150 self.binding,
151 access,
152 self.name,
153 self.dtype.to_wgsl()
154 )
155 }
156}
157
158#[derive(Debug, Clone)]
160pub enum TileStrategy {
161 Simple2D {
163 tile_x: u32,
165 tile_y: u32,
167 },
168 Cooperative {
170 m: u32,
172 n: u32,
174 k: u32,
176 },
177 Streaming {
179 window: u32,
181 },
182 None,
184}
185
186impl TileStrategy {
187 #[must_use]
189 pub fn optimal_workgroup_size(&self) -> (u32, u32, u32) {
190 match self {
191 Self::Simple2D { tile_x, tile_y } => (*tile_x, *tile_y, 1),
192 Self::Cooperative { m, n, .. } => (*m, *n, 1),
193 Self::Streaming { window } => (*window, 1, 1),
194 Self::None => (64, 1, 1),
195 }
196 }
197}
198
199#[derive(Debug, Clone, Copy, PartialEq, Eq)]
201pub enum ElementwiseOp {
202 Log,
204 Exp,
206 Sqrt,
208 Abs,
210 Relu,
212 Sigmoid,
214 Tanh,
216 AddScalar(i32),
218 MulScalar(i32),
220 Clamp,
222}
223
224impl ElementwiseOp {
225 #[must_use]
227 pub fn to_wgsl_expr(&self, operand: &str) -> String {
228 match self {
229 Self::Log => format!("log({})", operand),
230 Self::Exp => format!("exp({})", operand),
231 Self::Sqrt => format!("sqrt({})", operand),
232 Self::Abs => format!("abs({})", operand),
233 Self::Relu => format!("max({}, 0.0)", operand),
234 Self::Sigmoid => format!("1.0 / (1.0 + exp(-{}))", operand),
235 Self::Tanh => format!("tanh({})", operand),
236 Self::AddScalar(s) => format!("({} + {}.0)", operand, s),
237 Self::MulScalar(s) => format!("({} * {}.0)", operand, s),
238 Self::Clamp => format!("clamp({}, 0.0, 1.0)", operand),
239 }
240 }
241}
242
243#[derive(Debug, Clone)]
245pub enum TileOp {
246 LoadShared {
248 src: String,
250 tile_size: (u32, u32),
252 },
253 Mma {
255 a: String,
257 b: String,
259 c: String,
261 },
262 Elementwise {
264 op: ElementwiseOp,
266 operands: Vec<String>,
268 output: Option<String>,
270 },
271 StoreShared {
273 dst: String,
275 },
276 Barrier,
278 Reduce {
280 kind: ReduceKind,
282 input: String,
284 output: String,
286 },
287}
288
289#[derive(Debug, Clone, Copy, PartialEq, Eq)]
291pub enum ReduceKind {
292 Sum,
294 Max,
296 Min,
298 Mean,
300}
301
302impl ReduceKind {
303 #[must_use]
305 pub fn identity(&self) -> &'static str {
306 match self {
307 Self::Sum | Self::Mean => "0.0",
308 Self::Max => "-3.402823e+38", Self::Min => "3.402823e+38", }
311 }
312
313 #[must_use]
315 pub fn combine_op(&self) -> &'static str {
316 match self {
317 Self::Sum | Self::Mean => "+",
318 Self::Max => "max",
319 Self::Min => "min",
320 }
321 }
322}
323
324#[derive(Debug, Clone)]
326pub struct ComputeBrick {
327 name: String,
329 workgroup_size: (u32, u32, u32),
331 inputs: Vec<TensorBinding>,
333 outputs: Vec<TensorBinding>,
335 tile_strategy: TileStrategy,
337 operations: Vec<TileOp>,
339 shared_memory: Vec<(String, TensorType, u32)>,
341}
342
343impl ComputeBrick {
344 #[must_use]
346 pub fn new(name: impl Into<String>) -> Self {
347 Self {
348 name: name.into(),
349 workgroup_size: (64, 1, 1),
350 inputs: Vec::new(),
351 outputs: Vec::new(),
352 tile_strategy: TileStrategy::None,
353 operations: Vec::new(),
354 shared_memory: Vec::new(),
355 }
356 }
357
358 #[must_use]
360 pub fn workgroup_size(mut self, x: u32, y: u32, z: u32) -> Self {
361 self.workgroup_size = (x, y, z);
362 self
363 }
364
365 #[must_use]
367 pub fn input(mut self, name: impl Into<String>, dtype: TensorType, shape: &[u32]) -> Self {
368 let binding_idx = self.inputs.len() as u32;
369 self.inputs
370 .push(TensorBinding::new(name, dtype, shape).at(0, binding_idx));
371 self
372 }
373
374 #[must_use]
376 pub fn output(mut self, name: impl Into<String>, dtype: TensorType, shape: &[u32]) -> Self {
377 let binding_idx = self.outputs.len() as u32;
378 self.outputs.push(
379 TensorBinding::new(name, dtype, shape)
380 .at(1, binding_idx)
381 .writable(),
382 );
383 self
384 }
385
386 #[must_use]
388 pub fn tile_strategy(mut self, strategy: TileStrategy) -> Self {
389 self.tile_strategy = strategy;
390 self
391 }
392
393 #[must_use]
395 pub fn op(mut self, operation: TileOp) -> Self {
396 self.operations.push(operation);
397 self
398 }
399
400 #[must_use]
402 pub fn shared(mut self, name: impl Into<String>, dtype: TensorType, size: u32) -> Self {
403 self.shared_memory.push((name.into(), dtype, size));
404 self
405 }
406
407 #[must_use]
409 pub fn to_wgsl(&self) -> String {
410 let mut wgsl = String::new();
411
412 wgsl.push_str(&format!(
414 "// {} Compute Shader\n",
415 to_pascal_case(&self.name)
416 ));
417 wgsl.push_str("// Generated by probar ComputeBrick - DO NOT EDIT MANUALLY\n\n");
418
419 for input in &self.inputs {
421 wgsl.push_str(&input.to_wgsl_binding());
422 wgsl.push('\n');
423 }
424
425 for output in &self.outputs {
427 wgsl.push_str(&output.to_wgsl_binding());
428 wgsl.push('\n');
429 }
430
431 wgsl.push('\n');
432
433 for (name, dtype, size) in &self.shared_memory {
435 wgsl.push_str(&format!(
436 "var<workgroup> {}: array<{}, {}>;\n",
437 name,
438 dtype.to_wgsl(),
439 size
440 ));
441 }
442
443 if !self.shared_memory.is_empty() {
444 wgsl.push('\n');
445 }
446
447 let (wg_x, wg_y, wg_z) = self.workgroup_size;
449 wgsl.push_str(&format!(
450 "@compute @workgroup_size({}, {}, {})\n",
451 wg_x, wg_y, wg_z
452 ));
453 wgsl.push_str("fn main(\n");
454 wgsl.push_str(" @builtin(global_invocation_id) global_id: vec3<u32>,\n");
455 wgsl.push_str(" @builtin(local_invocation_id) local_id: vec3<u32>,\n");
456 wgsl.push_str(" @builtin(workgroup_id) workgroup_id: vec3<u32>,\n");
457 wgsl.push_str(") {\n");
458
459 wgsl.push_str(" let gid = global_id.x + global_id.y * ");
461 wgsl.push_str(&format!("{}u;\n", wg_x));
462 wgsl.push_str(" let lid = local_id.x + local_id.y * ");
463 wgsl.push_str(&format!("{}u;\n\n", wg_x));
464
465 for op in &self.operations {
467 match op {
468 TileOp::LoadShared { src, tile_size: _ } => {
469 wgsl.push_str(&format!(" // Load from {} to shared memory\n", src));
470 wgsl.push_str(&format!(" let val_{} = {}[gid];\n", src, src));
471 }
472 TileOp::Elementwise {
473 op: elem_op,
474 operands,
475 output,
476 } => {
477 let input = &operands[0];
478 let out_name = output.as_ref().unwrap_or(input);
479 let input_val = format!("val_{}", input);
480 let expr = elem_op.to_wgsl_expr(&input_val);
481 wgsl.push_str(&format!(" let val_{} = {};\n", out_name, expr));
482 }
483 TileOp::StoreShared { dst } => {
484 wgsl.push_str(&format!(" // Store to {}\n", dst));
485 let val_name = if self.operations.iter().any(
487 |o| matches!(o, TileOp::Elementwise { output: Some(n), .. } if n == dst),
488 ) {
489 format!("val_{}", dst)
490 } else if let Some(input) = self.inputs.first() {
491 format!("val_{}", input.name)
492 } else {
493 "0.0".to_string()
494 };
495 wgsl.push_str(&format!(" {}[gid] = {};\n", dst, val_name));
496 }
497 TileOp::Barrier => {
498 wgsl.push_str(" workgroupBarrier();\n");
499 }
500 TileOp::Mma { a, b, c } => {
501 wgsl.push_str(&format!(" // Matrix multiply: {} = {} @ {}\n", c, a, b));
502 wgsl.push_str(" // TODO: Implement cooperative matrix\n");
503 }
504 TileOp::Reduce {
505 kind,
506 input,
507 output,
508 } => {
509 wgsl.push_str(&format!(
510 " // Reduce {} -> {} ({:?})\n",
511 input, output, kind
512 ));
513 }
514 }
515 }
516
517 wgsl.push_str("}\n");
518
519 wgsl
520 }
521
522 #[must_use]
524 pub fn to_rust_bindings(&self) -> String {
525 let mut rust = String::new();
526
527 rust.push_str(&format!(
529 "//! {} Compute Bindings\n",
530 to_pascal_case(&self.name)
531 ));
532 rust.push_str("//! Generated by probar ComputeBrick - DO NOT EDIT MANUALLY\n\n");
533 rust.push_str(
534 "use wgpu::{BindGroupLayout, BindGroupLayoutDescriptor, BindGroupLayoutEntry};\n",
535 );
536 rust.push_str("use wgpu::{ShaderStages, BufferBindingType, BindingType};\n\n");
537
538 let struct_name = to_pascal_case(&self.name);
539
540 rust.push_str(&format!("pub struct {}Compute {{\n", struct_name));
542 rust.push_str(" pub pipeline: wgpu::ComputePipeline,\n");
543 rust.push_str(" pub bind_group_layout: wgpu::BindGroupLayout,\n");
544 rust.push_str("}\n\n");
545
546 rust.push_str(&format!("impl {}Compute {{\n", struct_name));
548 rust.push_str(" pub const WORKGROUP_SIZE: (u32, u32, u32) = ");
549 rust.push_str(&format!("{:?};\n\n", self.workgroup_size));
550
551 rust.push_str(" pub const SHADER_SOURCE: &'static str = r#\"\n");
553 rust.push_str(&self.to_wgsl());
554 rust.push_str("\"#;\n\n");
555
556 rust.push_str(
558 " pub fn create_bind_group_layout(device: &wgpu::Device) -> BindGroupLayout {\n",
559 );
560 rust.push_str(" device.create_bind_group_layout(&BindGroupLayoutDescriptor {\n");
561 rust.push_str(&format!(
562 " label: Some(\"{} bind group layout\"),\n",
563 self.name
564 ));
565 rust.push_str(" entries: &[\n");
566
567 for input in &self.inputs {
568 rust.push_str(&format!(" // Input: {}\n", input.name));
569 rust.push_str(&format!(
570 " BindGroupLayoutEntry {{\n binding: {},\n visibility: ShaderStages::COMPUTE,\n ty: BindingType::Buffer {{\n ty: BufferBindingType::Storage {{ read_only: true }},\n has_dynamic_offset: false,\n min_binding_size: None,\n }},\n count: None,\n }},\n",
571 input.binding
572 ));
573 }
574
575 for output in &self.outputs {
576 rust.push_str(&format!(" // Output: {}\n", output.name));
577 rust.push_str(&format!(
578 " BindGroupLayoutEntry {{\n binding: {},\n visibility: ShaderStages::COMPUTE,\n ty: BindingType::Buffer {{\n ty: BufferBindingType::Storage {{ read_only: false }},\n has_dynamic_offset: false,\n min_binding_size: None,\n }},\n count: None,\n }},\n",
579 output.binding
580 ));
581 }
582
583 rust.push_str(" ],\n");
584 rust.push_str(" })\n");
585 rust.push_str(" }\n");
586 rust.push_str("}\n");
587
588 rust
589 }
590
591 #[must_use]
593 pub fn to_dispatch_js(&self) -> String {
594 let mut js = String::new();
595
596 js.push_str(&format!(
597 "// {} Compute Dispatch\n",
598 to_pascal_case(&self.name)
599 ));
600 js.push_str("// Generated by probar ComputeBrick - DO NOT EDIT MANUALLY\n\n");
601
602 let (wg_x, wg_y, wg_z) = self.workgroup_size;
603 js.push_str(&format!(
604 "const WORKGROUP_SIZE = [{}, {}, {}];\n\n",
605 wg_x, wg_y, wg_z
606 ));
607
608 js.push_str(&format!(
609 "async function dispatch{}(device, inputs, outputs) {{\n",
610 to_pascal_case(&self.name)
611 ));
612
613 js.push_str(" // Create shader module\n");
614 js.push_str(" const shaderModule = device.createShaderModule({\n");
615 js.push_str(&format!(" label: '{} shader',\n", self.name));
616 js.push_str(" code: SHADER_SOURCE,\n");
617 js.push_str(" });\n\n");
618
619 js.push_str(" // Calculate dispatch size\n");
620 if let Some(output) = self.outputs.first() {
621 let total_size = output.element_count();
622 js.push_str(&format!(" const totalElements = {};\n", total_size));
623 js.push_str(&format!(
624 " const numWorkgroups = Math.ceil(totalElements / {});\n\n",
625 wg_x * wg_y * wg_z
626 ));
627 }
628
629 js.push_str(" // Dispatch\n");
630 js.push_str(" const commandEncoder = device.createCommandEncoder();\n");
631 js.push_str(" const passEncoder = commandEncoder.beginComputePass();\n");
632 js.push_str(" passEncoder.setPipeline(pipeline);\n");
633 js.push_str(" passEncoder.setBindGroup(0, bindGroup);\n");
634 js.push_str(" passEncoder.dispatchWorkgroups(numWorkgroups, 1, 1);\n");
635 js.push_str(" passEncoder.end();\n");
636 js.push_str(" device.queue.submit([commandEncoder.finish()]);\n");
637 js.push_str("}\n");
638
639 js
640 }
641
642 #[must_use]
644 pub fn name(&self) -> &str {
645 &self.name
646 }
647
648 #[must_use]
650 pub fn get_workgroup_size(&self) -> (u32, u32, u32) {
651 self.workgroup_size
652 }
653
654 #[must_use]
656 pub fn inputs(&self) -> &[TensorBinding] {
657 &self.inputs
658 }
659
660 #[must_use]
662 pub fn outputs(&self) -> &[TensorBinding] {
663 &self.outputs
664 }
665}
666
667impl Brick for ComputeBrick {
668 fn brick_name(&self) -> &'static str {
669 "ComputeBrick"
670 }
671
672 fn assertions(&self) -> &[BrickAssertion] {
673 &[]
674 }
675
676 fn budget(&self) -> BrickBudget {
677 BrickBudget::uniform(100)
679 }
680
681 fn verify(&self) -> BrickVerification {
682 let mut passed = Vec::new();
683 let mut failed = Vec::new();
684
685 let (x, y, z) = self.workgroup_size;
687 if x * y * z > 1024 {
688 failed.push((
689 BrickAssertion::Custom {
690 name: "workgroup_size_valid".into(),
691 validator_id: 1,
692 },
693 format!(
694 "Workgroup size {}x{}x{}={} exceeds maximum 1024",
695 x,
696 y,
697 z,
698 x * y * z
699 ),
700 ));
701 } else {
702 passed.push(BrickAssertion::Custom {
703 name: "workgroup_size_valid".into(),
704 validator_id: 1,
705 });
706 }
707
708 if self.inputs.is_empty() {
710 failed.push((
711 BrickAssertion::Custom {
712 name: "has_inputs".into(),
713 validator_id: 2,
714 },
715 "ComputeBrick has no input tensors".into(),
716 ));
717 } else {
718 passed.push(BrickAssertion::Custom {
719 name: "has_inputs".into(),
720 validator_id: 2,
721 });
722 }
723
724 if self.outputs.is_empty() {
725 failed.push((
726 BrickAssertion::Custom {
727 name: "has_outputs".into(),
728 validator_id: 3,
729 },
730 "ComputeBrick has no output tensors".into(),
731 ));
732 } else {
733 passed.push(BrickAssertion::Custom {
734 name: "has_outputs".into(),
735 validator_id: 3,
736 });
737 }
738
739 let tensor_names: Vec<_> = self
741 .inputs
742 .iter()
743 .chain(self.outputs.iter())
744 .map(|t| t.name.as_str())
745 .collect();
746
747 for op in &self.operations {
748 match op {
749 TileOp::LoadShared { src, .. } => {
750 if !tensor_names.contains(&src.as_str()) {
751 failed.push((
752 BrickAssertion::Custom {
753 name: "tensor_exists".into(),
754 validator_id: 4,
755 },
756 format!("LoadShared references unknown tensor: {}", src),
757 ));
758 }
759 }
760 TileOp::StoreShared { dst } if !tensor_names.contains(&dst.as_str()) => {
766 failed.push((
767 BrickAssertion::Custom {
768 name: "tensor_exists".into(),
769 validator_id: 4,
770 },
771 format!("StoreShared references unknown tensor: {}", dst),
772 ));
773 }
774 _ => {}
775 }
776 }
777
778 if failed.is_empty() {
779 passed.push(BrickAssertion::Custom {
780 name: "compute_brick_valid".into(),
781 validator_id: 5,
782 });
783 }
784
785 BrickVerification {
786 passed,
787 failed,
788 verification_time: Duration::from_micros(100),
789 }
790 }
791
792 fn to_html(&self) -> String {
793 String::new()
795 }
796
797 fn to_css(&self) -> String {
798 String::new()
800 }
801}
802
803fn to_pascal_case(s: &str) -> String {
805 let mut result = String::new();
806 let mut capitalize_next = true;
807
808 for c in s.chars() {
809 if c == '_' || c == '-' || c == ' ' {
810 capitalize_next = true;
811 } else if capitalize_next {
812 result.push(c.to_ascii_uppercase());
813 capitalize_next = false;
814 } else {
815 result.push(c);
816 }
817 }
818
819 result
820}
821
822#[cfg(test)]
823#[allow(clippy::unwrap_used, clippy::expect_used)]
824mod tests {
825 use super::*;
826
827 #[test]
828 fn test_compute_brick_basic() {
829 let brick = ComputeBrick::new("test")
830 .workgroup_size(256, 1, 1)
831 .input("audio", TensorType::F32, &[1024])
832 .output("mel", TensorType::F32, &[80, 100]);
833
834 assert_eq!(brick.name(), "test");
835 assert_eq!(brick.get_workgroup_size(), (256, 1, 1));
836 assert_eq!(brick.inputs().len(), 1);
837 assert_eq!(brick.outputs().len(), 1);
838 }
839
840 #[test]
841 fn test_compute_brick_wgsl_generation() {
842 let brick = ComputeBrick::new("log-transform")
843 .workgroup_size(64, 1, 1)
844 .input("input", TensorType::F32, &[1024])
845 .output("output", TensorType::F32, &[1024])
846 .op(TileOp::LoadShared {
847 src: "input".into(),
848 tile_size: (64, 1),
849 })
850 .op(TileOp::Elementwise {
851 op: ElementwiseOp::Log,
852 operands: vec!["input".into()],
853 output: Some("output".into()),
854 })
855 .op(TileOp::StoreShared {
856 dst: "output".into(),
857 });
858
859 let wgsl = brick.to_wgsl();
860
861 assert!(wgsl.contains("@compute @workgroup_size(64, 1, 1)"));
862 assert!(wgsl.contains("fn main("));
863 assert!(wgsl.contains("log("));
864 assert!(wgsl.contains("Generated by probar"));
865 }
866
867 #[test]
868 fn test_compute_brick_verification() {
869 let brick = ComputeBrick::new("test")
870 .workgroup_size(256, 1, 1)
871 .input("input", TensorType::F32, &[1024])
872 .output("output", TensorType::F32, &[1024]);
873
874 let result = brick.verify();
875 assert!(result.is_valid());
876 }
877
878 #[test]
879 fn test_compute_brick_verification_fails_no_inputs() {
880 let brick = ComputeBrick::new("test").workgroup_size(256, 1, 1).output(
881 "output",
882 TensorType::F32,
883 &[1024],
884 );
885
886 let result = brick.verify();
887 assert!(!result.is_valid());
888 }
889
890 #[test]
891 fn test_compute_brick_verification_fails_large_workgroup() {
892 let brick = ComputeBrick::new("test")
893 .workgroup_size(1024, 2, 1) .input("input", TensorType::F32, &[1024])
895 .output("output", TensorType::F32, &[1024]);
896
897 let result = brick.verify();
898 assert!(!result.is_valid());
899 }
900
901 #[test]
902 fn test_tensor_binding() {
903 let binding = TensorBinding::new("audio", TensorType::F32, &[1024, 80])
904 .at(0, 1)
905 .writable();
906
907 assert_eq!(binding.name, "audio");
908 assert_eq!(binding.element_count(), 1024 * 80);
909 assert_eq!(binding.byte_size(), 1024 * 80 * 4);
910 assert!(!binding.read_only);
911 }
912
913 #[test]
914 fn test_tensor_type_wgsl() {
915 assert_eq!(TensorType::F32.to_wgsl(), "f32");
916 assert_eq!(TensorType::F16.to_wgsl(), "f16");
917 assert_eq!(TensorType::I32.to_wgsl(), "i32");
918 assert_eq!(TensorType::U32.to_wgsl(), "u32");
919 }
920
921 #[test]
922 fn test_elementwise_ops() {
923 assert_eq!(ElementwiseOp::Log.to_wgsl_expr("x"), "log(x)");
924 assert_eq!(ElementwiseOp::Exp.to_wgsl_expr("x"), "exp(x)");
925 assert_eq!(ElementwiseOp::Relu.to_wgsl_expr("x"), "max(x, 0.0)");
926 assert_eq!(ElementwiseOp::AddScalar(5).to_wgsl_expr("x"), "(x + 5.0)");
927 }
928
929 #[test]
930 fn test_rust_bindings_generation() {
931 let brick = ComputeBrick::new("mel-transform")
932 .workgroup_size(256, 1, 1)
933 .input("audio", TensorType::F32, &[1024])
934 .output("mel", TensorType::F32, &[80]);
935
936 let rust = brick.to_rust_bindings();
937
938 assert!(rust.contains("pub struct MelTransformCompute"));
939 assert!(rust.contains("WORKGROUP_SIZE"));
940 assert!(rust.contains("SHADER_SOURCE"));
941 assert!(rust.contains("create_bind_group_layout"));
942 }
943
944 #[test]
945 fn test_js_dispatch_generation() {
946 let brick = ComputeBrick::new("fft")
947 .workgroup_size(64, 1, 1)
948 .input("signal", TensorType::F32, &[512])
949 .output("spectrum", TensorType::F32, &[512]);
950
951 let js = brick.to_dispatch_js();
952
953 assert!(js.contains("async function dispatchFft"));
954 assert!(js.contains("WORKGROUP_SIZE"));
955 assert!(js.contains("dispatchWorkgroups"));
956 }
957
958 #[test]
959 fn test_tile_strategy_workgroup_size() {
960 let simple = TileStrategy::Simple2D {
961 tile_x: 16,
962 tile_y: 16,
963 };
964 assert_eq!(simple.optimal_workgroup_size(), (16, 16, 1));
965
966 let coop = TileStrategy::Cooperative { m: 8, n: 8, k: 4 };
967 assert_eq!(coop.optimal_workgroup_size(), (8, 8, 1));
968
969 let streaming = TileStrategy::Streaming { window: 32 };
970 assert_eq!(streaming.optimal_workgroup_size(), (32, 1, 1));
971 }
972
973 #[test]
978 fn test_tensor_type_rust() {
979 assert_eq!(TensorType::F32.to_rust(), "f32");
980 assert_eq!(TensorType::F16.to_rust(), "half::f16");
981 assert_eq!(TensorType::I32.to_rust(), "i32");
982 assert_eq!(TensorType::U32.to_rust(), "u32");
983 }
984
985 #[test]
986 fn test_tensor_type_byte_size() {
987 assert_eq!(TensorType::F32.byte_size(), 4);
988 assert_eq!(TensorType::F16.byte_size(), 2);
989 assert_eq!(TensorType::I32.byte_size(), 4);
990 assert_eq!(TensorType::U32.byte_size(), 4);
991 }
992
993 #[test]
994 fn test_tensor_type_clone() {
995 let t = TensorType::F32;
996 let cloned = t;
997 assert_eq!(t, cloned);
998 }
999
1000 #[test]
1001 fn test_tensor_binding_default_values() {
1002 let binding = TensorBinding::new("test", TensorType::I32, &[10, 20]);
1003 assert_eq!(binding.group, 0);
1004 assert_eq!(binding.binding, 0);
1005 assert!(binding.read_only);
1006 }
1007
1008 #[test]
1009 fn test_tensor_binding_to_wgsl_binding_read_only() {
1010 let binding = TensorBinding::new("data", TensorType::F32, &[100]).at(1, 2);
1011 let wgsl = binding.to_wgsl_binding();
1012 assert!(wgsl.contains("@group(1) @binding(2)"));
1013 assert!(wgsl.contains("var<storage, read>"));
1014 assert!(wgsl.contains("data"));
1015 assert!(wgsl.contains("f32"));
1016 }
1017
1018 #[test]
1019 fn test_tensor_binding_to_wgsl_binding_read_write() {
1020 let binding = TensorBinding::new("output", TensorType::U32, &[50])
1021 .at(0, 0)
1022 .writable();
1023 let wgsl = binding.to_wgsl_binding();
1024 assert!(wgsl.contains("var<storage, read_write>"));
1025 }
1026
1027 #[test]
1028 fn test_tensor_binding_clone() {
1029 let binding = TensorBinding::new("test", TensorType::F32, &[1, 2, 3])
1030 .at(1, 2)
1031 .writable();
1032 let cloned = binding.clone();
1033 assert_eq!(binding.name, cloned.name);
1034 assert_eq!(binding.shape, cloned.shape);
1035 assert_eq!(binding.read_only, cloned.read_only);
1036 }
1037
1038 #[test]
1039 fn test_tile_strategy_none() {
1040 let strategy = TileStrategy::None;
1041 assert_eq!(strategy.optimal_workgroup_size(), (64, 1, 1));
1042 }
1043
1044 #[test]
1045 fn test_tile_strategy_clone() {
1046 let strategy = TileStrategy::Simple2D {
1047 tile_x: 8,
1048 tile_y: 8,
1049 };
1050 let cloned = strategy;
1051 assert!(matches!(
1052 cloned,
1053 TileStrategy::Simple2D {
1054 tile_x: 8,
1055 tile_y: 8
1056 }
1057 ));
1058 }
1059
1060 #[test]
1061 fn test_elementwise_op_sqrt() {
1062 assert_eq!(ElementwiseOp::Sqrt.to_wgsl_expr("val"), "sqrt(val)");
1063 }
1064
1065 #[test]
1066 fn test_elementwise_op_abs() {
1067 assert_eq!(ElementwiseOp::Abs.to_wgsl_expr("v"), "abs(v)");
1068 }
1069
1070 #[test]
1071 fn test_elementwise_op_sigmoid() {
1072 assert_eq!(
1073 ElementwiseOp::Sigmoid.to_wgsl_expr("x"),
1074 "1.0 / (1.0 + exp(-x))"
1075 );
1076 }
1077
1078 #[test]
1079 fn test_elementwise_op_tanh() {
1080 assert_eq!(ElementwiseOp::Tanh.to_wgsl_expr("x"), "tanh(x)");
1081 }
1082
1083 #[test]
1084 fn test_elementwise_op_mul_scalar() {
1085 assert_eq!(ElementwiseOp::MulScalar(3).to_wgsl_expr("y"), "(y * 3.0)");
1086 assert_eq!(ElementwiseOp::MulScalar(-2).to_wgsl_expr("x"), "(x * -2.0)");
1087 }
1088
1089 #[test]
1090 fn test_elementwise_op_clamp() {
1091 assert_eq!(ElementwiseOp::Clamp.to_wgsl_expr("x"), "clamp(x, 0.0, 1.0)");
1092 }
1093
1094 #[test]
1095 fn test_elementwise_op_eq() {
1096 assert_eq!(ElementwiseOp::Log, ElementwiseOp::Log);
1097 assert_ne!(ElementwiseOp::Log, ElementwiseOp::Exp);
1098 assert_eq!(ElementwiseOp::AddScalar(5), ElementwiseOp::AddScalar(5));
1099 assert_ne!(ElementwiseOp::AddScalar(5), ElementwiseOp::AddScalar(6));
1100 }
1101
1102 #[test]
1103 fn test_reduce_kind_identity() {
1104 assert_eq!(ReduceKind::Sum.identity(), "0.0");
1105 assert_eq!(ReduceKind::Mean.identity(), "0.0");
1106 assert_eq!(ReduceKind::Max.identity(), "-3.402823e+38");
1107 assert_eq!(ReduceKind::Min.identity(), "3.402823e+38");
1108 }
1109
1110 #[test]
1111 fn test_reduce_kind_combine_op() {
1112 assert_eq!(ReduceKind::Sum.combine_op(), "+");
1113 assert_eq!(ReduceKind::Mean.combine_op(), "+");
1114 assert_eq!(ReduceKind::Max.combine_op(), "max");
1115 assert_eq!(ReduceKind::Min.combine_op(), "min");
1116 }
1117
1118 #[test]
1119 fn test_reduce_kind_eq() {
1120 assert_eq!(ReduceKind::Sum, ReduceKind::Sum);
1121 assert_ne!(ReduceKind::Sum, ReduceKind::Max);
1122 }
1123
1124 #[test]
1125 fn test_tile_op_load_shared() {
1126 let op = TileOp::LoadShared {
1127 src: "audio".into(),
1128 tile_size: (32, 32),
1129 };
1130 match op {
1131 TileOp::LoadShared { src, tile_size } => {
1132 assert_eq!(src, "audio");
1133 assert_eq!(tile_size, (32, 32));
1134 }
1135 _ => panic!("Expected LoadShared"),
1136 }
1137 }
1138
1139 #[test]
1140 fn test_tile_op_mma() {
1141 let op = TileOp::Mma {
1142 a: "A".into(),
1143 b: "B".into(),
1144 c: "C".into(),
1145 };
1146 match op {
1147 TileOp::Mma { a, b, c } => {
1148 assert_eq!(a, "A");
1149 assert_eq!(b, "B");
1150 assert_eq!(c, "C");
1151 }
1152 _ => panic!("Expected Mma"),
1153 }
1154 }
1155
1156 #[test]
1157 fn test_tile_op_reduce() {
1158 let op = TileOp::Reduce {
1159 kind: ReduceKind::Max,
1160 input: "values".into(),
1161 output: "max_val".into(),
1162 };
1163 match op {
1164 TileOp::Reduce {
1165 kind,
1166 input,
1167 output,
1168 } => {
1169 assert_eq!(kind, ReduceKind::Max);
1170 assert_eq!(input, "values");
1171 assert_eq!(output, "max_val");
1172 }
1173 _ => panic!("Expected Reduce"),
1174 }
1175 }
1176
1177 #[test]
1178 fn test_tile_op_barrier() {
1179 let op = TileOp::Barrier;
1180 assert!(matches!(op, TileOp::Barrier));
1181 }
1182
1183 #[test]
1184 fn test_tile_op_clone() {
1185 let op = TileOp::Elementwise {
1186 op: ElementwiseOp::Relu,
1187 operands: vec!["x".into(), "y".into()],
1188 output: Some("z".into()),
1189 };
1190 let cloned = op;
1191 assert!(matches!(cloned, TileOp::Elementwise { .. }));
1192 }
1193
1194 #[test]
1195 fn test_compute_brick_tile_strategy() {
1196 let brick = ComputeBrick::new("test").tile_strategy(TileStrategy::Cooperative {
1197 m: 16,
1198 n: 16,
1199 k: 8,
1200 });
1201
1202 assert_eq!(brick.name(), "test");
1204 }
1205
1206 #[test]
1207 fn test_compute_brick_shared_memory() {
1208 let brick = ComputeBrick::new("test")
1209 .shared("tile_a", TensorType::F32, 256)
1210 .shared("tile_b", TensorType::F32, 128);
1211
1212 let wgsl = brick.to_wgsl();
1213 assert!(wgsl.contains("var<workgroup> tile_a"));
1214 assert!(wgsl.contains("var<workgroup> tile_b"));
1215 }
1216
1217 #[test]
1218 fn test_compute_brick_verification_no_outputs() {
1219 let brick = ComputeBrick::new("test").input("input", TensorType::F32, &[1024]);
1220
1221 let result = brick.verify();
1222 assert!(!result.is_valid());
1223 }
1224
1225 #[test]
1226 fn test_compute_brick_verification_invalid_load_tensor() {
1227 let brick = ComputeBrick::new("test")
1228 .input("input", TensorType::F32, &[1024])
1229 .output("output", TensorType::F32, &[1024])
1230 .op(TileOp::LoadShared {
1231 src: "nonexistent".into(),
1232 tile_size: (64, 1),
1233 });
1234
1235 let result = brick.verify();
1236 assert!(!result.is_valid());
1237 }
1238
1239 #[test]
1240 fn test_compute_brick_verification_invalid_store_tensor() {
1241 let brick = ComputeBrick::new("test")
1242 .input("input", TensorType::F32, &[1024])
1243 .output("output", TensorType::F32, &[1024])
1244 .op(TileOp::StoreShared {
1245 dst: "nonexistent".into(),
1246 });
1247
1248 let result = brick.verify();
1249 assert!(!result.is_valid());
1250 }
1251
1252 #[test]
1253 fn test_compute_brick_wgsl_barrier() {
1254 let brick = ComputeBrick::new("test")
1255 .input("input", TensorType::F32, &[64])
1256 .output("output", TensorType::F32, &[64])
1257 .op(TileOp::Barrier);
1258
1259 let wgsl = brick.to_wgsl();
1260 assert!(wgsl.contains("workgroupBarrier()"));
1261 }
1262
1263 #[test]
1264 fn test_compute_brick_wgsl_mma() {
1265 let brick = ComputeBrick::new("matmul")
1266 .input("A", TensorType::F32, &[64, 64])
1267 .input("B", TensorType::F32, &[64, 64])
1268 .output("C", TensorType::F32, &[64, 64])
1269 .op(TileOp::Mma {
1270 a: "A".into(),
1271 b: "B".into(),
1272 c: "C".into(),
1273 });
1274
1275 let wgsl = brick.to_wgsl();
1276 assert!(wgsl.contains("Matrix multiply"));
1277 }
1278
1279 #[test]
1280 fn test_compute_brick_wgsl_reduce() {
1281 let brick = ComputeBrick::new("reduce")
1282 .input("values", TensorType::F32, &[1024])
1283 .output("result", TensorType::F32, &[1])
1284 .op(TileOp::Reduce {
1285 kind: ReduceKind::Sum,
1286 input: "values".into(),
1287 output: "result".into(),
1288 });
1289
1290 let wgsl = brick.to_wgsl();
1291 assert!(wgsl.contains("Reduce"));
1292 }
1293
1294 #[test]
1295 fn test_compute_brick_wgsl_elementwise_no_output() {
1296 let brick = ComputeBrick::new("test")
1297 .input("x", TensorType::F32, &[64])
1298 .output("y", TensorType::F32, &[64])
1299 .op(TileOp::LoadShared {
1300 src: "x".into(),
1301 tile_size: (64, 1),
1302 })
1303 .op(TileOp::Elementwise {
1304 op: ElementwiseOp::Log,
1305 operands: vec!["x".into()],
1306 output: None, });
1308
1309 let wgsl = brick.to_wgsl();
1310 assert!(wgsl.contains("log(val_x)"));
1311 }
1312
1313 #[test]
1314 fn test_compute_brick_wgsl_store_fallback() {
1315 let brick = ComputeBrick::new("test")
1316 .input("input", TensorType::F32, &[64])
1317 .output("output", TensorType::F32, &[64])
1318 .op(TileOp::LoadShared {
1319 src: "input".into(),
1320 tile_size: (64, 1),
1321 })
1322 .op(TileOp::StoreShared {
1323 dst: "output".into(),
1324 });
1325
1326 let wgsl = brick.to_wgsl();
1327 assert!(wgsl.contains("output[gid]"));
1328 }
1329
1330 #[test]
1331 fn test_compute_brick_implements_brick() {
1332 let brick = ComputeBrick::new("test")
1333 .input("in", TensorType::F32, &[32])
1334 .output("out", TensorType::F32, &[32]);
1335
1336 assert_eq!(brick.brick_name(), "ComputeBrick");
1337 assert!(brick.assertions().is_empty());
1338 assert_eq!(brick.budget().total_ms, 100);
1339 assert!(brick.to_html().is_empty());
1340 assert!(brick.to_css().is_empty());
1341 }
1342
1343 #[test]
1344 fn test_to_pascal_case_variants() {
1345 assert_eq!(to_pascal_case("simple"), "Simple");
1346 assert_eq!(to_pascal_case("two_words"), "TwoWords");
1347 assert_eq!(to_pascal_case("three-part-name"), "ThreePartName");
1348 assert_eq!(to_pascal_case("mixed_style-here"), "MixedStyleHere");
1349 assert_eq!(to_pascal_case("with space"), "WithSpace");
1350 }
1351
1352 #[test]
1353 fn test_compute_brick_multiple_inputs() {
1354 let brick = ComputeBrick::new("multi")
1355 .input("a", TensorType::F32, &[100])
1356 .input("b", TensorType::I32, &[100])
1357 .input("c", TensorType::U32, &[100])
1358 .output("result", TensorType::F32, &[100]);
1359
1360 assert_eq!(brick.inputs().len(), 3);
1361 assert_eq!(brick.inputs()[0].binding, 0);
1362 assert_eq!(brick.inputs()[1].binding, 1);
1363 assert_eq!(brick.inputs()[2].binding, 2);
1364 }
1365
1366 #[test]
1367 fn test_compute_brick_multiple_outputs() {
1368 let brick = ComputeBrick::new("multi_out")
1369 .input("in", TensorType::F32, &[50])
1370 .output("out1", TensorType::F32, &[50])
1371 .output("out2", TensorType::F32, &[25]);
1372
1373 assert_eq!(brick.outputs().len(), 2);
1374 assert_eq!(brick.outputs()[0].binding, 0);
1375 assert_eq!(brick.outputs()[1].binding, 1);
1376 assert_eq!(brick.outputs()[0].group, 1);
1377 assert_eq!(brick.outputs()[1].group, 1);
1378 }
1379
1380 #[test]
1381 fn test_compute_brick_clone() {
1382 let brick = ComputeBrick::new("test")
1383 .workgroup_size(128, 4, 1)
1384 .input("in", TensorType::F16, &[256])
1385 .output("out", TensorType::F16, &[256])
1386 .shared("cache", TensorType::F16, 512);
1387
1388 let cloned = brick.clone();
1389 assert_eq!(brick.name(), cloned.name());
1390 assert_eq!(brick.get_workgroup_size(), cloned.get_workgroup_size());
1391 }
1392
1393 #[test]
1394 fn test_js_dispatch_no_outputs() {
1395 let brick = ComputeBrick::new("no_out").input("in", TensorType::F32, &[10]);
1396
1397 let js = brick.to_dispatch_js();
1398 assert!(js.contains("dispatchNoOut"));
1400 }
1401
1402 #[test]
1403 fn test_rust_bindings_multiple_io() {
1404 let brick = ComputeBrick::new("complex")
1405 .input("in1", TensorType::F32, &[100])
1406 .input("in2", TensorType::I32, &[50])
1407 .output("out1", TensorType::F32, &[100])
1408 .output("out2", TensorType::U32, &[25]);
1409
1410 let rust = brick.to_rust_bindings();
1411 assert!(rust.contains("Input: in1"));
1412 assert!(rust.contains("Input: in2"));
1413 assert!(rust.contains("Output: out1"));
1414 assert!(rust.contains("Output: out2"));
1415 }
1416
1417 #[test]
1418 fn test_tensor_binding_empty_shape() {
1419 let binding = TensorBinding::new("scalar", TensorType::F32, &[]);
1420 assert_eq!(binding.element_count(), 1); assert_eq!(binding.byte_size(), 4);
1422 }
1423}