Skip to main content

tract_transformers/
lib.rs

1pub 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    // expectations:
74    // - one input is for tokens, so integer dt (i64 ?) and typically of shape S or 1,S, or B,S
75    // - other inputs are kv cache, some kind of float. shape features both S and P, and B if B is present in tokens
76    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        // Look for KVCache Op
88        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}