tract_tensorflow/ops/
logic.rs1use tract_hir::internal::*;
2use tract_hir::ops;
3use tract_hir::ops::binary::BinIntoHir;
4use tract_hir::ops::logic::{CompEq, CompGT, CompGTE, CompLT, CompLTE};
5
6use crate::model::ParsingContext;
7use crate::model::TfOpRegister;
8use crate::tfpb::tensorflow::NodeDef;
9use std::collections::HashSet;
10
11pub fn register_all_ops(reg: &mut TfOpRegister) {
12 reg.insert("Equal", |_, _| Ok(CompEq.into_hir()));
13 reg.insert("Greater", |_, _| Ok(CompGT.into_hir()));
14 reg.insert("GreaterEqual", |_, _| Ok(CompGTE.into_hir()));
15 reg.insert("Less", |_, _| Ok(CompLT.into_hir()));
16 reg.insert("LessEqual", |_, _| Ok(CompLTE.into_hir()));
17 reg.insert("LogicalAnd", |_, _| Ok(ops::logic::And.into_hir()));
18 reg.insert("LogicalOr", |_, _| Ok(ops::logic::Or.into_hir()));
19 reg.insert("Merge", merge);
20 reg.insert("Switch", |_, _| Ok(Box::new(Switch)));
21}
22
23#[derive(Debug, Clone, new, Hash, PartialEq, Eq)]
24pub struct Switch;
25
26impl Op for Switch {
27 fn name(&self) -> StaticName {
28 "Switch".into()
29 }
30
31 not_a_typed_op!();
32}
33
34impl EvalOp for Switch {
35 op_out_of_plan!();
36
37 fn state(&self, _ctx: &EvalContext) -> TractResult<Option<Box<dyn OpState>>> {
38 Ok(None)
39 }
40}
41
42impl InferenceRulesOp for Switch {
43 fn rules<'r, 'p: 'r, 's: 'r>(
44 &'s self,
45 s: &mut Solver<'r>,
46 inputs: &'p [TensorProxy],
47 outputs: &'p [TensorProxy],
48 ) -> InferenceResult {
49 check_input_arity(inputs, 2)?;
50 check_output_arity(outputs, 2)?;
51 s.equals(&inputs[1].datum_type, DatumType::Bool)?;
52 s.equals(&inputs[1].shape, shapefactoid!())?;
53 for output in outputs {
54 s.equals(&inputs[0].datum_type, &output.datum_type)?;
55 s.equals(&inputs[0].shape, &output.shape)?;
56 }
57 Ok(())
58 }
59
60 fn incorporate(
61 &self,
62 model: &InferenceModel,
63 node: &InferenceNode,
64 ) -> TractResult<Option<InferenceModelPatch>> {
65 let pred = model.outlet_fact(node.inputs[1])?;
66 if let Some(pred) = pred.concretize() {
67 let pred = *pred.try_as_plain()?.to_scalar::<bool>()?;
68 let mut dead_to_visit = HashSet::new();
69 let mut dead_done = HashSet::new();
70 let mut patch = InferenceModelPatch::default();
71 dead_to_visit.insert(OutletId::new(node.id, !pred as usize));
72 while let Some(dead_outlet) = dead_to_visit.iter().cloned().next() {
73 dead_to_visit.remove(&dead_outlet);
74 dead_done.insert(dead_outlet);
75 for succ in model.outlet_successors(dead_outlet) {
76 if model.node(succ.node).op_is::<Merge>() {
77 let outlet = model.node(succ.node).inputs[(succ.slot == 0) as usize];
78 let tap = patch.tap_model(model, outlet)?;
79 patch.shunt_outside(model, succ.node.into(), tap)?;
80 } else {
81 for slot in 0..model.node(succ.node).outputs.len() {
82 let new = OutletId::new(succ.node, slot);
83 if !dead_done.contains(&new) {
84 dead_to_visit.insert(new);
85 }
86 }
87 }
88 }
89 }
90 let tap = patch.tap_model(model, node.inputs[0])?;
91 patch.shunt_outside(model, OutletId::new(node.id, 0), tap)?;
92 patch.shunt_outside(model, OutletId::new(node.id, 1), tap)?;
93 return Ok(Some(patch));
94 }
95 Ok(None)
96 }
97
98 fn nboutputs(&self) -> TractResult<usize> {
99 Ok(2)
100 }
101
102 as_op!();
103}
104
105fn merge(_ctx: &ParsingContext, pb: &NodeDef) -> TractResult<Box<dyn InferenceOp>> {
106 let inputs = pb.get_attr_int::<i32>("N")?;
107 Ok(Box::new(Merge::new(inputs as usize)))
108}
109
110#[derive(Debug, Clone, new, Hash, PartialEq, Eq)]
111pub struct Merge {
112 n: usize,
113}
114
115impl Op for Merge {
116 fn name(&self) -> StaticName {
117 "Merge".into()
118 }
119
120 op_as_typed_op!();
121}
122
123impl EvalOp for Merge {
124 op_out_of_plan!();
125
126 fn state(&self, _ctx: &EvalContext) -> TractResult<Option<Box<dyn OpState>>> {
127 Ok(None)
128 }
129}
130
131impl InferenceRulesOp for Merge {
132 fn rules<'r, 'p: 'r, 's: 'r>(
133 &'s self,
134 s: &mut Solver<'r>,
135 inputs: &'p [TensorProxy],
136 outputs: &'p [TensorProxy],
137 ) -> InferenceResult {
138 check_input_arity(inputs, self.n)?;
139 check_output_arity(outputs, 1)?;
140 for i in 1..self.n {
141 s.equals(&inputs[0].datum_type, &inputs[i].datum_type)?;
142 s.equals(&inputs[0].shape, &inputs[i].shape)?;
143 }
144 s.equals(&inputs[0].datum_type, &outputs[0].datum_type)?;
145 s.equals(&inputs[0].shape, &outputs[0].shape)?;
146 Ok(())
147 }
148
149 as_op!();
150 to_typed!();
151}
152
153impl TypedOp for Merge {
154 as_op!();
155
156 fn output_facts(&self, inputs: &[&TypedFact]) -> TractResult<TVec<TypedFact>> {
157 Ok(tvec!(f32::fact(inputs[0].shape.iter()), i32::fact([0; 0])))
158 }
159}