Skip to main content

gwr_models/processing_element/
task.rs

1// Copyright (c) 2026 Graphcore Ltd. All rights reserved.
2
3use 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    /// Only needed as a debug aid
19    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}