1use super::*;
4
5#[derive(Clone, Copy)]
6enum BarrierBehavior {
7 Ignore,
8 Suspend,
9 Reject,
10}
11
12enum FrameOutcome {
13 Barrier(usize),
14 Complete(Option<Value>),
15}
16
17struct ExecutionFrame {
19 function_index: usize,
20 registers: Vec<Option<Value>>,
21 locals: Vec<Option<Value>>,
22 instruction_index: usize,
23}
24
25struct TaskLane<'a> {
27 frame: ExecutionFrame,
28 state: ExecutionState<'a>,
29}
30
31impl ExecutableProgram {
32 pub fn run_main(&self, descriptors: &mut DescriptorBindings<'_>) -> Result<(), VmError> {
34 self.run_main_with_config(descriptors, &ExecutionConfig::default())
35 }
36
37 pub fn run_main_with_config(
42 &self,
43 descriptors: &mut DescriptorBindings<'_>,
44 config: &ExecutionConfig,
45 ) -> Result<(), VmError> {
46 let mut state = ExecutionState::new(config);
47 state.enter_call()?;
48 let mut frame = self.create_frame(self.main_function, &[])?;
49 let outcome = self.execute_frame(&mut frame, descriptors, &mut state, BarrierBehavior::Ignore);
50 state.leave_call();
51 let FrameOutcome::Complete(return_value) = outcome? else {
52 unreachable!(
53 "Unexpected suspended main frame. The most likely cause is that ignored workgroup barriers returned a scheduler-visible outcome."
54 )
55 };
56 if return_value.is_some() {
57 return Err(VmError::UnsupportedMainSignature {
58 message: "Main functions must not return a value".to_string(),
59 });
60 }
61 Ok(())
62 }
63
64 pub fn run_task_workgroup(
71 &self,
72 descriptors: &mut DescriptorBindings<'_>,
73 configs: &[ExecutionConfig],
74 ) -> Result<(), VmError> {
75 descriptors.begin_task_workgroup();
76 if configs.is_empty() {
77 return Ok(());
78 }
79 let mut lanes = configs
80 .iter()
81 .map(|config| {
82 let mut state = ExecutionState::new(config);
83 state.enter_call()?;
84 Ok(TaskLane {
85 frame: self.create_frame(self.main_function, &[])?,
86 state,
87 })
88 })
89 .collect::<Result<Vec<_>, VmError>>()?;
90
91 loop {
92 let mut expected_barrier = None;
93 let mut barrier_count = 0;
94 let mut completion_count = 0;
95 let mut first_completed_lane = None;
96 for (lane_index, lane) in lanes.iter_mut().enumerate() {
97 match self.execute_frame(&mut lane.frame, descriptors, &mut lane.state, BarrierBehavior::Suspend)? {
98 FrameOutcome::Barrier(instruction_index) => {
99 if let Some(expected) = expected_barrier {
100 if expected != instruction_index {
101 return Err(VmError::DivergentWorkgroupBarrier {
102 lane: lane_index,
103 expected_instruction: expected,
104 found_instruction: Some(instruction_index),
105 });
106 }
107 } else {
108 expected_barrier = Some(instruction_index);
109 }
110 barrier_count += 1;
111 }
112 FrameOutcome::Complete(return_value) => {
113 lane.state.leave_call();
114 if return_value.is_some() {
115 return Err(VmError::UnsupportedMainSignature {
116 message: "Main functions must not return a value".to_string(),
117 });
118 }
119 first_completed_lane.get_or_insert(lane_index);
120 completion_count += 1;
121 }
122 }
123 }
124
125 if barrier_count == lanes.len() {
126 continue;
127 }
128 if completion_count == lanes.len() {
129 return Ok(());
130 }
131 let (Some(divergent_lane), Some(expected_instruction)) = (first_completed_lane, expected_barrier) else {
132 unreachable!(
133 "Invalid task workgroup phase accounting. The most likely cause is that a lane outcome was not counted as either a barrier or completion."
134 )
135 };
136 return Err(VmError::DivergentWorkgroupBarrier {
137 lane: divergent_lane,
138 expected_instruction,
139 found_instruction: None,
140 });
141 }
142 }
143
144 fn execute_function(
146 &self,
147 function_index: usize,
148 arguments: &[Value],
149 descriptors: &mut DescriptorBindings<'_>,
150 state: &mut ExecutionState<'_>,
151 barrier_behavior: BarrierBehavior,
152 ) -> Result<Option<Value>, VmError> {
153 state.enter_call()?;
154 let result = (|| {
155 let mut frame = self.create_frame(function_index, arguments)?;
156 match self.execute_frame(&mut frame, descriptors, state, barrier_behavior)? {
157 FrameOutcome::Complete(value) => Ok(value),
158 FrameOutcome::Barrier(_) => unreachable!(
159 "Unexpected nested barrier suspension. The most likely cause is that nested execution stopped rejecting workgroup barriers."
160 ),
161 }
162 })();
163 state.leave_call();
164 result
165 }
166
167 fn create_frame(&self, function_index: usize, arguments: &[Value]) -> Result<ExecutionFrame, VmError> {
169 let function = self
170 .functions
171 .get(function_index)
172 .ok_or_else(|| VmError::UnsupportedExpression {
173 message: format!("Unknown function index {}", function_index),
174 })?;
175 if arguments.len() != function.parameter_count {
176 return Err(VmError::CallArgumentMismatch {
177 expected: function.parameter_count,
178 found: arguments.len(),
179 });
180 }
181 let mut locals = vec![None; function.local_types.len()];
182 for (index, argument) in arguments.iter().enumerate() {
183 locals[index] = Some(argument.clone());
184 }
185 Ok(ExecutionFrame {
186 function_index,
187 registers: vec![None; function.register_count],
188 locals,
189 instruction_index: 0,
190 })
191 }
192
193 fn execute_frame(
195 &self,
196 frame: &mut ExecutionFrame,
197 descriptors: &mut DescriptorBindings<'_>,
198 state: &mut ExecutionState<'_>,
199 barrier_behavior: BarrierBehavior,
200 ) -> Result<FrameOutcome, VmError> {
201 let function = self
202 .functions
203 .get(frame.function_index)
204 .ok_or_else(|| VmError::UnsupportedExpression {
205 message: format!("Unknown function index {}", frame.function_index),
206 })?;
207 let registers = &mut frame.registers;
208 let locals = &mut frame.locals;
209
210 while frame.instruction_index < function.instructions.len() {
211 state.consume_instruction()?;
212 let instruction = &function.instructions[frame.instruction_index];
213 match instruction {
214 Instruction::LoadLiteral { register, value } => {
215 registers[*register] = Some(value.clone());
216 }
217 Instruction::Construct {
218 register,
219 value_type,
220 components,
221 } => {
222 let values = components
223 .iter()
224 .map(|component| read_register(registers, *component))
225 .collect::<Result<Vec<_>, _>>()?;
226 registers[*register] = Some(construct_value(value_type, &values)?);
227 }
228 Instruction::Extract {
229 register,
230 source,
231 index,
232 value_type,
233 } => {
234 let source = read_register(registers, *source)?;
235 registers[*register] = Some(extract_value(&source, *index, value_type)?);
236 }
237 Instruction::ExtractDynamic {
238 register,
239 source,
240 index,
241 count,
242 value_type,
243 } => {
244 let source = read_register(registers, *source)?;
245 let index = expect_u32(read_register(registers, *index)?)? as usize;
246 if index >= *count {
247 return Err(VmError::BufferArrayIndexOutOfBounds { index, count: *count });
248 }
249 registers[*register] = Some(extract_value(&source, index, value_type)?);
250 }
251 Instruction::Arithmetic {
252 register,
253 operator,
254 left,
255 right,
256 } => {
257 let left = read_register(registers, *left)?;
258 let right = read_register(registers, *right)?;
259 registers[*register] = Some(apply_arithmetic(*operator, &left, &right)?);
260 }
261 Instruction::Compare {
262 register,
263 operator,
264 left,
265 right,
266 } => {
267 let left = read_register(registers, *left)?;
268 let right = read_register(registers, *right)?;
269 registers[*register] = Some(apply_comparison(*operator, &left, &right)?);
270 }
271 Instruction::JumpIfZero { register, target } => {
272 let value = read_register(registers, *register)?;
273 if is_zero_value(&value)? {
274 frame.instruction_index = *target;
275 continue;
276 }
277 }
278 Instruction::Jump { target } => {
279 frame.instruction_index = *target;
280 continue;
281 }
282 Instruction::DotProduct { register, left, right } => {
283 let left = read_register(registers, *left)?;
284 let right = read_register(registers, *right)?;
285 registers[*register] = Some(apply_dot_product(&left, &right)?);
286 }
287 Instruction::CrossProduct { register, left, right } => {
288 let left = read_register(registers, *left)?;
289 let right = read_register(registers, *right)?;
290 registers[*register] = Some(apply_cross_product(&left, &right)?);
291 }
292 Instruction::Length { register, value } => {
293 let value = read_register(registers, *value)?;
294 registers[*register] = Some(apply_length(&value)?);
295 }
296 Instruction::Normalize { register, value } => {
297 let value = read_register(registers, *value)?;
298 registers[*register] = Some(apply_normalize(&value)?);
299 }
300 Instruction::Reflect {
301 register,
302 incident,
303 normal,
304 } => {
305 let incident = read_register(registers, *incident)?;
306 let normal = read_register(registers, *normal)?;
307 registers[*register] = Some(apply_reflect(&incident, &normal)?);
308 }
309 Instruction::UnaryScalar {
310 register,
311 operator,
312 value,
313 } => {
314 let value = read_register(registers, *value)?;
315 registers[*register] = Some(apply_scalar_unary(*operator, &value)?);
316 }
317 Instruction::BinaryScalar {
318 register,
319 operator,
320 left,
321 right,
322 } => {
323 let left = read_register(registers, *left)?;
324 let right = read_register(registers, *right)?;
325 registers[*register] = Some(apply_scalar_binary(*operator, &left, &right)?);
326 }
327 Instruction::TernaryScalar {
328 register,
329 operator,
330 first,
331 second,
332 third,
333 } => {
334 let first = read_register(registers, *first)?;
335 let second = read_register(registers, *second)?;
336 let third = read_register(registers, *third)?;
337 registers[*register] = Some(apply_scalar_ternary(*operator, &first, &second, &third)?);
338 }
339 Instruction::ThreadIdx { register } => {
340 registers[*register] = Some(Value::U32(state.config.thread_idx()));
341 }
342 Instruction::ThreadPosition { register } => {
343 registers[*register] = Some(Value::U32(state.config.thread_position()));
344 }
345 Instruction::ThreadId { register } => {
346 registers[*register] = Some(Value::Vec2U(state.config.thread_id()));
347 }
348 Instruction::ThreadgroupPosition { register } => {
349 registers[*register] = Some(Value::U32(state.config.threadgroup_position()));
350 }
351 Instruction::LoadTaskPayload {
352 register,
353 name,
354 index,
355 count,
356 value_type,
357 } => {
358 let index = read_buffer_array_index(registers, *index, *count)?;
359 let value = descriptors.task_payload_value(name, index)?;
360 if !value.matches_type(value_type) {
361 return Err(VmError::TypeMismatch {
362 expected: value_type.name().to_string(),
363 found: value.value_type().name().to_string(),
364 });
365 }
366 registers[*register] = Some(value);
367 }
368 Instruction::StoreTaskPayload {
369 name,
370 index,
371 count,
372 value_type,
373 value,
374 } => {
375 let index = expect_u32(read_register(registers, *index)?)? as usize;
376 let value = read_register(registers, *value)?;
377 if !value.matches_type(value_type) {
378 return Err(VmError::TypeMismatch {
379 expected: value_type.name().to_string(),
380 found: value.value_type().name().to_string(),
381 });
382 }
383 descriptors.task_outputs_mut()?.write_payload(name, index, *count, value)?;
384 }
385 Instruction::LoadWorkgroup {
386 register,
387 name,
388 value_type,
389 } => {
390 let value = descriptors.workgroup_state_mut()?.load(name, value_type)?;
391 registers[*register] = Some(value);
392 }
393 Instruction::StoreWorkgroup { name, value_type, value } => {
394 let value = read_register(registers, *value)?;
395 descriptors.workgroup_state_mut()?.store(name, value_type, value)?;
396 }
397 Instruction::AtomicAddWorkgroup { register, name, value } => {
398 let value = expect_u32(read_register(registers, *value)?)?;
399 let previous = descriptors.workgroup_state_mut()?.atomic_add_u32(name, value)?;
400 registers[*register] = Some(Value::U32(previous));
401 }
402 Instruction::WorkgroupBarrier => match barrier_behavior {
403 BarrierBehavior::Ignore => {
404 }
406 BarrierBehavior::Suspend => {
407 let barrier_instruction = frame.instruction_index;
408 frame.instruction_index += 1;
409 return Ok(FrameOutcome::Barrier(barrier_instruction));
410 }
411 BarrierBehavior::Reject => {
412 return Err(VmError::UnsupportedStatement {
413 message: "Workgroup barriers inside called functions cannot participate in task rendezvous"
414 .to_string(),
415 });
416 }
417 },
418 Instruction::SetTaskMeshOutputCount { count } => {
419 let count = expect_u32(read_register(registers, *count)?)?;
420 if count > state.config.max_task_mesh_output_count() {
421 return Err(VmError::TaskMeshOutputCountLimitExceeded {
422 requested: count,
423 limit: state.config.max_task_mesh_output_count(),
424 });
425 }
426 descriptors.task_outputs_mut()?.set_mesh_output_count(count);
427 }
428 Instruction::SetMeshOutputCounts {
429 vertex_count,
430 primitive_count,
431 } => {
432 let vertex_count = expect_u32(read_register(registers, *vertex_count)?)?;
433 let primitive_count = expect_u32(read_register(registers, *primitive_count)?)?;
434 descriptors.mesh_outputs_mut()?.set_counts(
435 vertex_count,
436 primitive_count,
437 state.config.max_mesh_vertex_count(),
438 state.config.max_mesh_primitive_count(),
439 state.config.thread_idx() == 0,
440 )?;
441 }
442 Instruction::SetMeshVertexPosition { index, position } => {
443 let index = expect_u32(read_register(registers, *index)?)? as usize;
444 let position = read_register(registers, *position)?;
445 let Value::Vec4F(position) = position else {
446 return Err(VmError::TypeMismatch {
447 expected: ValueType::Vec4F.name().to_string(),
448 found: position.value_type().name().to_string(),
449 });
450 };
451 let outputs = descriptors.mesh_outputs_mut()?;
452 let count = outputs.vertex_positions.len();
453 let destination = outputs
454 .vertex_positions
455 .get_mut(index)
456 .ok_or(VmError::MeshOutputIndexOutOfBounds {
457 kind: "vertex",
458 index,
459 count,
460 })?;
461 *destination = position;
462 }
463 Instruction::SetMeshTriangle { index, triangle } => {
464 let index = expect_u32(read_register(registers, *index)?)? as usize;
465 let triangle = read_register(registers, *triangle)?;
466 let Value::Vec3U(triangle) = triangle else {
467 return Err(VmError::TypeMismatch {
468 expected: ValueType::Vec3U.name().to_string(),
469 found: triangle.value_type().name().to_string(),
470 });
471 };
472 let outputs = descriptors.mesh_outputs_mut()?;
473 let count = outputs.triangles.len();
474 let destination = outputs.triangles.get_mut(index).ok_or(VmError::MeshOutputIndexOutOfBounds {
475 kind: "primitive",
476 index,
477 count,
478 })?;
479 *destination = triangle;
480 }
481 Instruction::LoadLocal { register, local } => {
482 let value = locals
483 .get(*local)
484 .and_then(Option::clone)
485 .ok_or(VmError::UninitializedLocal { local: *local })?;
486 registers[*register] = Some(value);
487 }
488 Instruction::StoreLocal { local, register } => {
489 let value = read_register(registers, *register)?;
490 locals[*local] = Some(value.clone());
491 }
492 Instruction::LoadBuffer {
493 register,
494 slot,
495 offset,
496 value_type,
497 } => {
498 let value = if *slot == PUSH_CONSTANT_SLOT {
499 descriptors.push_constant_mut()?.read_value(*offset, value_type)?
500 } else {
501 descriptors.buffer_mut(*slot)?.read_value(*offset, value_type)?
502 };
503 registers[*register] = Some(value);
504 }
505 Instruction::LoadBufferIndexed {
506 register,
507 slot,
508 offset,
509 stride,
510 count,
511 index,
512 value_type,
513 } => {
514 let index = read_buffer_array_index(registers, *index, *count)?;
515 let value = if *slot == PUSH_CONSTANT_SLOT {
516 descriptors
517 .push_constant_mut()?
518 .read_value(*offset + *stride * index, value_type)?
519 } else {
520 descriptors
521 .buffer_mut(*slot)?
522 .read_value(*offset + *stride * index, value_type)?
523 };
524 registers[*register] = Some(value);
525 }
526 Instruction::FetchTexture { register, slot, coord } => {
527 let coord = read_register(registers, *coord)?;
528 let Value::Vec2U(coord) = coord else {
529 return Err(VmError::TypeMismatch {
530 expected: ValueType::Vec2U.name().to_string(),
531 found: coord.value_type().name().to_string(),
532 });
533 };
534
535 let slot = resolve_resource_slot(*slot, registers)?;
536 registers[*register] = Some(descriptors.texture_mut(slot)?.fetch(coord)?);
537 }
538 Instruction::FetchTextureU32 { register, slot, coord } => {
539 let coord = read_register(registers, *coord)?;
540 let Value::Vec2U(coord) = coord else {
541 return Err(VmError::TypeMismatch {
542 expected: ValueType::Vec2U.name().to_string(),
543 found: coord.value_type().name().to_string(),
544 });
545 };
546 let slot = resolve_resource_slot(*slot, registers)?;
547 registers[*register] = Some(descriptors.texture_mut(slot)?.fetch_u32(coord)?);
548 }
549 Instruction::SampleTexture { register, slot, uv } => {
550 let uv = read_register(registers, *uv)?;
551 let Value::Vec2F(uv) = uv else {
552 return Err(VmError::TypeMismatch {
553 expected: ValueType::Vec2F.name().to_string(),
554 found: uv.value_type().name().to_string(),
555 });
556 };
557
558 let slot = resolve_resource_slot(*slot, registers)?;
559 registers[*register] = Some(descriptors.texture_mut(slot)?.sample(uv)?);
560 }
561 Instruction::SampleTexture3D { register, slot, uvw } => {
562 let uvw = read_register(registers, *uvw)?;
563 let Value::Vec3F(uvw) = uvw else {
564 return Err(VmError::TypeMismatch {
565 expected: ValueType::Vec3F.name().to_string(),
566 found: uvw.value_type().name().to_string(),
567 });
568 };
569 let slot = resolve_resource_slot(*slot, registers)?;
570 registers[*register] = Some(descriptors.texture_mut(slot)?.sample_3d(uvw)?);
571 }
572 Instruction::TextureSize { register, slot } => {
573 let slot = resolve_resource_slot(*slot, registers)?;
574 let texture = descriptors.texture_mut(slot)?;
575 registers[*register] = Some(Value::Vec2U([texture.width, texture.height]));
576 }
577 Instruction::ImageSize { register, slot } => {
578 let slot = resolve_resource_slot(*slot, registers)?;
579 let image = descriptors.image_mut(slot)?;
580 registers[*register] = Some(Value::Vec2U([image.width, image.height]));
581 }
582 Instruction::LoadImage { register, slot, coord } => {
583 let coord = expect_vec2u(read_register(registers, *coord)?)?;
584 let slot = resolve_resource_slot(*slot, registers)?;
585 registers[*register] = Some(descriptors.image_mut(slot)?.fetch(coord)?);
586 }
587 Instruction::LoadImageU32 { register, slot, coord } => {
588 let coord = expect_vec2u(read_register(registers, *coord)?)?;
589 let slot = resolve_resource_slot(*slot, registers)?;
590 registers[*register] = Some(descriptors.image_mut(slot)?.fetch_u32(coord)?);
591 }
592 Instruction::GuardImageBounds { slot, coord } => {
593 let coord = expect_vec2u(read_register(registers, *coord)?)?;
594 let slot = resolve_resource_slot(*slot, registers)?;
595 if !descriptors.image_mut(slot)?.contains_2d(coord) {
596 return Ok(FrameOutcome::Complete(None));
597 }
598 }
599 Instruction::ImageAtomicOr {
600 register,
601 slot,
602 coord,
603 value,
604 } => {
605 let coord = expect_vec2u(read_register(registers, *coord)?)?;
606 let value = expect_u32(read_register(registers, *value)?)?;
607 let slot = resolve_resource_slot(*slot, registers)?;
608 let previous = descriptors.image_mut(slot)?.atomic_or(coord, value)?;
609 registers[*register] = Some(Value::U32(previous));
610 }
611 Instruction::WriteImage { slot, coord, value } => {
612 let coord = read_register(registers, *coord)?;
613 let Value::Vec2U(coord) = coord else {
614 return Err(VmError::TypeMismatch {
615 expected: ValueType::Vec2U.name().to_string(),
616 found: coord.value_type().name().to_string(),
617 });
618 };
619
620 let value = read_register(registers, *value)?;
621 let Value::Vec4F(value) = value else {
622 return Err(VmError::TypeMismatch {
623 expected: ValueType::Vec4F.name().to_string(),
624 found: value.value_type().name().to_string(),
625 });
626 };
627
628 let slot = resolve_resource_slot(*slot, registers)?;
629 descriptors.image_mut(slot)?.write(coord, value)?;
630 }
631 Instruction::StoreBuffer {
632 slot,
633 offset,
634 value_type,
635 register,
636 } => {
637 let value = read_register(registers, *register)?;
638 descriptors.buffer_mut(*slot)?.write_value(*offset, value_type, &value)?;
639 }
640 Instruction::StoreBufferIndexed {
641 slot,
642 offset,
643 stride,
644 count,
645 index,
646 value_type,
647 register,
648 } => {
649 let index = read_buffer_array_index(registers, *index, *count)?;
650 let value = read_register(registers, *register)?;
651 descriptors
652 .buffer_mut(*slot)?
653 .write_value(*offset + *stride * index, value_type, &value)?;
654 }
655 Instruction::AtomicAddBuffer {
656 register,
657 slot,
658 offset,
659 stride,
660 count,
661 index,
662 value,
663 } => {
664 let index = match index {
665 Some(index) => read_buffer_array_index(registers, *index, *count)?,
666 None => 0,
667 };
668 let value = expect_u32(read_register(registers, *value)?)?;
669 let buffer = descriptors.buffer_mut(*slot)?;
670 let address = *offset + *stride * index;
671 let previous = expect_u32(buffer.read_value(address, &ValueType::U32)?)?;
672 buffer.write_value(address, &ValueType::U32, &Value::U32(previous.wrapping_add(value)))?;
673 registers[*register] = Some(Value::U32(previous));
674 }
675 Instruction::Call {
676 register,
677 function,
678 arguments,
679 } => {
680 let arguments = arguments
681 .iter()
682 .map(|argument| read_register(registers, *argument))
683 .collect::<Result<Vec<_>, _>>()?;
684 let nested_barrier_behavior = match barrier_behavior {
686 BarrierBehavior::Ignore => BarrierBehavior::Ignore,
687 BarrierBehavior::Suspend | BarrierBehavior::Reject => BarrierBehavior::Reject,
688 };
689 let value = self.execute_function(*function, &arguments, descriptors, state, nested_barrier_behavior)?;
690 if let Some(register) = register {
691 registers[*register] = value;
692 }
693 }
694 Instruction::Return { register } => {
695 return match register {
696 Some(register) => Ok(FrameOutcome::Complete(Some(read_register(registers, *register)?))),
697 None => Ok(FrameOutcome::Complete(None)),
698 };
699 }
700 }
701
702 frame.instruction_index += 1;
703 }
704
705 match &function.return_type {
706 Some(return_type) => Err(VmError::UnsupportedStatement {
707 message: format!(
708 "Function with return type `{}` ended without returning a value",
709 return_type.name()
710 ),
711 }),
712 None => Ok(FrameOutcome::Complete(None)),
713 }
714 }
715}