tract_transformers/
lib.rs1pub mod ops;
2mod rewriter;
3use std::collections::HashSet;
4
5use rewriter::*;
6use tract_nnef::internal::*;
7
8register_simple_model_transform!("detect_apply_rope", ApplyRopeTransform);
9register_simple_model_transform!("detect_diag_gather", DetectDiagGatherTransform);
10register_simple_model_transform!("detect_scaled_masked_softmax", ScaledMaskedSoftmaxTransform);
11register_simple_model_transform!("detect_kv_cache", KeyValueCacheTransform);
12register_simple_model_transform!(
13 "detect_sdpa_kv_cache_broadcast",
14 SdpaFuseKvCacheBroadcastTransform
15);
16register_simple_model_transform!("unfold_kv_cache", UnfoldKeyValueCacheTransform);
17register_simple_model_transform!(
18 "fuse_inplace_kv_sdpa",
19 ops::inplace_kv_cache::InPlaceKvSdpaTransform
20);
21register_simple_model_transform!("transformers_detect_all", TransformersTransform);
22register_simple_model_transform!(
23 "quantize_kv_storage",
24 ops::quant_dyn_kv_cache::QuantizeKvStorageTransform { bits: 8 }
25);
26register_simple_model_transform!(
27 "quantize_kv_storage_int4",
28 ops::quant_dyn_kv_cache::QuantizeKvStorageTransform { bits: 4 }
29);
30
31pub fn register(registry: &mut Registry) {
32 ops::causal_conv1d_update::register(registry);
33 ops::apply_rope::register(registry);
34 ops::scaled_masked_softmax::register(registry);
35 ops::sdpa::register(registry);
36 ops::dyn_kv_cache::register(registry);
37 ops::window_kv_cache::register(registry);
38 ops::kv_quant::register(registry);
39 ops::quant_dyn_kv_cache::register(registry);
40 ops::gdn_recurrent::register(registry);
41}
42
43pub trait WithTractTransformers {
44 fn enable_tract_transformers(&mut self);
45 fn with_tract_transformers(self) -> Self;
46}
47
48impl WithTractTransformers for tract_nnef::framework::Nnef {
49 fn enable_tract_transformers(&mut self) {
50 self.registries.push(tract_transformers_registry());
51 }
52
53 fn with_tract_transformers(mut self) -> Self {
54 self.enable_tract_transformers();
55 self
56 }
57}
58
59pub fn tract_transformers_registry() -> Registry {
60 let mut reg = Registry::new("tract_transformers")
61 .with_doc("Extension `tract_transformers` extends NNEF with operators")
62 .with_doc("for transformer networks.")
63 .with_doc("")
64 .with_doc("Add `extension tract_transformers` to `graph.nnef`");
65
66 register(&mut reg);
67 reg
68}
69
70pub fn figure_out_causal_llm_b_s_p(
71 model: &TypedModel,
72) -> TractResult<(Option<Symbol>, Option<Symbol>, Option<Symbol>)> {
73 let token_input = model
77 .inputs
78 .iter()
79 .position(|i| model.outlet_fact(*i).unwrap().datum_type.is_integer())
80 .context("No token input found")?;
81 let tokens_symbols = model.input_fact(token_input)?.shape.volume().symbols();
82 let kv_symbols = if let Some(kv_input) =
83 model.inputs.iter().position(|i| model.outlet_fact(*i).unwrap().datum_type.is_float())
84 {
85 model.input_fact(kv_input)?.shape.volume().symbols()
86 } else {
87 let mut symbols = HashSet::new();
89 for node in &model.nodes {
90 if let Some((_, fact)) = node
91 .op
92 .state(&EvalContext::out_of_plan())?
93 .and_then(|state| state.init_tensor_fact())
94 {
95 symbols = fact.shape.volume().symbols();
96 break;
97 }
98 }
99 symbols
100 };
101
102 let b = tokens_symbols.intersection(&kv_symbols).cloned().collect::<HashSet<_>>();
103 let s = tokens_symbols.difference(&b).cloned().collect::<HashSet<_>>();
104 let p = kv_symbols.difference(&b).cloned().collect::<HashSet<_>>();
105 Ok((b.into_iter().next(), s.into_iter().next(), p.into_iter().next()))
106}
107
108pub fn memory_arena_hints_for_causal_llm(model: &TypedModel) -> TractResult<SymbolValues> {
109 let (b, s, p) = figure_out_causal_llm_b_s_p(model)?;
110 let mut values = SymbolValues::default()
111 .with(&s.context("Could not determine sequence_len (S)")?, 1024)
112 .with(&p.context("Could not determine past_sequence_len (P)")?, 0);
113 if let Some(b) = b {
114 values = values.with(&b, 1);
115 }
116 Ok(values)
117}