1#![cfg_attr(not(va_vulkan), allow(dead_code))]
23
24use std::collections::HashMap;
25use std::fmt;
26
27use virtio_accel_tosa::{
28 AnalysisError, AnalyzedValueKind, CapabilityDescriptor, DType, DTypeCapability,
29 DTypeConstraints, Error as ParseError, ExtensionSet, GraphCapabilities, Level,
30 NanPropagationMode, Op, OpAttributes, OperatorCapability, OperatorConstraints, OperatorId,
31 OptimizationHints, ProfileSet, RuntimeCondition, RuntimeConditionSupport, Target, TosaAnalysis,
32 ValueId, ValueRoles, Version, parse,
33};
34
35use crate::shader::{
36 ElementwiseOp, ElementwiseSpec, Fp8Format, MAX_RANK, MoveGeometry, NanMode, Operand,
37 PoolGeometry, ReduceOp, Storage, matmul_spec, max_pool_spec, move_spec, reduce_spec,
38};
39
40pub const VULKAN_TOSA_TARGET: Target = Target::new(
42 Version::TOSA_1_0,
43 ProfileSet::FLOATING_POINT,
44 Level::Level8K,
45 ExtensionSet::NONE,
46);
47
48pub const VULKAN_TOSA_FP8_TARGET: Target = Target::new(
55 Version::TOSA_1_0,
56 ProfileSet::FLOATING_POINT,
57 Level::Level8K,
58 ExtensionSet::NONE
59 .union(ExtensionSet::FP8E4M3)
60 .union(ExtensionSet::FP8E5M2),
61);
62
63pub const VULKAN_TOSA_INTEGER_TARGET: Target = Target::new(
68 Version::TOSA_1_0,
69 ProfileSet::INTEGER,
70 Level::Level8K,
71 ExtensionSet::NONE,
72);
73
74const FLOAT_DTYPES: &[DTypeCapability] = &[
75 DTypeCapability::new(DType::FP32, ValueRoles::ALL),
76 DTypeCapability::new(DType::BOOL, ValueRoles::ALL),
77 DTypeCapability::new(DType::INT32, ValueRoles::ALL),
78 DTypeCapability::constrained(
81 DType::INT8,
82 ValueRoles::CONSTANT,
83 DTypeConstraints::PARAMETER_ONLY,
84 ),
85];
86
87const FLOAT16_DTYPES: &[DTypeCapability] = &[
91 DTypeCapability::new(DType::FP32, ValueRoles::ALL),
92 DTypeCapability::new(DType::FP16, ValueRoles::ALL),
93 DTypeCapability::new(DType::BOOL, ValueRoles::ALL),
94 DTypeCapability::new(DType::INT32, ValueRoles::ALL),
95 DTypeCapability::constrained(
96 DType::INT8,
97 ValueRoles::CONSTANT,
98 DTypeConstraints::PARAMETER_ONLY,
99 ),
100];
101
102const FLOAT_OPERATORS: &[OperatorCapability] = &[
106 OperatorCapability::new(Op::ARGMAX),
107 OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
108 OperatorCapability::new(Op::MAX_POOL2D),
109 OperatorCapability::new(Op::CLAMP),
110 OperatorCapability::new(Op::ERF),
111 OperatorCapability::new(Op::SIGMOID),
112 OperatorCapability::new(Op::TANH),
113 OperatorCapability::new(Op::ADD),
114 OperatorCapability::new(Op::LOGICAL_AND),
115 OperatorCapability::new(Op::LOGICAL_OR),
116 OperatorCapability::new(Op::LOGICAL_XOR),
117 OperatorCapability::new(Op::MAXIMUM),
118 OperatorCapability::new(Op::MINIMUM),
119 OperatorCapability::constrained(Op::MUL, OperatorConstraints::ZERO_SHIFT),
120 OperatorCapability::new(Op::POW),
121 OperatorCapability::new(Op::SUB),
122 OperatorCapability::new(Op::ABS),
123 OperatorCapability::new(Op::CEIL),
124 OperatorCapability::new(Op::COS),
125 OperatorCapability::new(Op::EXP),
126 OperatorCapability::new(Op::FLOOR),
127 OperatorCapability::new(Op::LOG),
128 OperatorCapability::new(Op::LOGICAL_NOT),
129 OperatorCapability::constrained(Op::NEGATE, OperatorConstraints::ZERO_ZERO_POINTS),
130 OperatorCapability::new(Op::RECIPROCAL),
131 OperatorCapability::new(Op::RSQRT),
132 OperatorCapability::new(Op::SIN),
133 OperatorCapability::new(Op::SELECT),
134 OperatorCapability::new(Op::EQUAL),
135 OperatorCapability::new(Op::GREATER),
136 OperatorCapability::new(Op::GREATER_EQUAL),
137 OperatorCapability::new(Op::REDUCE_MAX),
138 OperatorCapability::new(Op::REDUCE_MIN),
139 OperatorCapability::new(Op::REDUCE_PRODUCT),
140 OperatorCapability::new(Op::REDUCE_SUM),
141 OperatorCapability::new(Op::CONCAT),
142 OperatorCapability::constrained(Op::RESHAPE, OperatorConstraints::CONSTANT_PARAMETERS),
143 OperatorCapability::new(Op::REVERSE),
144 OperatorCapability::new(Op::TRANSPOSE),
145 OperatorCapability::new(Op::CONST),
146 OperatorCapability::new(Op::CONST_SHAPE),
147 OperatorCapability::new(Op::IDENTITY),
148];
149
150pub const VULKAN_TOSA_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
152 target: VULKAN_TOSA_TARGET,
153 dtypes: FLOAT_DTYPES,
154 operators: FLOAT_OPERATORS,
155 graph: GraphCapabilities {
156 max_regions: 1,
157 max_blocks: 1,
158 dynamic_shapes: false,
159 runtime_conditions: RuntimeConditionSupport::None,
160 },
161};
162
163pub const VULKAN_TOSA_FP16_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
169 target: VULKAN_TOSA_TARGET,
170 dtypes: FLOAT16_DTYPES,
171 operators: FLOAT_OPERATORS,
172 graph: GraphCapabilities {
173 max_regions: 1,
174 max_blocks: 1,
175 dynamic_shapes: false,
176 runtime_conditions: RuntimeConditionSupport::None,
177 },
178};
179
180const FLOAT8_DTYPES: &[DTypeCapability] = &[
183 DTypeCapability::new(DType::FP8E4M3, ValueRoles::ALL),
184 DTypeCapability::new(DType::FP8E5M2, ValueRoles::ALL),
185 DTypeCapability::new(DType::FP16, ValueRoles::ALL),
186 DTypeCapability::new(DType::INT32, ValueRoles::ALL),
187];
188
189const FLOAT8_OPERATORS: &[OperatorCapability] = &[
199 OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
200 OperatorCapability::new(Op::CONCAT),
201 OperatorCapability::constrained(Op::RESHAPE, OperatorConstraints::CONSTANT_PARAMETERS),
202 OperatorCapability::new(Op::REVERSE),
203 OperatorCapability::new(Op::TRANSPOSE),
204 OperatorCapability::new(Op::CONST),
205 OperatorCapability::new(Op::CONST_SHAPE),
206 OperatorCapability::new(Op::IDENTITY),
207 OperatorCapability::new(Op::CAST),
208 OperatorCapability::new(Op::MAX_POOL2D),
209 OperatorCapability::new(Op::ARGMAX),
210];
211
212pub const VULKAN_TOSA_FP8_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
217 target: VULKAN_TOSA_FP8_TARGET,
218 dtypes: FLOAT8_DTYPES,
219 operators: FLOAT8_OPERATORS,
220 graph: GraphCapabilities {
221 max_regions: 1,
222 max_blocks: 1,
223 dynamic_shapes: false,
224 runtime_conditions: RuntimeConditionSupport::None,
225 },
226};
227
228pub const fn supports_tosa_operator(op: Op) -> bool {
230 VULKAN_TOSA_CAPABILITY.supports_operator(op)
231}
232
233pub const fn supports_tosa_dtype(dtype: DType) -> bool {
235 VULKAN_TOSA_FP16_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
236 || VULKAN_TOSA_FP16_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
237 || VULKAN_TOSA_FP8_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
238 || VULKAN_TOSA_FP8_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
239}
240
241#[derive(Clone, Copy, Debug, PartialEq, Eq)]
243pub enum LoweringError {
244 Parse(ParseError),
245 Analysis(AnalysisError),
246 UnsupportedTarget,
248 UnsupportedGraph,
251 UnsupportedType(DType),
252 UnsupportedOperator(Op),
253 ResourceLimit,
256}
257
258impl fmt::Display for LoweringError {
259 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
260 write!(formatter, "{self:?}")
261 }
262}
263
264impl std::error::Error for LoweringError {}
265
266pub(crate) const ARENA_ALIGNMENT: u64 = 256;
269
270#[derive(Clone, Copy, Debug, PartialEq, Eq)]
273pub(crate) enum KernelSpec {
274 Nvfp4Matmul {
276 cooperative: bool,
279 },
280 Elementwise {
281 op: ElementwiseOp,
282 float: Storage,
283 broadcast: bool,
284 },
285 Reduce {
286 op: ReduceOp,
287 float: Storage,
288 },
289 Matmul {
290 input: Storage,
291 output: Storage,
292 },
293 MatmulStream {
296 rhs: Storage,
297 output: Storage,
298 },
299 Cast {
300 input: Storage,
301 output: Storage,
302 },
303 MaxPool {
304 nan_mode: NanMode,
305 float: Storage,
306 },
307 Move {
308 storage: Storage,
309 contiguous: bool,
310 },
311}
312
313#[derive(Clone, Copy, Debug, PartialEq, Eq)]
315pub(crate) enum Work {
316 Linear(u32),
318 Matmul { m: u32, n: u32, batch: u32 },
320 MatmulStream { n: u32, batch: u32 },
322 Nvfp4Matmul { m: u32, n: u32 },
324}
325
326#[derive(Clone, Debug, PartialEq, Eq)]
328pub(crate) struct DispatchPlan {
329 pub kernel: KernelSpec,
330 pub spec: Vec<u32>,
332 pub work: Work,
333 pub barrier_before: bool,
336}
337
338#[derive(Clone, Copy, Debug, PartialEq, Eq)]
340pub(crate) struct SlotPlan {
341 pub slot: u32,
342 pub role: SlotRole,
343 pub byte_len: u64,
345 pub storage: Storage,
346}
347
348#[derive(Clone, Copy, Debug, PartialEq, Eq)]
349pub(crate) enum SlotRole {
350 Input,
351 Output,
352}
353
354#[derive(Clone, Debug, PartialEq, Eq)]
356pub(crate) struct ConstantPlan {
357 pub offset: u64,
358 pub bytes: Vec<u8>,
359}
360
361#[derive(Clone, Debug, PartialEq, Eq)]
367pub(crate) struct ProgramPlan {
368 pub slots: Vec<SlotPlan>,
369 pub arena_bytes: u64,
371 pub constants: Vec<ConstantPlan>,
372 pub dispatches: Vec<DispatchPlan>,
373}
374
375impl ProgramPlan {
376 pub fn arena_buffer_index(&self) -> u32 {
378 self.slots.len() as u32
379 }
380
381 pub fn buffer_count(&self) -> u32 {
383 self.slots.len() as u32 + u32::from(self.arena_bytes != 0)
384 }
385
386 #[cfg(test)]
387 fn slot(&self, slot: u32) -> Option<&SlotPlan> {
388 self.slots.iter().find(|plan| plan.slot == slot)
389 }
390}
391
392pub(crate) fn lower_tosa(bytes: &[u8], target: Target) -> Result<ProgramPlan, LoweringError> {
394 if target != VULKAN_TOSA_TARGET && target != VULKAN_TOSA_FP8_TARGET {
395 return Err(LoweringError::UnsupportedTarget);
396 }
397 let model = parse(bytes).map_err(LoweringError::Parse)?;
398 let analysis = model.analyze_for(target).map_err(LoweringError::Analysis)?;
399 let capability = if target == VULKAN_TOSA_FP8_TARGET {
400 VULKAN_TOSA_FP8_CAPABILITY
401 } else {
402 VULKAN_TOSA_CAPABILITY
403 };
404 Lowering::new(&analysis, capability)?.run()
405}
406
407#[derive(Clone, Debug, PartialEq, Eq)]
413struct TensorShape {
414 dtype: DType,
415 dims: Vec<u32>,
416 elements: u32,
417}
418
419impl TensorShape {
420 fn storage(&self) -> Storage {
421 storage_of(self.dtype)
422 }
423
424 fn byte_len(&self) -> u64 {
425 u64::from(self.elements) * scalar_bytes(self.dtype)
426 }
427
428 fn rank(&self) -> usize {
429 self.dims.len()
430 }
431
432 fn strides(&self) -> Vec<u32> {
434 let mut strides = vec![1_u32; self.dims.len()];
435 let mut acc = 1_u32;
436 for d in (0..self.dims.len()).rev() {
437 strides[d] = acc;
438 acc = acc.wrapping_mul(self.dims[d]);
439 }
440 strides
441 }
442
443 fn padded_dims(&self) -> [u32; MAX_RANK] {
445 pad_leading(&self.dims, 1)
446 }
447}
448
449fn pad_leading(values: &[u32], fill: u32) -> [u32; MAX_RANK] {
450 let mut padded = [fill; MAX_RANK];
451 let offset = MAX_RANK - values.len();
452 padded[offset..].copy_from_slice(values);
453 padded
454}
455
456fn storage_of(dtype: DType) -> Storage {
457 match dtype {
458 DType::BOOL => Storage::Byte,
459 DType::FP8E4M3 => Storage::Quarter(Fp8Format::E4M3),
460 DType::FP8E5M2 => Storage::Quarter(Fp8Format::E5M2),
461 DType::FP16 => Storage::Half,
462 _ => Storage::Word,
463 }
464}
465
466fn scalar_bytes(dtype: DType) -> u64 {
467 match dtype {
468 DType::BOOL | DType::FP8E4M3 | DType::FP8E5M2 => 1,
469 DType::FP16 => 2,
470 _ => 4,
471 }
472}
473
474#[derive(Clone, Copy, Debug, PartialEq, Eq)]
476enum Location {
477 Slot(u32),
478 Region(usize),
480}
481
482#[derive(Clone, Copy, Debug, PartialEq, Eq)]
487enum MemKey {
488 Slot(u32),
489 Arena { offset: u64, end: u64 },
490}
491
492impl MemKey {
493 fn overlaps(self, other: Self) -> bool {
494 match (self, other) {
495 (Self::Slot(a), Self::Slot(b)) => a == b,
496 (
497 Self::Arena { offset, end },
498 Self::Arena {
499 offset: other_offset,
500 end: other_end,
501 },
502 ) => offset < other_end && other_offset < end,
503 _ => false,
504 }
505 }
506}
507
508#[derive(Clone, Copy, Debug)]
509struct Region {
510 offset: u64,
511 bytes: u64,
512 live_end: u32,
514}
515
516struct Lowering<'a, 'b> {
517 analysis: &'b TosaAnalysis<'a>,
518 capability: CapabilityDescriptor,
521 inputs: &'b [ValueId],
522 outputs: &'b [ValueId],
523 order: &'b [OperatorId],
524 shapes: HashMap<ValueId, TensorShape>,
525 locations: HashMap<ValueId, Location>,
526 last_use: HashMap<ValueId, u32>,
528 regions: Vec<Region>,
529 arena_bytes: u64,
530 constants: Vec<ConstantPlan>,
531 dispatches: Vec<DispatchPlan>,
532 slots: Vec<SlotPlan>,
533 written: Vec<(MemKey, ValueId)>,
535 read: Vec<MemKey>,
536 position: u32,
537}
538
539impl<'a, 'b> Lowering<'a, 'b> {
540 fn new(
541 analysis: &'b TosaAnalysis<'a>,
542 capability: CapabilityDescriptor,
543 ) -> Result<Self, LoweringError> {
544 if analysis.regions().len() != 1 || analysis.blocks().len() != 1 {
545 return Err(LoweringError::UnsupportedGraph);
546 }
547 if analysis
550 .conditions()
551 .iter()
552 .any(|condition| !matches!(condition, RuntimeCondition::PowDomain { .. }))
553 {
554 return Err(LoweringError::UnsupportedGraph);
555 }
556 let block = analysis.blocks()[0].id();
557 let inputs = analysis.block_inputs(block);
558 let outputs = analysis.block_outputs(block);
559 let order = analysis.execution_order(block);
560 if inputs.is_empty()
561 || outputs.is_empty()
562 || inputs.iter().any(|input| outputs.contains(input))
563 || outputs
564 .iter()
565 .any(|output| analysis.serialized_constant(*output).is_some())
566 {
567 return Err(LoweringError::UnsupportedGraph);
568 }
569 let mut duplicates = inputs.iter().chain(outputs).collect::<Vec<_>>();
570 duplicates.sort_unstable_by_key(|value| value.get());
571 if duplicates.windows(2).any(|pair| pair[0] == pair[1]) {
572 return Err(LoweringError::UnsupportedGraph);
573 }
574 if inputs.len() + outputs.len() > u32::MAX as usize {
575 return Err(LoweringError::ResourceLimit);
576 }
577 Ok(Self {
578 analysis,
579 capability,
580 inputs,
581 outputs,
582 order,
583 shapes: HashMap::new(),
584 locations: HashMap::new(),
585 last_use: HashMap::new(),
586 regions: Vec::new(),
587 arena_bytes: 0,
588 constants: Vec::new(),
589 dispatches: Vec::new(),
590 slots: Vec::new(),
591 written: Vec::new(),
592 read: Vec::new(),
593 position: 0,
594 })
595 }
596
597 fn run(mut self) -> Result<ProgramPlan, LoweringError> {
598 for (index, value) in self.inputs.iter().chain(self.outputs).enumerate() {
600 let slot = index as u32;
601 let shape = self.shape(*value)?;
602 self.slots.push(SlotPlan {
603 slot,
604 role: if index < self.inputs.len() {
605 SlotRole::Input
606 } else {
607 SlotRole::Output
608 },
609 byte_len: shape.byte_len(),
610 storage: shape.storage(),
611 });
612 self.locations.insert(*value, Location::Slot(slot));
613 }
614
615 let mut consumers: HashMap<ValueId, u32> = HashMap::new();
617 for (position, operator_id) in self.order.iter().enumerate() {
618 let operator = self.analysis.operator(*operator_id);
619 if self.skipped(operator) {
620 continue;
621 }
622 for input in self.analysis.operator_inputs(*operator_id) {
623 self.last_use.insert(*input, position as u32);
624 *consumers.entry(*input).or_default() += 1;
625 }
626 }
627 self.alias_outputs(&consumers)?;
628
629 for (position, operator_id) in self.order.iter().enumerate() {
630 self.position = position as u32;
631 let operator = self.analysis.operator(*operator_id);
632 if self.skipped(operator) {
633 continue;
634 }
635 let op = operator.op();
636 if !self.capability.supports_operator(op) {
637 return Err(LoweringError::UnsupportedOperator(op));
638 }
639 let operator_inputs = self.analysis.operator_inputs(*operator_id);
640 let operator_outputs = self.analysis.operator_outputs(*operator_id);
641 let [output] = operator_outputs else {
642 return Err(LoweringError::UnsupportedGraph);
643 };
644 let output = *output;
645 if self.analysis.serialized_constant(output).is_some() {
646 return Err(LoweringError::UnsupportedGraph);
647 }
648 match op {
649 Op::IDENTITY => self.lower_copy(operator_inputs, output, 1)?,
650 Op::RESHAPE => self.lower_reshape(operator_inputs, output)?,
651 Op::TRANSPOSE => self.lower_transpose(operator, operator_inputs, output)?,
652 Op::REVERSE => self.lower_reverse(operator, operator_inputs, output)?,
653 Op::CONCAT => self.lower_concat(operator, operator_inputs, output)?,
654 Op::CAST => self.lower_cast(operator_inputs, output)?,
655 Op::MATMUL => self.lower_matmul(operator_inputs, output)?,
656 Op::MAX_POOL2D => self.lower_max_pool(operator, operator_inputs, output)?,
657 Op::ARGMAX
658 | Op::REDUCE_MAX
659 | Op::REDUCE_MIN
660 | Op::REDUCE_PRODUCT
661 | Op::REDUCE_SUM => self.lower_reduce(operator, operator_inputs, output)?,
662 _ => self.lower_elementwise(operator, operator_inputs, output)?,
663 }
664 }
665
666 for output in self.outputs {
668 if self.analysis.value(*output).producer().is_none() {
669 return Err(LoweringError::UnsupportedGraph);
670 }
671 }
672 Ok(ProgramPlan {
673 slots: self.slots,
674 arena_bytes: self.arena_bytes,
675 constants: self.constants,
676 dispatches: self.dispatches,
677 })
678 }
679
680 fn skipped(&self, operator: &virtio_accel_tosa::AnalyzedOperator<'_>) -> bool {
682 matches!(operator.op(), Op::CONST | Op::CONST_SHAPE)
683 || operator.hints().contains(OptimizationHints::DEAD)
684 }
685
686 fn shape(&mut self, value: ValueId) -> Result<TensorShape, LoweringError> {
689 if let Some(shape) = self.shapes.get(&value) {
690 return Ok(shape.clone());
691 }
692 let AnalyzedValueKind::Tensor(tensor) = self.analysis.value(value).kind() else {
693 return Err(LoweringError::UnsupportedGraph);
694 };
695 let dtype = tensor.dtype();
696 if !matches!(
697 dtype,
698 DType::FP32
699 | DType::FP16
700 | DType::BOOL
701 | DType::INT32
702 | DType::FP8E4M3
703 | DType::FP8E5M2
704 ) {
705 return Err(LoweringError::UnsupportedType(dtype));
706 }
707 let rank = tensor.rank().ok_or(LoweringError::UnsupportedGraph)?;
708 if rank > MAX_RANK {
709 return Err(LoweringError::UnsupportedGraph);
710 }
711 let mut dims = Vec::with_capacity(rank);
712 let mut elements = 1_u64;
713 for dimension in tensor.dimensions() {
714 let dimension = u32::try_from(dimension)
715 .ok()
716 .filter(|dimension| *dimension > 0)
717 .ok_or(LoweringError::UnsupportedGraph)?;
718 dims.push(dimension);
719 elements = elements
720 .checked_mul(u64::from(dimension))
721 .ok_or(LoweringError::ResourceLimit)?;
722 }
723 let elements = u32::try_from(elements).map_err(|_| LoweringError::ResourceLimit)?;
724 let shape = TensorShape {
725 dtype,
726 dims,
727 elements,
728 };
729 self.shapes.insert(value, shape.clone());
730 Ok(shape)
731 }
732
733 fn float_shape(&mut self, value: ValueId) -> Result<TensorShape, LoweringError> {
735 let shape = self.shape(value)?;
736 if !matches!(shape.dtype, DType::FP32 | DType::FP16) {
737 return Err(LoweringError::UnsupportedType(shape.dtype));
738 }
739 Ok(shape)
740 }
741
742 fn convertible_float_shape(&mut self, value: ValueId) -> Result<TensorShape, LoweringError> {
746 let shape = self.shape(value)?;
747 if !matches!(
748 shape.dtype,
749 DType::FP32 | DType::FP16 | DType::FP8E4M3 | DType::FP8E5M2
750 ) {
751 return Err(LoweringError::UnsupportedType(shape.dtype));
752 }
753 Ok(shape)
754 }
755
756 fn typed_shape(&mut self, value: ValueId, dtype: DType) -> Result<TensorShape, LoweringError> {
757 let shape = self.shape(value)?;
758 if shape.dtype != dtype {
759 return Err(LoweringError::UnsupportedType(shape.dtype));
760 }
761 Ok(shape)
762 }
763
764 fn allocate_region(&mut self, bytes: u64, live_end: u32) -> Result<usize, LoweringError> {
770 self.allocate_region_written_at(bytes, live_end, self.position)
771 }
772
773 fn allocate_region_written_at(
778 &mut self,
779 bytes: u64,
780 live_end: u32,
781 written_at: u32,
782 ) -> Result<usize, LoweringError> {
783 let bytes = bytes.max(1).div_ceil(ARENA_ALIGNMENT) * ARENA_ALIGNMENT;
784 let mut candidates = vec![0_u64];
785 let overlapping: Vec<Region> = self
786 .regions
787 .iter()
788 .filter(|region| region.live_end >= written_at)
789 .copied()
790 .collect();
791 for region in &overlapping {
792 candidates.push(region.offset + region.bytes);
793 }
794 candidates.sort_unstable();
795 let offset = candidates
796 .into_iter()
797 .find(|candidate| {
798 let end = candidate + bytes;
799 overlapping.iter().all(|region| {
800 end <= region.offset || *candidate >= region.offset + region.bytes
801 })
802 })
803 .ok_or(LoweringError::ResourceLimit)?;
804 let end = offset
805 .checked_add(bytes)
806 .ok_or(LoweringError::ResourceLimit)?;
807 if end / 4 > u64::from(u32::MAX) {
809 return Err(LoweringError::ResourceLimit);
810 }
811 self.arena_bytes = self.arena_bytes.max(end);
812 self.regions.push(Region {
813 offset,
814 bytes,
815 live_end,
816 });
817 Ok(self.regions.len() - 1)
818 }
819
820 fn alias_outputs(&mut self, consumers: &HashMap<ValueId, u32>) -> Result<(), LoweringError> {
826 for output in self.outputs {
827 let Some(producer) = self.analysis.value(*output).producer() else {
828 continue;
829 };
830 let operator = self.analysis.operator(producer);
831 if self.skipped(operator) || !matches!(operator.op(), Op::IDENTITY | Op::RESHAPE) {
832 continue;
833 }
834 let Some(source) = self.analysis.operator_inputs(producer).first().copied() else {
835 continue;
836 };
837 if self.locations.contains_key(&source)
838 || self.analysis.serialized_constant(source).is_some()
839 || consumers.get(&source) != Some(&1)
840 {
841 continue;
842 }
843 let (from, to) = (self.shape(source)?, self.shape(*output)?);
844 if from.dtype != to.dtype || from.elements != to.elements {
845 continue;
846 }
847 let slot = self.locations[output];
848 self.locations.insert(source, slot);
849 }
850 Ok(())
851 }
852
853 fn value_last_use(&self, value: ValueId) -> u32 {
854 self.last_use.get(&value).copied().unwrap_or(self.position)
855 }
856
857 fn input_location(&mut self, value: ValueId) -> Result<Location, LoweringError> {
860 if let Some(location) = self.locations.get(&value) {
861 return Ok(*location);
862 }
863 let shape = self.shape(value)?;
864 let bytes = self
865 .analysis
866 .serialized_constant(value)
867 .ok_or(LoweringError::UnsupportedGraph)?;
868 if bytes.len() as u64 != shape.byte_len() {
869 return Err(LoweringError::UnsupportedGraph);
870 }
871 let region = self.allocate_region_written_at(shape.byte_len(), u32::MAX, 0)?;
876 self.constants.push(ConstantPlan {
877 offset: self.regions[region].offset,
878 bytes: bytes.to_vec(),
879 });
880 let location = Location::Region(region);
881 self.locations.insert(value, location);
882 Ok(location)
883 }
884
885 fn output_location(&mut self, value: ValueId) -> Result<Location, LoweringError> {
887 if let Some(location) = self.locations.get(&value) {
888 return Ok(*location);
889 }
890 let shape = self.shape(value)?;
891 let live_end = self.value_last_use(value);
892 let region = self.allocate_region(shape.byte_len(), live_end)?;
893 let location = Location::Region(region);
894 self.locations.insert(value, location);
895 Ok(location)
896 }
897
898 fn operand(&self, location: Location) -> Operand {
899 match location {
900 Location::Slot(slot) => Operand {
901 buffer: slot,
902 base: 0,
903 },
904 Location::Region(index) => Operand {
905 buffer: self.slots.len() as u32,
906 base: (self.regions[index].offset / 4) as u32,
907 },
908 }
909 }
910
911 fn mem_key(&self, location: Location) -> MemKey {
914 match location {
915 Location::Slot(slot) => MemKey::Slot(slot),
916 Location::Region(index) => {
917 let region = &self.regions[index];
918 MemKey::Arena {
919 offset: region.offset,
920 end: region.offset + region.bytes,
921 }
922 }
923 }
924 }
925
926 fn dispatch(
927 &mut self,
928 kernel: KernelSpec,
929 spec: Vec<u32>,
930 work: Work,
931 reads: &[Location],
932 writes: Location,
933 written_value: ValueId,
934 ) {
935 let reads: Vec<MemKey> = reads.iter().map(|read| self.mem_key(*read)).collect();
936 let writes = self.mem_key(writes);
937 let raw = reads.iter().any(|read| {
938 self.written
939 .iter()
940 .any(|(written, _)| written.overlaps(*read))
941 });
942 let war = self.read.iter().any(|read| read.overlaps(writes));
943 let waw = self
946 .written
947 .iter()
948 .any(|(written, value)| written.overlaps(writes) && *value != written_value);
949 let barrier_before = raw || war || waw;
950 if barrier_before {
951 self.written.clear();
952 self.read.clear();
953 }
954 self.read.extend_from_slice(&reads);
955 self.written.push((writes, written_value));
956 self.dispatches.push(DispatchPlan {
957 kernel,
958 spec,
959 work,
960 barrier_before,
961 });
962 }
963
964 fn lower_copy(
971 &mut self,
972 inputs: &[ValueId],
973 output: ValueId,
974 expected_inputs: usize,
975 ) -> Result<(), LoweringError> {
976 if inputs.len() != expected_inputs {
977 return Err(LoweringError::UnsupportedGraph);
978 }
979 let source = self.shape(inputs[0])?;
980 let target = self.shape(output)?;
981 if source.dtype != target.dtype || source.elements != target.elements {
982 return Err(LoweringError::UnsupportedGraph);
983 }
984 let from = self.input_location(inputs[0])?;
985 if !self.locations.contains_key(&output) {
986 if let Location::Region(region) = from {
989 let live_end = self.value_last_use(output);
990 self.regions[region].live_end = self.regions[region].live_end.max(live_end);
991 }
992 self.locations.insert(output, from);
993 return Ok(());
994 }
995 let to = self.output_location(output)?;
996 if to == from {
997 return Ok(());
1000 }
1001 let storage = source.storage();
1002 let geometry = MoveGeometry {
1003 count: source.elements,
1004 dims: source.padded_dims(),
1005 in_strides: pad_leading(&source.strides(), 0),
1006 in_offset: 0,
1007 out_strides: pad_leading(&source.strides(), 0),
1008 out_offset: 0,
1009 };
1010 let spec = move_spec(self.operand(from), self.operand(to), geometry, true);
1011 self.dispatch(
1014 KernelSpec::Move {
1015 storage,
1016 contiguous: true,
1017 },
1018 spec,
1019 Work::Linear(source.elements.div_ceil(storage.lanes())),
1020 &[from],
1021 to,
1022 output,
1023 );
1024 Ok(())
1025 }
1026
1027 fn lower_reshape(&mut self, inputs: &[ValueId], output: ValueId) -> Result<(), LoweringError> {
1028 let [source, shape] = inputs else {
1029 return Err(LoweringError::UnsupportedGraph);
1030 };
1031 let AnalyzedValueKind::Shape(shape) = self.analysis.value(*shape).kind() else {
1032 return Err(LoweringError::UnsupportedGraph);
1033 };
1034 let values = shape.values().ok_or(LoweringError::UnsupportedGraph)?;
1035 let target = self.shape(output)?;
1036 let declared: Vec<i64> = values.collect();
1037 if declared.len() != target.rank()
1038 || declared
1039 .iter()
1040 .zip(&target.dims)
1041 .any(|(declared, dim)| *declared != i64::from(*dim))
1042 {
1043 return Err(LoweringError::UnsupportedGraph);
1044 }
1045 self.lower_copy(&[*source], output, 1)
1046 }
1047
1048 fn lower_transpose(
1049 &mut self,
1050 operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1051 inputs: &[ValueId],
1052 output: ValueId,
1053 ) -> Result<(), LoweringError> {
1054 let [input] = inputs else {
1055 return Err(LoweringError::UnsupportedGraph);
1056 };
1057 let OpAttributes::Transpose { perms } = operator.source().attributes() else {
1058 return Err(LoweringError::UnsupportedGraph);
1059 };
1060 let source = self.shape(*input)?;
1061 let target = self.shape(output)?;
1062 let perms: Vec<usize> = perms
1063 .iter()
1064 .map(|perm| usize::try_from(perm).map_err(|_| LoweringError::UnsupportedGraph))
1065 .collect::<Result<_, _>>()?;
1066 if perms.len() != source.rank()
1067 || target.rank() != source.rank()
1068 || source.dtype != target.dtype
1069 {
1070 return Err(LoweringError::UnsupportedGraph);
1071 }
1072 let mut seen = vec![false; perms.len()];
1073 for (d, perm) in perms.iter().enumerate() {
1074 if *perm >= perms.len() || seen[*perm] || target.dims[d] != source.dims[*perm] {
1075 return Err(LoweringError::UnsupportedGraph);
1076 }
1077 seen[*perm] = true;
1078 }
1079 let source_strides = source.strides();
1080 let in_strides: Vec<u32> = perms.iter().map(|perm| source_strides[*perm]).collect();
1081 let geometry = MoveGeometry {
1082 count: target.elements,
1083 dims: target.padded_dims(),
1084 in_strides: pad_leading(&in_strides, 0),
1085 in_offset: 0,
1086 out_strides: pad_leading(&target.strides(), 0),
1087 out_offset: 0,
1088 };
1089 self.lower_move(*input, output, source.storage(), geometry)
1090 }
1091
1092 fn lower_reverse(
1093 &mut self,
1094 operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1095 inputs: &[ValueId],
1096 output: ValueId,
1097 ) -> Result<(), LoweringError> {
1098 let [input] = inputs else {
1099 return Err(LoweringError::UnsupportedGraph);
1100 };
1101 let OpAttributes::Reverse { axis } = operator.source().attributes() else {
1102 return Err(LoweringError::UnsupportedGraph);
1103 };
1104 let source = self.shape(*input)?;
1105 let target = self.shape(output)?;
1106 if source != target {
1107 return Err(LoweringError::UnsupportedGraph);
1108 }
1109 let axis = usize::try_from(axis)
1110 .ok()
1111 .filter(|axis| *axis < source.rank())
1112 .ok_or(LoweringError::UnsupportedGraph)?;
1113 let strides = source.strides();
1114 let mut in_strides = strides.clone();
1115 in_strides[axis] = strides[axis].wrapping_neg();
1116 let in_offset = (source.dims[axis] - 1).wrapping_mul(strides[axis]);
1117 let geometry = MoveGeometry {
1118 count: source.elements,
1119 dims: source.padded_dims(),
1120 in_strides: pad_leading(&in_strides, 0),
1121 in_offset,
1122 out_strides: pad_leading(&strides, 0),
1123 out_offset: 0,
1124 };
1125 self.lower_move(*input, output, source.storage(), geometry)
1126 }
1127
1128 fn lower_concat(
1129 &mut self,
1130 operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1131 inputs: &[ValueId],
1132 output: ValueId,
1133 ) -> Result<(), LoweringError> {
1134 if inputs.is_empty() {
1135 return Err(LoweringError::UnsupportedGraph);
1136 }
1137 let OpAttributes::Concat { axis } = operator.source().attributes() else {
1138 return Err(LoweringError::UnsupportedGraph);
1139 };
1140 let target = self.shape(output)?;
1141 let axis = usize::try_from(axis)
1142 .ok()
1143 .filter(|axis| *axis < target.rank())
1144 .ok_or(LoweringError::UnsupportedGraph)?;
1145 let mut total = 0_u64;
1146 let mut sources = Vec::with_capacity(inputs.len());
1147 for input in inputs {
1148 let source = self.shape(*input)?;
1149 if source.dtype != target.dtype
1150 || source.rank() != target.rank()
1151 || source
1152 .dims
1153 .iter()
1154 .zip(&target.dims)
1155 .enumerate()
1156 .any(|(d, (source, target))| d != axis && source != target)
1157 {
1158 return Err(LoweringError::UnsupportedGraph);
1159 }
1160 total += u64::from(source.dims[axis]);
1161 sources.push(source);
1162 }
1163 if total != u64::from(target.dims[axis]) {
1164 return Err(LoweringError::UnsupportedGraph);
1165 }
1166 let out_strides = target.strides();
1167 let to = self.output_location(output)?;
1168 let mut offset = 0_u32;
1169 for (input, source) in inputs.iter().zip(&sources) {
1170 let geometry = MoveGeometry {
1171 count: source.elements,
1172 dims: source.padded_dims(),
1173 in_strides: pad_leading(&source.strides(), 0),
1174 in_offset: 0,
1175 out_strides: pad_leading(&out_strides, 0),
1176 out_offset: offset.wrapping_mul(out_strides[axis]),
1177 };
1178 offset += source.dims[axis];
1179 let from = self.input_location(*input)?;
1180 let spec = move_spec(self.operand(from), self.operand(to), geometry, false);
1181 self.dispatch(
1182 KernelSpec::Move {
1183 storage: source.storage(),
1184 contiguous: false,
1185 },
1186 spec,
1187 Work::Linear(source.elements),
1188 &[from],
1189 to,
1190 output,
1191 );
1192 }
1193 Ok(())
1194 }
1195
1196 fn lower_move(
1197 &mut self,
1198 input: ValueId,
1199 output: ValueId,
1200 storage: Storage,
1201 geometry: MoveGeometry,
1202 ) -> Result<(), LoweringError> {
1203 let from = self.input_location(input)?;
1204 let to = self.output_location(output)?;
1205 let spec = move_spec(self.operand(from), self.operand(to), geometry, false);
1206 self.dispatch(
1207 KernelSpec::Move {
1208 storage,
1209 contiguous: false,
1210 },
1211 spec,
1212 Work::Linear(geometry.count),
1213 &[from],
1214 to,
1215 output,
1216 );
1217 Ok(())
1218 }
1219
1220 fn lower_cast(&mut self, inputs: &[ValueId], output: ValueId) -> Result<(), LoweringError> {
1224 let [input] = inputs else {
1225 return Err(LoweringError::UnsupportedGraph);
1226 };
1227 let source = self.shape(*input)?;
1228 let destination = self.shape(output)?;
1229 if source.dims != destination.dims {
1230 return Err(LoweringError::UnsupportedGraph);
1231 }
1232 for dtype in [source.dtype, destination.dtype] {
1233 if !matches!(
1234 dtype,
1235 DType::FP32 | DType::FP16 | DType::FP8E4M3 | DType::FP8E5M2
1236 ) {
1237 return Err(LoweringError::UnsupportedType(dtype));
1238 }
1239 }
1240 let from = self.input_location(*input)?;
1241 let to = self.output_location(output)?;
1242 let strides = pad_leading(&source.strides(), 0);
1243 let geometry = MoveGeometry {
1244 count: source.elements,
1245 dims: source.padded_dims(),
1246 in_strides: strides,
1247 in_offset: 0,
1248 out_strides: strides,
1249 out_offset: 0,
1250 };
1251 let spec = move_spec(self.operand(from), self.operand(to), geometry, true);
1252 self.dispatch(
1254 KernelSpec::Cast {
1255 input: source.storage(),
1256 output: destination.storage(),
1257 },
1258 spec,
1259 Work::Linear(source.elements.div_ceil(destination.storage().lanes())),
1260 &[from],
1261 to,
1262 output,
1263 );
1264 Ok(())
1265 }
1266
1267 fn lower_matmul(&mut self, inputs: &[ValueId], output: ValueId) -> Result<(), LoweringError> {
1268 let [lhs, rhs, lhs_zp, rhs_zp] = inputs else {
1269 return Err(LoweringError::UnsupportedGraph);
1270 };
1271 let lhs_shape = self.convertible_float_shape(*lhs)?;
1272 let rhs_shape = self.typed_shape(*rhs, lhs_shape.dtype)?;
1273 let result_dtype = match lhs_shape.dtype {
1276 DType::FP8E4M3 | DType::FP8E5M2 => DType::FP16,
1277 dtype => dtype,
1278 };
1279 let out_shape = self.typed_shape(output, result_dtype)?;
1280 for zero_point in [lhs_zp, rhs_zp] {
1283 self.require_zero_constant(*zero_point, lhs_shape.dtype)?;
1284 }
1285 let ([batch_l, m, k], [batch_r, k_r, n], [batch_o, m_o, n_o]) = (
1286 lhs_shape.dims.as_slice(),
1287 rhs_shape.dims.as_slice(),
1288 out_shape.dims.as_slice(),
1289 ) else {
1290 return Err(LoweringError::UnsupportedGraph);
1291 };
1292 if batch_l != batch_r || batch_l != batch_o || k != k_r || m != m_o || n != n_o {
1293 return Err(LoweringError::UnsupportedGraph);
1294 }
1295 let (batch, m, n, k) = (*batch_l, *m, *n, *k);
1296 let from_lhs = self.input_location(*lhs)?;
1297 let from_rhs = self.input_location(*rhs)?;
1298 let to = self.output_location(output)?;
1299 let (input, output_storage) = (lhs_shape.storage(), out_shape.storage());
1300 if m <= crate::shader::STREAM_ROWS {
1301 let lhs_words = if input == Storage::Word {
1306 from_lhs
1307 } else {
1308 let count = batch
1309 .checked_mul(m)
1310 .and_then(|rows| rows.checked_mul(k))
1311 .ok_or(LoweringError::ResourceLimit)?;
1312 let region = self.allocate_region(u64::from(count) * 4, self.position)?;
1313 let widened = Location::Region(region);
1314 let geometry = MoveGeometry {
1315 count,
1316 dims: [1; MAX_RANK],
1317 in_strides: [0; MAX_RANK],
1318 in_offset: 0,
1319 out_strides: [0; MAX_RANK],
1320 out_offset: 0,
1321 };
1322 let spec = move_spec(
1323 self.operand(from_lhs),
1324 self.operand(widened),
1325 geometry,
1326 true,
1327 );
1328 self.dispatch(
1329 KernelSpec::Cast {
1330 input,
1331 output: Storage::Word,
1332 },
1333 spec,
1334 Work::Linear(count),
1335 &[from_lhs],
1336 widened,
1337 output,
1338 );
1339 widened
1340 };
1341 let spec = matmul_spec(
1342 self.operand(lhs_words),
1343 self.operand(from_rhs),
1344 self.operand(to),
1345 m,
1346 n,
1347 k,
1348 batch,
1349 );
1350 self.dispatch(
1351 KernelSpec::MatmulStream {
1352 rhs: input,
1353 output: output_storage,
1354 },
1355 spec,
1356 Work::MatmulStream { n, batch },
1357 &[lhs_words, from_rhs],
1358 to,
1359 output,
1360 );
1361 } else {
1362 let spec = matmul_spec(
1363 self.operand(from_lhs),
1364 self.operand(from_rhs),
1365 self.operand(to),
1366 m,
1367 n,
1368 k,
1369 batch,
1370 );
1371 self.dispatch(
1372 KernelSpec::Matmul {
1373 input,
1374 output: output_storage,
1375 },
1376 spec,
1377 Work::Matmul { m, n, batch },
1378 &[from_lhs, from_rhs],
1379 to,
1380 output,
1381 );
1382 }
1383 Ok(())
1384 }
1385
1386 fn lower_max_pool(
1387 &mut self,
1388 operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1389 inputs: &[ValueId],
1390 output: ValueId,
1391 ) -> Result<(), LoweringError> {
1392 let [input] = inputs else {
1393 return Err(LoweringError::UnsupportedGraph);
1394 };
1395 let OpAttributes::MaxPool2d {
1396 kernel,
1397 stride,
1398 pad,
1399 nan_mode,
1400 } = operator.source().attributes()
1401 else {
1402 return Err(LoweringError::UnsupportedGraph);
1403 };
1404 let nan_mode = nan_mode_of(nan_mode)?;
1405 let source = self.convertible_float_shape(*input)?;
1408 let target = self.typed_shape(output, source.dtype)?;
1409 let ([batch, height, width, channels], [batch_o, out_height, out_width, channels_o]) =
1410 (source.dims.as_slice(), target.dims.as_slice())
1411 else {
1412 return Err(LoweringError::UnsupportedGraph);
1413 };
1414 let attribute = |list: virtio_accel_tosa::I32List<'_>, len: usize| {
1415 let values: Vec<u32> = list
1416 .iter()
1417 .map(|value| u32::try_from(value).map_err(|_| LoweringError::UnsupportedGraph))
1418 .collect::<Result<_, _>>()?;
1419 if values.len() != len {
1420 return Err(LoweringError::UnsupportedGraph);
1421 }
1422 Ok(values)
1423 };
1424 let kernel = attribute(kernel, 2)?;
1425 let stride = attribute(stride, 2)?;
1426 let pad = attribute(pad, 4)?;
1427 if kernel.contains(&0) || stride.contains(&0) {
1428 return Err(LoweringError::UnsupportedGraph);
1429 }
1430 let [pad_top, pad_bottom, pad_left, pad_right] = pad[..] else {
1433 return Err(LoweringError::UnsupportedGraph);
1434 };
1435 if pad_top >= kernel[0]
1436 || pad_bottom >= kernel[0]
1437 || pad_left >= kernel[1]
1438 || pad_right >= kernel[1]
1439 {
1440 return Err(LoweringError::UnsupportedGraph);
1441 }
1442 let output_extent = |input: u32, pad_a: u32, pad_b: u32, kernel: u32, stride: u32| {
1443 let padded = u64::from(input) + u64::from(pad_a) + u64::from(pad_b);
1444 let span = padded.checked_sub(u64::from(kernel))?;
1445 if span % u64::from(stride) != 0 {
1446 return None;
1447 }
1448 u32::try_from(span / u64::from(stride) + 1).ok()
1449 };
1450 let expected_height = output_extent(*height, pad_top, pad_bottom, kernel[0], stride[0])
1451 .ok_or(LoweringError::UnsupportedGraph)?;
1452 let expected_width = output_extent(*width, pad_left, pad_right, kernel[1], stride[1])
1453 .ok_or(LoweringError::UnsupportedGraph)?;
1454 if batch != batch_o
1455 || channels != channels_o
1456 || expected_height != *out_height
1457 || expected_width != *out_width
1458 {
1459 return Err(LoweringError::UnsupportedGraph);
1460 }
1461 let last_row = u64::from(*out_height - 1) * u64::from(stride[0]) + u64::from(kernel[0]);
1463 let last_col = u64::from(*out_width - 1) * u64::from(stride[1]) + u64::from(kernel[1]);
1464 if last_row > u64::from(u32::MAX) || last_col > u64::from(u32::MAX) {
1465 return Err(LoweringError::ResourceLimit);
1466 }
1467 let geometry = PoolGeometry {
1468 batch: *batch,
1469 height: *height,
1470 width: *width,
1471 channels: *channels,
1472 out_height: *out_height,
1473 out_width: *out_width,
1474 kernel: [kernel[0], kernel[1]],
1475 stride: [stride[0], stride[1]],
1476 pad_top,
1477 pad_left,
1478 };
1479 let from = self.input_location(*input)?;
1480 let to = self.output_location(output)?;
1481 let spec = max_pool_spec(self.operand(from), self.operand(to), geometry);
1482 self.dispatch(
1483 KernelSpec::MaxPool {
1484 nan_mode,
1485 float: source.storage(),
1486 },
1487 spec,
1488 Work::Linear(target.elements),
1489 &[from],
1490 to,
1491 output,
1492 );
1493 Ok(())
1494 }
1495
1496 fn lower_reduce(
1497 &mut self,
1498 operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1499 inputs: &[ValueId],
1500 output: ValueId,
1501 ) -> Result<(), LoweringError> {
1502 let [input] = inputs else {
1503 return Err(LoweringError::UnsupportedGraph);
1504 };
1505 let (axis, op, argmax) = match operator.source().attributes() {
1506 OpAttributes::ArgMax { axis, nan_mode } => {
1507 (axis, ReduceOp::ArgMax(nan_mode_of(nan_mode)?), true)
1508 }
1509 OpAttributes::ReduceMax { axis, nan_mode } => {
1510 (axis, ReduceOp::Max(nan_mode_of(nan_mode)?), false)
1511 }
1512 OpAttributes::ReduceMin { axis, nan_mode } => {
1513 (axis, ReduceOp::Min(nan_mode_of(nan_mode)?), false)
1514 }
1515 OpAttributes::ReduceProduct { axis } => (axis, ReduceOp::Product, false),
1516 OpAttributes::ReduceSum { axis } => (axis, ReduceOp::Sum, false),
1517 _ => return Err(LoweringError::UnsupportedGraph),
1518 };
1519 let source = if argmax {
1522 self.convertible_float_shape(*input)?
1523 } else {
1524 self.float_shape(*input)?
1525 };
1526 let target = self.typed_shape(output, if argmax { DType::INT32 } else { source.dtype })?;
1527 let axis = usize::try_from(axis)
1528 .ok()
1529 .filter(|axis| *axis < source.rank())
1530 .ok_or(LoweringError::UnsupportedGraph)?;
1531 let mut expected = source.dims.clone();
1532 if argmax {
1533 expected.remove(axis);
1534 } else {
1535 expected[axis] = 1;
1536 }
1537 if target.dims != expected {
1538 return Err(LoweringError::UnsupportedGraph);
1539 }
1540 let outer = source.dims[..axis].iter().product::<u32>();
1541 let inner = source.dims[axis + 1..].iter().product::<u32>();
1542 let from = self.input_location(*input)?;
1543 let to = self.output_location(output)?;
1544 let spec = reduce_spec(
1545 self.operand(from),
1546 self.operand(to),
1547 outer,
1548 source.dims[axis],
1549 inner,
1550 );
1551 self.dispatch(
1552 KernelSpec::Reduce {
1553 op,
1554 float: source.storage(),
1555 },
1556 spec,
1557 Work::Linear(target.elements),
1558 &[from],
1559 to,
1560 output,
1561 );
1562 Ok(())
1563 }
1564
1565 fn lower_elementwise(
1566 &mut self,
1567 operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1568 inputs: &[ValueId],
1569 output: ValueId,
1570 ) -> Result<(), LoweringError> {
1571 let op = operator.op();
1572 let attributes = operator.source().attributes();
1573 let mut clamp = None;
1574 let (lane, tensor_inputs): (ElementwiseOp, Vec<ValueId>) = match op {
1575 Op::ABS => (ElementwiseOp::Abs, inputs.to_vec()),
1576 Op::CEIL => (ElementwiseOp::Ceil, inputs.to_vec()),
1577 Op::COS => (ElementwiseOp::Cos, inputs.to_vec()),
1578 Op::ERF => (ElementwiseOp::Erf, inputs.to_vec()),
1579 Op::EXP => (ElementwiseOp::Exp, inputs.to_vec()),
1580 Op::FLOOR => (ElementwiseOp::Floor, inputs.to_vec()),
1581 Op::LOG => (ElementwiseOp::Log, inputs.to_vec()),
1582 Op::RECIPROCAL => (ElementwiseOp::Reciprocal, inputs.to_vec()),
1583 Op::RSQRT => (ElementwiseOp::Rsqrt, inputs.to_vec()),
1584 Op::SIN => (ElementwiseOp::Sin, inputs.to_vec()),
1585 Op::SIGMOID => (ElementwiseOp::Sigmoid, inputs.to_vec()),
1586 Op::TANH => (ElementwiseOp::Tanh, inputs.to_vec()),
1587 Op::NEGATE => {
1588 let [value, input_zp, output_zp] = inputs else {
1589 return Err(LoweringError::UnsupportedGraph);
1590 };
1591 let dtype = self.shape(*value)?.dtype;
1592 self.require_zero_constant(*input_zp, dtype)?;
1593 self.require_zero_constant(*output_zp, dtype)?;
1594 (ElementwiseOp::Negate, vec![*value])
1595 }
1596 Op::CLAMP => {
1597 let OpAttributes::Clamp {
1598 min_val,
1599 max_val,
1600 nan_mode,
1601 } = attributes
1602 else {
1603 return Err(LoweringError::UnsupportedGraph);
1604 };
1605 let bound = |bytes: &[u8]| -> Result<u32, LoweringError> {
1610 let value = match bytes.len() {
1611 4 => f32::from_le_bytes(
1612 bytes
1613 .try_into()
1614 .map_err(|_| LoweringError::UnsupportedGraph)?,
1615 ),
1616 2 => crate::shader::f16_to_f32(u16::from_le_bytes(
1617 bytes
1618 .try_into()
1619 .map_err(|_| LoweringError::UnsupportedGraph)?,
1620 )),
1621 _ => return Err(LoweringError::UnsupportedGraph),
1622 };
1623 if value.is_nan() {
1624 return Err(LoweringError::UnsupportedGraph);
1625 }
1626 Ok(value.to_bits())
1627 };
1628 let lo = bound(min_val)?;
1629 let hi = bound(max_val)?;
1630 if f32::from_bits(hi) < f32::from_bits(lo) {
1631 return Err(LoweringError::UnsupportedGraph);
1632 }
1633 clamp = Some([lo, hi]);
1634 (
1635 ElementwiseOp::Clamp(nan_mode_of(nan_mode)?),
1636 inputs.to_vec(),
1637 )
1638 }
1639 Op::ADD => (ElementwiseOp::Add, inputs.to_vec()),
1640 Op::SUB => (ElementwiseOp::Sub, inputs.to_vec()),
1641 Op::POW => (ElementwiseOp::Pow, inputs.to_vec()),
1642 Op::MUL => {
1643 let [lhs, rhs, shift] = inputs else {
1644 return Err(LoweringError::UnsupportedGraph);
1645 };
1646 self.require_zero_constant(*shift, DType::INT8)?;
1647 (ElementwiseOp::Mul, vec![*lhs, *rhs])
1648 }
1649 Op::MAXIMUM => {
1650 let OpAttributes::Maximum { nan_mode } = attributes else {
1651 return Err(LoweringError::UnsupportedGraph);
1652 };
1653 (
1654 ElementwiseOp::Maximum(nan_mode_of(nan_mode)?),
1655 inputs.to_vec(),
1656 )
1657 }
1658 Op::MINIMUM => {
1659 let OpAttributes::Minimum { nan_mode } = attributes else {
1660 return Err(LoweringError::UnsupportedGraph);
1661 };
1662 (
1663 ElementwiseOp::Minimum(nan_mode_of(nan_mode)?),
1664 inputs.to_vec(),
1665 )
1666 }
1667 Op::EQUAL => (ElementwiseOp::Equal, inputs.to_vec()),
1668 Op::GREATER => (ElementwiseOp::Greater, inputs.to_vec()),
1669 Op::GREATER_EQUAL => (ElementwiseOp::GreaterEqual, inputs.to_vec()),
1670 Op::LOGICAL_AND => (ElementwiseOp::LogicalAnd, inputs.to_vec()),
1671 Op::LOGICAL_OR => (ElementwiseOp::LogicalOr, inputs.to_vec()),
1672 Op::LOGICAL_XOR => (ElementwiseOp::LogicalXor, inputs.to_vec()),
1673 Op::LOGICAL_NOT => (ElementwiseOp::LogicalNot, inputs.to_vec()),
1674 Op::SELECT => (ElementwiseOp::Select, inputs.to_vec()),
1675 other => return Err(LoweringError::UnsupportedOperator(other)),
1676 };
1677 self.emit_elementwise(lane, &tensor_inputs, output, clamp)
1678 }
1679
1680 fn emit_elementwise(
1681 &mut self,
1682 lane: ElementwiseOp,
1683 inputs: &[ValueId],
1684 output: ValueId,
1685 clamp: Option<[u32; 2]>,
1686 ) -> Result<(), LoweringError> {
1687 let lanes = lane.inputs();
1688 if inputs.len() != lanes.len() {
1689 return Err(LoweringError::UnsupportedGraph);
1690 }
1691 let target = self.shape(output)?;
1692 let mut float_dtype: Option<DType> = None;
1695 let mut unify = |dtype: DType| -> Result<(), LoweringError> {
1696 if !matches!(dtype, DType::FP32 | DType::FP16) {
1697 return Err(LoweringError::UnsupportedType(dtype));
1698 }
1699 match float_dtype {
1700 Some(existing) if existing != dtype => Err(LoweringError::UnsupportedGraph),
1701 _ => {
1702 float_dtype = Some(dtype);
1703 Ok(())
1704 }
1705 }
1706 };
1707 let mut shapes = Vec::with_capacity(inputs.len());
1708 for (input, storage) in inputs.iter().zip(lanes) {
1709 let shape = self.shape(*input)?;
1710 match storage {
1711 Storage::Word => unify(shape.dtype)?,
1712 Storage::Byte if shape.dtype != DType::BOOL => {
1713 return Err(LoweringError::UnsupportedType(shape.dtype));
1714 }
1715 _ => {}
1716 }
1717 if shape.rank() != target.rank()
1718 || shape
1719 .dims
1720 .iter()
1721 .zip(&target.dims)
1722 .any(|(dim, out)| *dim != *out && *dim != 1)
1723 {
1724 return Err(LoweringError::UnsupportedGraph);
1725 }
1726 shapes.push(shape);
1727 }
1728 match lane.output() {
1729 Storage::Word => unify(target.dtype)?,
1730 Storage::Byte if target.dtype != DType::BOOL => {
1731 return Err(LoweringError::UnsupportedType(target.dtype));
1732 }
1733 _ => {}
1734 }
1735 let float = storage_of(float_dtype.unwrap_or(DType::FP32));
1736 let broadcast = shapes.iter().any(|shape| shape.dims != target.dims);
1737 let mut strides = Vec::with_capacity(inputs.len());
1738 for shape in &shapes {
1739 let own = shape.strides();
1740 let mut broadcast_strides = vec![0_u32; shape.rank()];
1741 for d in 0..shape.rank() {
1742 broadcast_strides[d] = if shape.dims[d] == 1 && target.dims[d] != 1 {
1743 0
1744 } else {
1745 own[d]
1746 };
1747 }
1748 strides.push(pad_leading(&broadcast_strides, 0));
1749 }
1750 let mut reads = Vec::with_capacity(inputs.len());
1751 let mut operands = Vec::with_capacity(inputs.len());
1752 for input in inputs {
1753 let location = self.input_location(*input)?;
1754 reads.push(location);
1755 operands.push(self.operand(location));
1756 }
1757 let to = self.output_location(output)?;
1758 let spec = ElementwiseSpec {
1759 count: target.elements,
1760 inputs: &operands,
1761 output: self.operand(to),
1762 dims: target.padded_dims(),
1763 strides: &strides,
1764 clamp,
1765 }
1766 .words(broadcast);
1767 self.dispatch(
1768 KernelSpec::Elementwise {
1769 op: lane,
1770 float,
1771 broadcast,
1772 },
1773 spec,
1774 Work::Linear(target.elements),
1775 &reads,
1776 to,
1777 output,
1778 );
1779 Ok(())
1780 }
1781
1782 fn require_zero_constant(&mut self, value: ValueId, dtype: DType) -> Result<(), LoweringError> {
1784 let AnalyzedValueKind::Tensor(tensor) = self.analysis.value(value).kind() else {
1785 return Err(LoweringError::UnsupportedGraph);
1786 };
1787 if tensor.dtype() != dtype {
1788 return Err(LoweringError::UnsupportedGraph);
1789 }
1790 let bytes = self
1791 .analysis
1792 .serialized_constant(value)
1793 .ok_or(LoweringError::UnsupportedGraph)?;
1794 let zero = match dtype {
1795 DType::FP32 => {
1796 bytes.len() % 4 == 0
1797 && bytes.chunks_exact(4).all(|chunk| {
1798 u32::from_le_bytes(chunk.try_into().expect("four bytes")) & 0x7fff_ffff == 0
1799 })
1800 }
1801 DType::FP16 => {
1802 bytes.len() % 2 == 0
1803 && bytes.chunks_exact(2).all(|chunk| {
1804 u16::from_le_bytes(chunk.try_into().expect("two bytes")) & 0x7fff == 0
1805 })
1806 }
1807 DType::FP8E4M3 | DType::FP8E5M2 => bytes.iter().all(|byte| byte & 0x7f == 0),
1810 _ => bytes.iter().all(|byte| *byte == 0),
1811 };
1812 if bytes.is_empty() || !zero {
1813 return Err(LoweringError::UnsupportedGraph);
1814 }
1815 Ok(())
1816 }
1817}
1818
1819fn nan_mode_of(mode: NanPropagationMode) -> Result<NanMode, LoweringError> {
1820 if mode == NanPropagationMode::PROPAGATE {
1821 Ok(NanMode::Propagate)
1822 } else if mode == NanPropagationMode::IGNORE {
1823 Ok(NanMode::Ignore)
1824 } else {
1825 Err(LoweringError::UnsupportedGraph)
1826 }
1827}
1828
1829#[cfg(test)]
1830mod tests {
1831 use super::*;
1832 use virtio_accel_conformance::numerics::{
1833 HEXAGON_UNARY_FP16_CASES, IDENTITY_EDGES_FP16, IDENTITY_EDGES_FP32, IDENTITY_INT8,
1834 MATMUL_FP16, MATMUL_FP32, MAX_POOL2D_FP16, MAX_POOL2D_FP32,
1835 };
1836
1837 const IDENTITY_FP32_LOCAL: &[u8] = include_bytes!("../tests/data/identity-fp32-v1.0.0.tosa");
1838
1839 #[test]
1840 fn targets_validate_and_round_trip() {
1841 for target in [VULKAN_TOSA_TARGET, VULKAN_TOSA_INTEGER_TARGET] {
1842 assert_eq!(target.validate(), Ok(target));
1843 assert_eq!(Target::from_identity(target.to_identity()), Ok(target));
1844 }
1845 assert_ne!(VULKAN_TOSA_TARGET, VULKAN_TOSA_INTEGER_TARGET);
1846 }
1847
1848 #[test]
1849 fn capability_names_the_shared_fp32_operator_set() {
1850 assert_eq!(FLOAT_OPERATORS.len(), 42);
1851 for op in [
1852 Op::IDENTITY,
1853 Op::MATMUL,
1854 Op::MAX_POOL2D,
1855 Op::ARGMAX,
1856 Op::ERF,
1857 Op::CONCAT,
1858 Op::TRANSPOSE,
1859 ] {
1860 assert!(supports_tosa_operator(op), "{op:?}");
1861 }
1862 for op in [Op::CONV2D, Op::AVG_POOL2D, Op::CAST, Op::RESCALE, Op::PAD] {
1863 assert!(!supports_tosa_operator(op), "{op:?}");
1864 }
1865 assert!(supports_tosa_dtype(DType::FP32));
1866 assert!(supports_tosa_dtype(DType::FP16));
1867 assert!(supports_tosa_dtype(DType::BOOL));
1868 assert!(supports_tosa_dtype(DType::INT32));
1869 assert!(!supports_tosa_dtype(DType::INT8));
1871 assert!(VULKAN_TOSA_CAPABILITY.supports_dtype(DType::INT8, ValueRoles::CONSTANT));
1872 assert!(!VULKAN_TOSA_CAPABILITY.supports_dtype(DType::INT8, ValueRoles::INTERMEDIATE));
1873 assert_eq!(VULKAN_TOSA_CAPABILITY.target, VULKAN_TOSA_TARGET);
1874 }
1875
1876 #[test]
1877 fn lowers_the_local_fp32_identity_artifact() {
1878 let plan = lower_tosa(IDENTITY_FP32_LOCAL, VULKAN_TOSA_TARGET).unwrap();
1879 assert_eq!(plan.slots.len(), 2);
1880 assert_eq!(plan.slot(0).unwrap().role, SlotRole::Input);
1881 assert_eq!(plan.slot(1).unwrap().role, SlotRole::Output);
1882 assert_eq!(plan.slot(0).unwrap().byte_len, 4);
1883 assert_eq!(plan.slot(1).unwrap().storage, Storage::Word);
1884 assert!(plan.slot(2).is_none());
1885 assert_eq!(plan.arena_bytes, 0);
1886 assert!(plan.constants.is_empty());
1887 assert_eq!(plan.dispatches.len(), 1);
1888 let dispatch = &plan.dispatches[0];
1889 assert_eq!(
1890 dispatch.kernel,
1891 KernelSpec::Move {
1892 storage: Storage::Word,
1893 contiguous: true
1894 }
1895 );
1896 assert_eq!(dispatch.work, Work::Linear(1));
1897 assert!(!dispatch.barrier_before);
1898 assert_eq!(dispatch.spec, vec![0, 0, 1, 0, 1]);
1900 }
1901
1902 #[test]
1903 fn lowers_the_shared_fp32_edge_identity_artifact() {
1904 let plan = lower_tosa(IDENTITY_EDGES_FP32.artifact, VULKAN_TOSA_TARGET).unwrap();
1905 let expected = IDENTITY_EDGES_FP32.inputs[0].values.len();
1906 assert_eq!(plan.dispatches[0].work, Work::Linear(expected as u32));
1907 assert_eq!(plan.slot(1).unwrap().byte_len as usize, expected * 4);
1908 }
1909
1910 #[test]
1911 fn rejects_other_targets_before_parsing() {
1912 assert_eq!(
1913 lower_tosa(IDENTITY_FP32_LOCAL, VULKAN_TOSA_INTEGER_TARGET),
1914 Err(LoweringError::UnsupportedTarget)
1915 );
1916 assert_eq!(
1917 lower_tosa(IDENTITY_INT8.artifact, VULKAN_TOSA_INTEGER_TARGET),
1918 Err(LoweringError::UnsupportedTarget)
1919 );
1920 }
1921
1922 #[test]
1923 fn rejects_mistyped_identity_graphs_loudly() {
1924 assert!(matches!(
1926 lower_tosa(IDENTITY_INT8.artifact, VULKAN_TOSA_TARGET),
1927 Err(LoweringError::UnsupportedType(DType::INT8) | LoweringError::Analysis(_))
1928 ));
1929 }
1930
1931 #[test]
1932 fn fp16_capability_extends_the_fp32_boundary() {
1933 for dtype in [DType::FP32, DType::FP16, DType::BOOL, DType::INT32] {
1934 assert!(
1935 VULKAN_TOSA_FP16_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT),
1936 "{dtype:?}"
1937 );
1938 assert!(
1939 VULKAN_TOSA_FP16_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT),
1940 "{dtype:?}"
1941 );
1942 }
1943 assert_eq!(VULKAN_TOSA_FP16_CAPABILITY.target, VULKAN_TOSA_TARGET);
1944 assert_eq!(VULKAN_TOSA_FP16_CAPABILITY.operators, FLOAT_OPERATORS);
1945 assert!(!VULKAN_TOSA_CAPABILITY.supports_dtype(DType::FP16, ValueRoles::INPUT));
1946 }
1947
1948 #[test]
1949 fn lowers_the_shared_fp16_artifacts() {
1950 let plan = lower_tosa(IDENTITY_EDGES_FP16.artifact, VULKAN_TOSA_TARGET).unwrap();
1951 let expected = IDENTITY_EDGES_FP16.inputs[0].bits.len() as u32;
1952 assert_eq!(plan.slot(0).unwrap().byte_len, u64::from(expected) * 2);
1953 assert_eq!(plan.slot(0).unwrap().storage, Storage::Half);
1954 assert_eq!(plan.dispatches.len(), 1);
1955 assert_eq!(
1956 plan.dispatches[0].kernel,
1957 KernelSpec::Move {
1958 storage: Storage::Half,
1959 contiguous: true
1960 }
1961 );
1962 assert_eq!(plan.dispatches[0].work, Work::Linear(expected.div_ceil(2)));
1964
1965 let plan = lower_tosa(MATMUL_FP16.artifact, VULKAN_TOSA_TARGET).unwrap();
1967 assert_eq!(plan.dispatches.len(), 2);
1968 assert_eq!(
1969 plan.dispatches[0].kernel,
1970 KernelSpec::Cast {
1971 input: Storage::Half,
1972 output: Storage::Word
1973 }
1974 );
1975 assert_eq!(plan.dispatches[0].work, Work::Linear(6));
1976 assert_eq!(
1977 plan.dispatches[1].kernel,
1978 KernelSpec::MatmulStream {
1979 rhs: Storage::Half,
1980 output: Storage::Half
1981 }
1982 );
1983 assert!(plan.dispatches[1].barrier_before, "reads the widened lhs");
1984 assert_eq!(plan.arena_bytes, ARENA_ALIGNMENT);
1985 assert_eq!(plan.slot(0).unwrap().byte_len, 6 * 2);
1986 assert_eq!(plan.slot(2).unwrap().byte_len, 4 * 2);
1987
1988 let plan = lower_tosa(MAX_POOL2D_FP16.artifact, VULKAN_TOSA_TARGET).unwrap();
1989 assert_eq!(
1990 plan.dispatches[0].kernel,
1991 KernelSpec::MaxPool {
1992 nan_mode: NanMode::Propagate,
1993 float: Storage::Half
1994 }
1995 );
1996 }
1997
1998 #[test]
1999 fn fp16_clamp_bounds_arrive_widened_to_binary32() {
2000 let clamp = HEXAGON_UNARY_FP16_CASES
2004 .iter()
2005 .find(|case| case.name == "clamp-fp16")
2006 .expect("the shared clamp-fp16 case");
2007 let plan = lower_tosa(clamp.artifact, VULKAN_TOSA_TARGET).unwrap();
2008 let dispatch = &plan.dispatches[0];
2009 assert_eq!(
2010 dispatch.kernel,
2011 KernelSpec::Elementwise {
2012 op: ElementwiseOp::Clamp(NanMode::Propagate),
2013 float: Storage::Half,
2014 broadcast: false
2015 }
2016 );
2017 assert_eq!(dispatch.spec[dispatch.spec.len() - 2], (-1.0_f32).to_bits());
2018 assert_eq!(dispatch.spec[dispatch.spec.len() - 1], 1.0_f32.to_bits());
2019 }
2020
2021 #[test]
2022 fn fp16_negate_zero_points_are_consumed_at_admission() {
2023 let negate = HEXAGON_UNARY_FP16_CASES
2024 .iter()
2025 .find(|case| case.name == "negate-fp16")
2026 .expect("the shared negate-fp16 case");
2027 let plan = lower_tosa(negate.artifact, VULKAN_TOSA_TARGET).unwrap();
2028 assert_eq!(
2029 plan.dispatches[0].kernel,
2030 KernelSpec::Elementwise {
2031 op: ElementwiseOp::Negate,
2032 float: Storage::Half,
2033 broadcast: false
2034 }
2035 );
2036 assert_eq!(plan.arena_bytes, 0);
2037 }
2038
2039 #[test]
2040 fn lowers_the_shared_fp32_matmul_artifact() {
2041 let plan = lower_tosa(MATMUL_FP32.artifact, VULKAN_TOSA_TARGET).unwrap();
2042 assert_eq!(plan.slots.len(), 3);
2043 assert_eq!(plan.slot(0).unwrap().byte_len, 6 * 4);
2044 assert_eq!(plan.slot(1).unwrap().byte_len, 6 * 4);
2045 assert_eq!(plan.slot(2).unwrap().role, SlotRole::Output);
2046 assert_eq!(plan.slot(2).unwrap().byte_len, 4 * 4);
2047 assert_eq!(plan.dispatches.len(), 1);
2048 let dispatch = &plan.dispatches[0];
2049 assert_eq!(
2051 dispatch.kernel,
2052 KernelSpec::MatmulStream {
2053 rhs: Storage::Word,
2054 output: Storage::Word
2055 }
2056 );
2057 assert_eq!(dispatch.work, Work::MatmulStream { n: 2, batch: 1 });
2058 assert_eq!(dispatch.spec, vec![0, 0, 1, 0, 2, 0, 2, 2, 3, 1]);
2060 assert_eq!(plan.arena_bytes, 0);
2062 }
2063
2064 #[test]
2065 fn lowers_the_shared_fp32_max_pool_artifact() {
2066 let plan = lower_tosa(MAX_POOL2D_FP32.artifact, VULKAN_TOSA_TARGET).unwrap();
2067 let dispatch = &plan.dispatches[0];
2068 assert_eq!(
2069 dispatch.kernel,
2070 KernelSpec::MaxPool {
2071 nan_mode: NanMode::Propagate,
2072 float: Storage::Word
2073 }
2074 );
2075 assert_eq!(dispatch.work, Work::Linear(8));
2076 assert_eq!(
2078 dispatch.spec,
2079 vec![0, 0, 1, 0, 1, 4, 4, 2, 2, 2, 2, 2, 2, 2, 0, 0]
2080 );
2081 }
2082
2083 #[test]
2084 fn rejects_garbage_as_a_parse_error() {
2085 assert!(matches!(
2086 lower_tosa(b"not a flatbuffer", VULKAN_TOSA_TARGET),
2087 Err(LoweringError::Parse(_))
2088 ));
2089 }
2090
2091 #[test]
2092 fn arena_regions_pack_by_lifetime() {
2093 let bytes = IDENTITY_FP32_LOCAL;
2095 let model = parse(bytes).unwrap();
2096 let analysis = model.analyze_for(VULKAN_TOSA_TARGET).unwrap();
2097 let mut lowering = Lowering::new(&analysis, VULKAN_TOSA_CAPABILITY).unwrap();
2098 lowering.position = 0;
2099 let a = lowering.allocate_region(100, 1).unwrap();
2100 let b = lowering.allocate_region(100, 5).unwrap();
2101 assert_eq!(lowering.regions[a].offset, 0);
2102 assert_eq!(lowering.regions[b].offset, ARENA_ALIGNMENT);
2103 lowering.position = 2;
2105 let c = lowering.allocate_region(ARENA_ALIGNMENT, 9).unwrap();
2106 assert_eq!(lowering.regions[c].offset, 0);
2107 let d = lowering.allocate_region(1, 9).unwrap();
2108 assert_eq!(lowering.regions[d].offset, 2 * ARENA_ALIGNMENT);
2109 assert_eq!(lowering.arena_bytes, 3 * ARENA_ALIGNMENT);
2110 }
2111
2112 #[test]
2116 fn constants_never_share_bytes_with_earlier_intermediates() {
2117 let model = parse(IDENTITY_FP32_LOCAL).unwrap();
2118 let analysis = model.analyze_for(VULKAN_TOSA_TARGET).unwrap();
2119 let mut lowering = Lowering::new(&analysis, VULKAN_TOSA_CAPABILITY).unwrap();
2120 lowering.position = 0;
2121 let a = lowering.allocate_region(64, 1).unwrap();
2122 assert_eq!(lowering.regions[a].offset, 0);
2123 lowering.position = 2;
2124 let constant = lowering
2125 .allocate_region_written_at(64, u32::MAX, 0)
2126 .unwrap();
2127 assert_eq!(lowering.regions[constant].offset, ARENA_ALIGNMENT);
2128 let intermediate = lowering.allocate_region(64, 3).unwrap();
2129 assert_eq!(lowering.regions[intermediate].offset, 0);
2130 }
2131
2132 #[test]
2136 fn reused_arena_bytes_force_a_barrier_before_the_new_writer() {
2137 let model = parse(IDENTITY_FP32_LOCAL).unwrap();
2138 let analysis = model.analyze_for(VULKAN_TOSA_TARGET).unwrap();
2139 let mut lowering = Lowering::new(&analysis, VULKAN_TOSA_CAPABILITY).unwrap();
2140 let values: Vec<ValueId> = analysis.values().iter().map(|value| value.id()).collect();
2141 let (a_value, b_value) = (values[0], values[1]);
2142 let kernel = KernelSpec::Move {
2143 storage: Storage::Word,
2144 contiguous: true,
2145 };
2146
2147 lowering.position = 0;
2149 let a = lowering.allocate_region(64, 1).unwrap();
2150 lowering.dispatch(
2151 kernel,
2152 Vec::new(),
2153 Work::Linear(16),
2154 &[Location::Slot(0)],
2155 Location::Region(a),
2156 a_value,
2157 );
2158 lowering.position = 1;
2160 lowering.dispatch(
2161 kernel,
2162 Vec::new(),
2163 Work::Linear(16),
2164 &[Location::Region(a)],
2165 Location::Slot(1),
2166 b_value,
2167 );
2168 lowering.position = 2;
2170 let b = lowering.allocate_region(64, 3).unwrap();
2171 assert_eq!(lowering.regions[b].offset, lowering.regions[a].offset);
2172 assert_ne!(a, b, "a fresh region index over the same bytes");
2173 lowering.dispatch(
2174 kernel,
2175 Vec::new(),
2176 Work::Linear(16),
2177 &[Location::Slot(0)],
2178 Location::Region(b),
2179 b_value,
2180 );
2181
2182 let barriers: Vec<bool> = lowering
2183 .dispatches
2184 .iter()
2185 .map(|dispatch| dispatch.barrier_before)
2186 .collect();
2187 assert_eq!(
2188 barriers,
2189 [false, true, true],
2190 "the reader depends on the producer (RAW); the new writer on the reader (WAR)"
2191 );
2192
2193 lowering.position = 3;
2195 let c = lowering.allocate_region(64, 4).unwrap();
2196 assert_ne!(lowering.regions[c].offset, lowering.regions[b].offset);
2197 lowering.dispatch(
2198 kernel,
2199 Vec::new(),
2200 Work::Linear(16),
2201 &[Location::Slot(0)],
2202 Location::Region(c),
2203 a_value,
2204 );
2205 assert!(!lowering.dispatches[3].barrier_before);
2206 assert!(
2207 MemKey::Arena {
2208 offset: 0,
2209 end: 256
2210 }
2211 .overlaps(MemKey::Arena {
2212 offset: 255,
2213 end: 512
2214 })
2215 );
2216 assert!(
2217 !MemKey::Arena {
2218 offset: 0,
2219 end: 256
2220 }
2221 .overlaps(MemKey::Arena {
2222 offset: 256,
2223 end: 512
2224 })
2225 );
2226 assert!(!MemKey::Slot(0).overlaps(MemKey::Arena {
2227 offset: 0,
2228 end: 256
2229 }));
2230 }
2231
2232 #[test]
2233 fn padding_and_strides_follow_row_major_order() {
2234 let shape = TensorShape {
2235 dtype: DType::FP32,
2236 dims: vec![2, 3, 4],
2237 elements: 24,
2238 };
2239 assert_eq!(shape.strides(), vec![12, 4, 1]);
2240 assert_eq!(shape.padded_dims(), [1, 1, 1, 2, 3, 4]);
2241 assert_eq!(pad_leading(&shape.strides(), 0), [0, 0, 0, 12, 4, 1]);
2242 }
2243}