Skip to main content

tract_core/
floats.rs

1use crate::internal::translator::Translate;
2use crate::internal::*;
3use crate::ops::binary::TypedBinOp;
4use crate::ops::cast::{Cast, cast};
5use crate::ops::einsum::EinSum;
6use crate::ops::element_wise::ElementWiseOp;
7use crate::ops::konst::Const;
8use crate::ops::scan::Scan;
9use crate::ops::source::TypedSource;
10use crate::transform::ModelTransform;
11
12pub struct FloatPrecisionTranslator {
13    from_dt: DatumType,
14    to_dt: DatumType,
15    #[allow(clippy::type_complexity)]
16    node_predicate: Option<Box<dyn Fn(&TypedNode) -> bool>>,
17}
18
19impl FloatPrecisionTranslator {
20    pub fn new(from_dt: DatumType, to_dt: DatumType) -> Self {
21        Self { from_dt, to_dt, node_predicate: None }
22    }
23
24    pub fn with_filter(
25        from_dt: DatumType,
26        to_dt: DatumType,
27        node_predicate: impl Fn(&TypedNode) -> bool + 'static,
28    ) -> Self {
29        Self { from_dt, to_dt, node_predicate: Some(Box::new(node_predicate)) }
30    }
31
32    fn should_translate_node(&self, node: &TypedNode) -> bool {
33        self.node_predicate.as_ref().map(|it| (it)(node)).unwrap_or(true)
34    }
35
36    /// Cast node inputs to the working float precision for the operator
37    /// Only input using float datumtype are impacted. This will add cast operations
38    /// in the model. The function return the new input outlet ids.
39    fn cast_inputs_if_required(
40        &self,
41        model: &mut TypedModel,
42        node: &TypedNode,
43        mapping: &HashMap<OutletId, OutletId>,
44        op_float_dt: DatumType,
45    ) -> TractResult<TVec<OutletId>> {
46        let original_op_float_dt =
47            if op_float_dt == self.from_dt { self.to_dt } else { self.from_dt };
48
49        let mut mapped_inputs = tvec![];
50        for (i_idx, i) in node.inputs.iter().enumerate() {
51            let fact = model.outlet_fact(mapping[i])?;
52            if fact.datum_type == original_op_float_dt && fact.is_plain() {
53                let casted_mapped_input = model.wire_node(
54                    format!("{}.cast-{i_idx}", node.name),
55                    Cast { to: op_float_dt },
56                    &[mapping[i]],
57                )?[0];
58                mapped_inputs.push(casted_mapped_input);
59            } else {
60                mapped_inputs.push(mapping[i])
61            }
62        }
63        Ok(mapped_inputs)
64    }
65
66    /// Cast node output outlet ids to the destination float precision,
67    /// after insertion in the target mode. This preserves the model output float
68    /// precision.
69    fn cast_model_outputs_if_required(
70        &self,
71        source: &TypedModel,
72        node: &TypedNode,
73        target: &mut TypedModel,
74        target_node_outlet_ids: TVec<OutletId>,
75    ) -> TractResult<TVec<OutletId>> {
76        let mut outputs = tvec![];
77        for (o_idx, o) in target_node_outlet_ids.into_iter().enumerate() {
78            // Add Cast op for model output
79            let is_source_output = source.outputs.contains(&OutletId::new(node.id, o_idx));
80            let fact = target.outlet_fact(o)?;
81            if fact.datum_type == self.from_dt && fact.is_plain() && is_source_output {
82                let casted_output = target.wire_node(
83                    format!("{}.cast-out-{o_idx}", node.name),
84                    Cast { to: self.to_dt },
85                    &[o],
86                )?[0];
87                outputs.push(casted_output);
88            } else {
89                outputs.push(o)
90            }
91        }
92        Ok(outputs)
93    }
94}
95
96impl std::fmt::Debug for FloatPrecisionTranslator {
97    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
98        f.debug_struct("FloatPrecisionTranslator")
99            .field("from", &self.from_dt)
100            .field("to", &self.to_dt)
101            .finish()
102    }
103}
104
105impl ModelTransform for FloatPrecisionTranslator {
106    fn name(&self) -> StaticName {
107        format!("{:?}-to-{:?}", self.from_dt, self.to_dt).into()
108    }
109
110    fn transform(&self, model: &mut TypedModel) -> TractResult<()> {
111        let new = self.translate_model(model)?;
112        *model = new;
113        Ok(())
114    }
115}
116
117impl Translate<TypedFact, Box<dyn TypedOp>, TypedFact, Box<dyn TypedOp>>
118    for FloatPrecisionTranslator
119{
120    fn translate_node(
121        &self,
122        source: &TypedModel,
123        node: &TypedNode,
124        target: &mut TypedModel,
125        mapping: &HashMap<OutletId, OutletId>,
126    ) -> TractResult<TVec<OutletId>> {
127        let is_source = node.op_as::<TypedSource>().is_some();
128        if !self.should_translate_node(node) && !is_source {
129            let new_op = node.op.clone();
130
131            let casted_inputs =
132                self.cast_inputs_if_required(target, node, mapping, self.from_dt)?;
133            let target_node_outlet_ids = target.wire_node(&node.name, new_op, &casted_inputs)?;
134
135            self.cast_model_outputs_if_required(source, node, target, target_node_outlet_ids)
136        } else {
137            let casted_inputs = self.cast_inputs_if_required(target, node, mapping, self.to_dt)?;
138
139            let new_op = if let Some(source_op) = node.op_as::<TypedSource>() {
140                let mut fact = source_op.fact.clone();
141                if fact.datum_type == self.from_dt {
142                    fact.datum_type = self.to_dt;
143                }
144                Box::new(TypedSource::new(fact))
145            } else if let Some(konst) = node.op_as::<Const>() {
146                if konst.val().datum_type() == self.from_dt && konst.val().is_plain() {
147                    let wire = target.add_const(
148                        format!("{}.{:?}", node.name, self.from_dt),
149                        konst.val().clone(),
150                    )?;
151                    return target.wire_node(&node.name, cast(self.to_dt), &[wire]);
152                } else {
153                    node.op.clone()
154                }
155            } else if let Some(cast_op) = node.op_as::<Cast>() {
156                if cast_op.to == self.from_dt {
157                    Box::new(Cast { to: self.to_dt })
158                } else {
159                    node.op.clone()
160                }
161            } else if let Some(ew) = node.op_as::<ElementWiseOp>() {
162                if ew.1 == Some(self.from_dt) {
163                    Box::new(ElementWiseOp(ew.0.clone(), Some(self.to_dt)))
164                } else {
165                    node.op.clone()
166                }
167            } else if let Some(bin) = node.op_as::<TypedBinOp>() {
168                if bin.1 == Some(self.from_dt) {
169                    Box::new(TypedBinOp(bin.0.clone(), Some(self.to_dt)))
170                } else {
171                    node.op.clone()
172                }
173            } else if let Some(op) = node.op_as::<Scan>() {
174                let body = FloatPrecisionTranslator::new(self.from_dt, self.to_dt)
175                    .translate_model(&op.body)?;
176                Box::new(Scan { body, ..op.clone() })
177            } else if let Some(op) = node.op_as::<EinSum>() {
178                let operating_dt =
179                    if op.operating_dt == self.from_dt { self.to_dt } else { op.operating_dt };
180                Box::new(EinSum { operating_dt, ..op.clone() })
181            } else {
182                node.op.clone()
183            };
184            target.wire_node(&node.name, new_op, &casted_inputs)
185        }
186    }
187}
188
189#[cfg(test)]
190mod test {
191    use super::*;
192    use crate::ops::math;
193    use tract_data::prelude::f16;
194
195    fn build_f32_model() -> TractResult<TypedModel> {
196        // F32 model definition
197        let mut model = TypedModel::default();
198        let a = model.add_source("source", f32::fact([1])).unwrap();
199        let multiplier = model.add_const("multiplier", tensor1(&[1.0f32]))?;
200        let neg_infinity = model.add_const("neg_infinity", tensor1(&[f32::NEG_INFINITY]))?;
201        let pow_factor = model.add_const("pow_factor", tensor1(&[10.0f32]))?;
202        let add = model.wire_node("layer.0/add", math::add(), &[a, a]).unwrap()[0];
203        let mul = model.wire_node("layer.0/mul", math::mul(), &[add, multiplier]).unwrap()[0];
204        let pow = model.wire_node("layer.1/pow", math::pow(), &[mul, pow_factor]).unwrap()[0];
205        let _output = model
206            .wire_node("layer.1/add_neg_infinity", math::add(), &[pow, neg_infinity])
207            .unwrap()[0];
208        model.auto_outputs()?;
209        Ok(model)
210    }
211
212    #[test]
213    fn test_high_level_f16_transform_with_filter() -> TractResult<()> {
214        // F32 model definition
215        let model = build_f32_model()?;
216
217        // Execution in F32
218        let runnable_model = model.clone().into_runnable()?;
219        assert_eq!(
220            runnable_model.run(tvec![tensor1(&[5.0f32]).into()])?[0],
221            tensor1(&[f32::NEG_INFINITY]).into()
222        );
223
224        // Execution in F16 with returns NaN
225        let runnable_model = &crate::transform::get_transform("f32_to_f16")?
226            .unwrap()
227            .transform_into(model.clone())?
228            .into_runnable()?;
229        assert!(
230            runnable_model.run(tvec![tensor1(&[f16::from_f32(5.0)]).into()])?[0]
231                .try_as_plain()?
232                .to_scalar::<f16>()?
233                .is_nan()
234        );
235
236        // Execution in F16 with filter that returns the good output.
237        let runnable_model = &crate::transform::build_float_translator(
238            f32::datum_type(),
239            f16::datum_type(),
240            crate::transform::NodeFilter {
241                exclude: Some(vec!["layer.1".into()]),
242                ..Default::default()
243            },
244        )
245        .transform_into(model.clone())?
246        .into_runnable()?;
247        assert_eq!(
248            runnable_model.run(tvec![tensor1(&[f16::from_f32(5.0)]).into()])?[0],
249            tensor1(&[f16::NEG_INFINITY]).into()
250        );
251
252        // Execution in F16 with returns NaN despite the filter.
253        let runnable_model = &crate::transform::build_float_translator(
254            f32::datum_type(),
255            f16::datum_type(),
256            crate::transform::NodeFilter {
257                exclude: Some(vec!["layer.0".into()]),
258                ..Default::default()
259            },
260        )
261        .transform_into(model)?
262        .into_runnable()?;
263        assert!(
264            runnable_model.run(tvec![tensor1(&[f16::from_f32(5.0)]).into()])?[0]
265                .try_as_plain()?
266                .to_scalar::<f16>()?
267                .is_nan()
268        );
269
270        Ok(())
271    }
272
273    #[test]
274    fn test_f16_transform_with_filter() -> TractResult<()> {
275        // F32 model definition
276        let model = build_f32_model()?;
277
278        // Execution in F32
279        let runnable_model = model.clone().into_runnable()?;
280        assert_eq!(
281            runnable_model.run(tvec![tensor1(&[5.0f32]).into()])?[0],
282            tensor1(&[f32::NEG_INFINITY]).into()
283        );
284
285        // Execution in F16 with returns NaN
286        let mut model_f16 = model.clone();
287        model_f16
288            .transform(&FloatPrecisionTranslator::new(f32::datum_type(), f16::datum_type()))?;
289        let runnable_model_f16 = model_f16.clone().into_runnable()?;
290        assert!(
291            runnable_model_f16.run(tvec![tensor1(&[f16::from_f32(5.0)]).into()])?[0]
292                .try_as_plain()?
293                .to_scalar::<f16>()?
294                .is_nan()
295        );
296
297        // Execution in F16 with filter that returns the good output.
298        let mut model_f16_with_filter = model.clone();
299        model_f16_with_filter.transform(&FloatPrecisionTranslator::with_filter(
300            f32::datum_type(),
301            f16::datum_type(),
302            |node| !node.name.contains("layer.1"),
303        ))?;
304        let runnable_model_f16 = model_f16_with_filter.clone().into_runnable()?;
305        assert_eq!(
306            runnable_model_f16.run(tvec![tensor1(&[f16::from_f32(5.0)]).into()])?[0],
307            tensor1(&[f16::NEG_INFINITY]).into()
308        );
309        let mut model_f16_with_filter = model.clone();
310        model_f16_with_filter.transform(&FloatPrecisionTranslator::with_filter(
311            f32::datum_type(),
312            f16::datum_type(),
313            |node| !node.name.contains("layer.0"),
314        ))?;
315        let runnable_model_f16 = model_f16_with_filter.clone().into_runnable()?;
316        assert!(
317            runnable_model_f16.run(tvec![tensor1(&[f16::from_f32(5.0)]).into()])?[0]
318                .try_as_plain()?
319                .to_scalar::<f16>()?
320                .is_nan()
321        );
322        Ok(())
323    }
324}