gwr_models/processing_element/operators/
custom.rs1use std::rc::Rc;
6
7use gwr_engine::types::{SimError, SimResult};
8use serde::{Deserialize, Serialize};
9
10use super::{Operator, Tensor, TensorPartition, TensorView};
11use crate::processing_element::{ComputeCapabilities, MachineOp, MachineOpCounts};
12
13#[derive(Clone, Debug, Deserialize, Serialize)]
14#[serde(deny_unknown_fields)]
15pub struct OperatorCustom {
16 #[serde(default, skip_serializing_if = "Option::is_none")]
17 pub name: Option<String>,
18 pub machine_ops: MachineOpCounts,
19}
20
21impl Operator for OperatorCustom {
22 fn validate_tensors(
23 &self,
24 _inputs: &[Option<Tensor>],
25 _outputs: &[Option<Tensor>],
26 ) -> SimResult {
27 Ok(())
28 }
29
30 fn compute_delay_ticks(
31 &self,
32 compute_capabilities: &Rc<ComputeCapabilities>,
33 _inputs: &[Option<TensorView>],
34 _outputs: &[Option<TensorView>],
35 ) -> Result<usize, SimError> {
36 Ok(
37 compute_capabilities.cycles_for_ops(self.machine_ops.adds, MachineOp::Add)?
38 + compute_capabilities.cycles_for_ops(self.machine_ops.muls, MachineOp::Mul)?
39 + compute_capabilities
40 .cycles_for_ops(self.machine_ops.compares, MachineOp::Compare)?,
41 )
42 }
43
44 fn compute_machine_ops(
45 &self,
46 _inputs: &[Option<TensorView>],
47 _outputs: &[Option<TensorView>],
48 ) -> Result<MachineOpCounts, SimError> {
49 Ok(self.machine_ops)
50 }
51
52 fn partition_views(
53 &self,
54 input_views: &[Option<TensorView>],
55 output_views: &[Option<TensorView>],
56 _num_partitions: usize,
57 ) -> Result<Vec<TensorPartition>, SimError> {
58 Ok(vec![TensorPartition {
59 inputs: input_views.to_vec(),
60 outputs: output_views.to_vec(),
61 }])
62 }
63}