1use std::borrow::Cow;
2
3use crate::internal::*;
4use crate::ops::matmul::de_block_quant::BlockQuantTransform;
5use std::fmt::Debug;
6
7use tract_data::TractResult;
8
9use crate::floats::FloatPrecisionTranslator;
10use crate::ops::nn::{Softmax, SoftmaxExp, SoftmaxKind, TypedModel};
11
12#[macro_export]
13macro_rules! rule_if {
14 ($cond:expr) => {
15 if !$cond {
16 return Ok(None);
17 }
18 };
19}
20
21#[macro_export]
22macro_rules! rule_if_let {
23 ($pat:pat = $expr:expr) => {
24 let $pat = $expr else {
25 return Ok(None);
26 };
27 };
28}
29
30#[macro_export]
31macro_rules! rule_if_some {
32 ($pat:pat = $expr:expr) => {
33 let Some($pat) = $expr else {
34 return Ok(None);
35 };
36 };
37}
38
39#[derive(Debug, Clone, Default)]
44pub struct NodeFilter {
45 pub include: Option<Vec<String>>,
46 pub exclude: Option<Vec<String>>,
47}
48
49impl NodeFilter {
50 pub fn matches(&self, name: &str) -> bool {
52 let dominated = match &self.include {
53 Some(patterns) => patterns.iter().any(|p| name.contains(p)),
54 None => true,
55 };
56 if !dominated {
57 return false;
58 }
59 match &self.exclude {
60 Some(patterns) => !patterns.iter().any(|p| name.contains(p)),
61 None => true,
62 }
63 }
64
65 pub fn is_pass_through(&self) -> bool {
67 self.include.is_none() && self.exclude.is_none()
68 }
69}
70
71pub fn parse_legacy_filter(filter: Option<&str>) -> TractResult<NodeFilter> {
73 let Some(filter) = filter.filter(|f| !f.is_empty()) else {
74 return Ok(NodeFilter::default());
75 };
76 if let Some(patterns) = filter.strip_prefix("!=") {
77 let patterns = patterns.split(',').map(|it| it.trim().to_string()).collect();
78 Ok(NodeFilter { exclude: Some(patterns), ..Default::default() })
79 } else if let Some(patterns) = filter.strip_prefix("==") {
80 let patterns = patterns.split(',').map(|it| it.trim().to_string()).collect();
81 Ok(NodeFilter { include: Some(patterns), ..Default::default() })
82 } else {
83 Ok(NodeFilter::default())
84 }
85}
86
87pub fn build_float_translator(
90 from_dt: DatumType,
91 to_dt: DatumType,
92 filter: NodeFilter,
93) -> Box<dyn ModelTransform> {
94 if filter.is_pass_through() {
95 return Box::new(FloatPrecisionTranslator::new(from_dt, to_dt));
96 }
97 Box::new(FloatPrecisionTranslator::with_filter(from_dt, to_dt, move |node| {
98 filter.matches(&node.name)
99 }))
100}
101
102pub trait ModelTransform: Debug {
103 fn name(&self) -> StaticName;
104 fn transform(&self, model: &mut TypedModel) -> TractResult<()>;
105 fn transform_into(&self, mut model: TypedModel) -> TractResult<TypedModel> {
106 self.transform(&mut model)?;
107 Ok(model)
108 }
109}
110
111#[derive(Debug)]
112struct SoftmaxFastCompact;
113
114impl ModelTransform for SoftmaxFastCompact {
115 fn name(&self) -> StaticName {
116 "softmax_fast_compact".into()
117 }
118
119 fn transform(&self, model: &mut TypedModel) -> TractResult<()> {
120 for node in &mut model.nodes {
121 if let Some(softmax) = node.op_as_mut::<Softmax>()
122 && let SoftmaxKind::Softmax(kind) = &mut softmax.kind
123 {
124 *kind = SoftmaxExp::FastCompact
125 }
126 }
127 Ok(())
128 }
129}
130
131#[derive(Debug, Default, serde::Deserialize)]
133pub struct FloatTranslatorConfig {
134 #[serde(default)]
136 pub filter: Option<String>,
137 #[serde(default)]
139 pub include: Option<Vec<String>>,
140 #[serde(default)]
142 pub exclude: Option<Vec<String>>,
143}
144
145impl FloatTranslatorConfig {
146 pub fn into_node_filter(self) -> TractResult<NodeFilter> {
147 if self.include.is_some() || self.exclude.is_some() {
148 Ok(NodeFilter { include: self.include, exclude: self.exclude })
149 } else {
150 parse_legacy_filter(self.filter.as_deref())
151 }
152 }
153}
154
155#[derive(Debug, serde::Deserialize)]
157pub struct FloatPrecisionConfig {
158 pub from: String,
159 pub to: String,
160 #[serde(default)]
162 pub include: Option<Vec<String>>,
163 #[serde(default)]
165 pub exclude: Option<Vec<String>>,
166}
167
168pub struct ModelTransformFactory {
169 pub name: &'static str,
170 pub build_default: fn() -> TractResult<Box<dyn ModelTransform>>,
172 pub build: fn(&mut dyn erased_serde::Deserializer) -> TractResult<Box<dyn ModelTransform>>,
174}
175
176inventory::collect!(ModelTransformFactory);
177
178#[macro_export]
179macro_rules! register_simple_model_transform {
180 ($name: expr, $type: expr) => {
181 $crate::internal::inventory::submit! {
182 $crate::transform::ModelTransformFactory {
183 name: $name,
184 build_default: || Ok(Box::new($type)),
185 build: |_de| Ok(Box::new($type)),
186 }
187 }
188 };
189}
190
191#[macro_export]
192macro_rules! register_model_transform {
193 ($name:expr, $config:ty, $builder:expr) => {
194 $crate::internal::inventory::submit! {
195 $crate::transform::ModelTransformFactory {
196 name: $name,
197 build_default: || {
198 let config = <$config>::default();
199 let builder: fn($config) -> $crate::prelude::TractResult<Box<dyn $crate::transform::ModelTransform>> = $builder;
200 builder(config)
201 },
202 build: |de: &mut dyn erased_serde::Deserializer| {
203 let config: $config = erased_serde::deserialize(de)
204 .map_err(|e| $crate::internal::anyhow!("deserializing transform config: {e}"))?;
205 let builder: fn($config) -> $crate::prelude::TractResult<Box<dyn $crate::transform::ModelTransform>> = $builder;
206 builder(config)
207 },
208 }
209 }
210 };
211}
212
213pub fn split_spec(spec: &str) -> (Cow<'_, str>, &str) {
215 if let Some(pos) = spec.find('(') {
216 (Cow::Borrowed(&spec[..pos]), &spec[pos..])
217 } else if spec.contains('-') {
218 (Cow::Owned(spec.replace('-', "_")), "")
220 } else {
221 (Cow::Borrowed(spec), "")
222 }
223}
224
225pub fn get_transform(name: &str) -> TractResult<Option<Box<dyn ModelTransform>>> {
227 let (name, _) = split_spec(name);
228 for factory in inventory::iter::<ModelTransformFactory>() {
229 if factory.name == &*name {
230 return Ok(Some((factory.build_default)()?));
231 }
232 }
233 Ok(None)
234}
235
236pub fn get_transform_with_params(
238 name: &str,
239 de: &mut dyn erased_serde::Deserializer,
240) -> TractResult<Option<Box<dyn ModelTransform>>> {
241 for factory in inventory::iter::<ModelTransformFactory>() {
242 if factory.name == name {
243 return Ok(Some((factory.build)(de)?));
244 }
245 }
246 Ok(None)
247}
248
249#[derive(Debug, serde::Deserialize)]
252#[serde(untagged)]
253pub enum SymbolValueSpec {
254 Int(i64),
255 Expr(String),
256}
257
258#[derive(Debug, Default, serde::Deserialize)]
259pub struct SetSymbolsConfig {
260 pub values: std::collections::HashMap<String, SymbolValueSpec>,
261}
262
263#[derive(Debug)]
264struct SetSymbolsTransform(SetSymbolsConfig);
265
266impl ModelTransform for SetSymbolsTransform {
267 fn name(&self) -> StaticName {
268 "set_symbols".into()
269 }
270
271 fn transform(&self, model: &mut TypedModel) -> TractResult<()> {
272 let mut subs = std::collections::HashMap::new();
273 for (k, spec) in &self.0.values {
274 let sym = model.symbols.sym(k);
275 let dim = match spec {
276 SymbolValueSpec::Int(v) => TDim::Val(*v),
277 SymbolValueSpec::Expr(s) => model
278 .symbols
279 .parse_tdim(s)
280 .with_context(|| format!("Parsing TDim expression {s:?} for symbol {k}"))?,
281 };
282 subs.insert(sym, dim);
283 }
284 *model = model.set_symbols(&subs)?;
285 Ok(())
286 }
287}
288
289register_model_transform!("set_symbols", SetSymbolsConfig, |config| Ok(Box::new(
290 SetSymbolsTransform(config)
291)));
292
293#[derive(Debug)]
305struct ForceScanExternalState;
306
307impl ModelTransform for ForceScanExternalState {
308 fn name(&self) -> StaticName {
309 "force_scan_external_state".into()
310 }
311
312 fn transform(&self, model: &mut TypedModel) -> TractResult<()> {
313 use crate::ops::scan::Scan;
314 for node in &mut model.nodes {
315 if let Some(scan) = node.op_as_mut::<Scan>() {
316 scan.external_state = true;
317 }
318 }
319 Ok(())
320 }
321}
322
323register_simple_model_transform!("force_scan_external_state", ForceScanExternalState);
324
325#[derive(Debug)]
351struct ExportScanState;
352
353impl ModelTransform for ExportScanState {
354 fn name(&self) -> StaticName {
355 "export_scan_state".into()
356 }
357
358 fn transform(&self, model: &mut TypedModel) -> TractResult<()> {
359 use crate::ops::konst::Const;
360 use crate::ops::scan::{InputMapping, Scan};
361
362 let scans: Vec<usize> =
363 model.nodes().iter().filter(|n| n.op_is::<Scan>()).map(|n| n.id).collect();
364
365 for id in scans {
366 let eligible = {
367 let inputs = model.node_input_facts(id)?;
368 let scan = model.node(id).op_as::<Scan>().unwrap();
369 scan.skip == 0 && scan.iteration_count(&inputs).map(|i| i.is_one()).unwrap_or(false)
370 };
371 if !eligible {
372 continue;
373 }
374
375 let scan = model.node(id).op_as::<Scan>().unwrap();
376 let Some(state_output) =
377 scan.output_mapping.iter().find(|om| om.state).and_then(|om| om.scan.map(|s| s.0))
378 else {
379 continue;
380 };
381 let seeded: Vec<usize> = scan
382 .input_mapping
383 .iter()
384 .enumerate()
385 .filter(|(_, im)| matches!(im, InputMapping::State))
386 .map(|(slot, _)| slot)
387 .filter(|&slot| model.node(model.node(id).inputs[slot].node).op_is::<Const>())
388 .collect();
389 if seeded.is_empty() {
390 continue;
391 }
392
393 for slot in seeded {
394 let fact = model.outlet_fact(model.node(id).inputs[slot])?.clone().without_value();
395 let name = format!("{}.state_in", model.node(id).name);
396 let source = model.add_source(name, fact)?;
397 model.add_edge(source, InletId::new(id, slot))?;
398 }
399 model.outputs.push(OutletId::new(id, state_output));
400 model.node_mut(id).op_as_mut::<Scan>().unwrap().external_state = true;
401 }
402 Ok(())
403 }
404}
405
406register_simple_model_transform!("export_scan_state", ExportScanState);
407
408register_simple_model_transform!("softmax_fast_compact", SoftmaxFastCompact);
409register_simple_model_transform!("block_quant", BlockQuantTransform);
410
411#[derive(Debug, serde::Deserialize, Default)]
412pub struct SelectOutputsConfig {
413 pub outputs: Vec<String>,
414}
415
416#[derive(Debug)]
417struct SelectOutputsTransform(SelectOutputsConfig);
418
419impl ModelTransform for SelectOutputsTransform {
420 fn name(&self) -> StaticName {
421 "select_outputs".into()
422 }
423
424 fn transform(&self, model: &mut TypedModel) -> TractResult<()> {
425 model.select_outputs_by_name(self.0.outputs.iter())
426 }
427}
428
429register_model_transform!("select_outputs", SelectOutputsConfig, |config| Ok(Box::new(
430 SelectOutputsTransform(config)
431)));
432
433#[derive(Debug, serde::Deserialize, Default)]
434pub struct SelectInputsConfig {
435 pub inputs: Vec<String>,
436}
437
438#[derive(Debug)]
439struct SelectInputsTransform(SelectInputsConfig);
440
441impl ModelTransform for SelectInputsTransform {
442 fn name(&self) -> StaticName {
443 "select_inputs".into()
444 }
445
446 fn transform(&self, model: &mut TypedModel) -> TractResult<()> {
447 model.select_inputs_by_name(self.0.inputs.iter())
448 }
449}
450
451register_model_transform!("select_inputs", SelectInputsConfig, |config| Ok(Box::new(
452 SelectInputsTransform(config)
453)));
454
455inventory::submit! {
456 ModelTransformFactory {
457 name: "f32_to_f16",
458 build_default: || Ok(build_float_translator(DatumType::F32, DatumType::F16, NodeFilter::default())),
459 build: |de| {
460 let config: FloatTranslatorConfig = erased_serde::deserialize(de)
461 .map_err(|e| anyhow::anyhow!("deserializing f32_to_f16 config: {e}"))?;
462 Ok(build_float_translator(DatumType::F32, DatumType::F16, config.into_node_filter()?))
463 },
464 }
465}
466
467inventory::submit! {
468 ModelTransformFactory {
469 name: "f16_to_f32",
470 build_default: || Ok(build_float_translator(DatumType::F16, DatumType::F32, NodeFilter::default())),
471 build: |de| {
472 let config: FloatTranslatorConfig = erased_serde::deserialize(de)
473 .map_err(|e| anyhow::anyhow!("deserializing f16_to_f32 config: {e}"))?;
474 Ok(build_float_translator(DatumType::F16, DatumType::F32, config.into_node_filter()?))
475 },
476 }
477}
478
479inventory::submit! {
480 ModelTransformFactory {
481 name: "float_precision",
482 build_default: || {
483 anyhow::bail!("float_precision transform requires 'from' and 'to' parameters")
484 },
485 build: |de| {
486 let config: FloatPrecisionConfig = erased_serde::deserialize(de)
487 .map_err(|e| anyhow::anyhow!("deserializing float_precision config: {e}"))?;
488 let from_dt: DatumType = config.from.parse()
489 .map_err(|e| anyhow::anyhow!("parsing 'from' datum type: {e}"))?;
490 let to_dt: DatumType = config.to.parse()
491 .map_err(|e| anyhow::anyhow!("parsing 'to' datum type: {e}"))?;
492 let filter = NodeFilter { include: config.include, exclude: config.exclude };
493 Ok(build_float_translator(from_dt, to_dt, filter))
494 },
495 }
496}