gwr_models/processing_element/
task.rs1use std::rc::Rc;
4
5use gwr_engine::types::SimError;
6use serde::ser::SerializeMap;
7use serde::{Deserialize, Serialize, Serializer};
8
9use crate::processing_element::operators::add::OperatorAdd;
10use crate::processing_element::operators::custom::OperatorCustom;
11use crate::processing_element::operators::gemm::OperatorGemm;
12use crate::processing_element::operators::maxpool::OperatorMaxPool;
13use crate::processing_element::operators::{Operator, TensorPartition, TensorView};
14use crate::processing_element::{ComputeCapabilities, MachineOpCounts};
15
16#[derive(Debug, Clone)]
17pub struct ComputeTaskConfig {
18 pub id: String,
20 pub op: ComputeOp,
21 pub inputs: Vec<Option<TensorView>>,
22 pub outputs: Vec<Option<TensorView>>,
23}
24
25impl ComputeTaskConfig {
26 #[must_use]
27 pub fn activity_name(&self) -> &str {
28 match &self.op {
29 ComputeOp::Custom(operator) => operator.name.as_deref().unwrap_or(&self.id),
30 _ => &self.id,
31 }
32 }
33}
34
35#[derive(Clone, Debug, Deserialize)]
36#[serde(rename_all = "lowercase")]
37pub enum ComputeOp {
38 Add,
39 Gemm,
40 MaxPool(OperatorMaxPool),
41 Custom(OperatorCustom),
42}
43
44impl Serialize for ComputeOp {
45 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
46 where
47 S: Serializer,
48 {
49 match self {
50 Self::Add => serializer.serialize_str("add"),
51 Self::Gemm => serializer.serialize_str("gemm"),
52 Self::MaxPool(operator) => {
53 let mut map = serializer.serialize_map(Some(1))?;
54 map.serialize_entry("maxpool", operator)?;
55 map.end()
56 }
57 Self::Custom(operator) => {
58 let mut map = serializer.serialize_map(Some(1))?;
59 map.serialize_entry("custom", operator)?;
60 map.end()
61 }
62 }
63 }
64}
65
66impl ComputeOp {
67 #[must_use]
68 pub fn trace_name(&self) -> &str {
69 match self {
70 ComputeOp::Add => "add",
71 ComputeOp::Gemm => "gemm",
72 ComputeOp::MaxPool(_) => "maxpool",
73 ComputeOp::Custom(operator) => operator.name.as_deref().unwrap_or("custom"),
74 }
75 }
76
77 pub fn compute_delay_ticks(
78 &self,
79 compute_capabilities: &Rc<ComputeCapabilities>,
80 input_views: &[Option<TensorView>],
81 output_views: &[Option<TensorView>],
82 ) -> Result<usize, SimError> {
83 match self {
84 ComputeOp::Add => {
85 OperatorAdd {}.compute_delay_ticks(compute_capabilities, input_views, output_views)
86 }
87 ComputeOp::Gemm => {
88 OperatorGemm {}.compute_delay_ticks(compute_capabilities, input_views, output_views)
89 }
90 ComputeOp::MaxPool(operator) => {
91 operator.compute_delay_ticks(compute_capabilities, input_views, output_views)
92 }
93 ComputeOp::Custom(operator) => {
94 operator.compute_delay_ticks(compute_capabilities, input_views, output_views)
95 }
96 }
97 }
98
99 pub fn compute_flops(
100 &self,
101 input_views: &[Option<TensorView>],
102 output_views: &[Option<TensorView>],
103 ) -> Result<usize, SimError> {
104 match self {
105 ComputeOp::Add => OperatorAdd {}.compute_flops(input_views, output_views),
106 ComputeOp::Gemm => OperatorGemm {}.compute_flops(input_views, output_views),
107 ComputeOp::MaxPool(operator) => operator.compute_flops(input_views, output_views),
108 ComputeOp::Custom(operator) => operator.compute_flops(input_views, output_views),
109 }
110 }
111
112 pub fn compute_machine_ops(
113 &self,
114 input_views: &[Option<TensorView>],
115 output_views: &[Option<TensorView>],
116 ) -> Result<MachineOpCounts, SimError> {
117 match self {
118 ComputeOp::Add => OperatorAdd {}.compute_machine_ops(input_views, output_views),
119 ComputeOp::Gemm => OperatorGemm {}.compute_machine_ops(input_views, output_views),
120 ComputeOp::MaxPool(operator) => operator.compute_machine_ops(input_views, output_views),
121 ComputeOp::Custom(operator) => operator.compute_machine_ops(input_views, output_views),
122 }
123 }
124
125 pub fn create_partitions(
126 &self,
127 input_views: &[Option<TensorView>],
128 output_views: &[Option<TensorView>],
129 num_partitions: usize,
130 ) -> Result<Vec<TensorPartition>, SimError> {
131 match self {
132 ComputeOp::Add => {
133 OperatorAdd {}.partition_views(input_views, output_views, num_partitions)
134 }
135 ComputeOp::Gemm => {
136 OperatorGemm {}.partition_views(input_views, output_views, num_partitions)
137 }
138 ComputeOp::MaxPool(operator) => {
139 operator.partition_views(input_views, output_views, num_partitions)
140 }
141 ComputeOp::Custom(operator) => {
142 operator.partition_views(input_views, output_views, num_partitions)
143 }
144 }
145 }
146}
147
148#[derive(Debug, Clone, Copy)]
149pub enum SyncRegion {
150 Local,
151 Global,
152}
153
154#[derive(Debug, Clone)]
155pub enum Task {
156 ComputeTask { config: ComputeTaskConfig },
157 SyncTask { region: SyncRegion },
158}