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 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 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 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 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 let model = build_f32_model()?;
216
217 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 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 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 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 let model = build_f32_model()?;
277
278 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 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 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}