tract_core/ops/nn/
gelu_exact.rs1use crate::internal::*;
2use crate::ops::binary::TypedBinOp;
3use crate::ops::math::{Add, Mul};
4use tract_linalg::routines::Func;
5
6const CHUNK: usize = 1024;
7
8fn inv_sqrt2() -> f32 {
9 (2.0f32).sqrt().recip()
10}
11
12crate::element_wise!(gelu_exact, GeluExact,
13 [f16] => |_, xs| {
14 let erf = Func::Erf.ew_f32()?;
15 let c = f16::from_f32(inv_sqrt2());
16 let half = f16::from_f32(0.5);
17 let one = f16::from_f32(1.0);
18 let mut scratch = vec![0f32; xs.len().min(CHUNK)];
19 for chunk in xs.chunks_mut(CHUNK) {
20 let scaled = &mut scratch[..chunk.len()];
21 scaled.iter_mut().zip(chunk.iter()).for_each(|(s, x)| *s = (*x * c).to_f32());
22 erf.run(scaled)?;
23 chunk.iter_mut().zip(scaled.iter()).for_each(|(x, e)| {
24 *x = (*x * half) * (f16::from_f32(*e) + one);
25 });
26 }
27 Ok(())
28 },
29 [f32] => |_, xs| {
30 let erf = Func::Erf.ew_f32()?;
31 let c = inv_sqrt2();
32 let mut scratch = vec![0f32; xs.len().min(CHUNK)];
33 for chunk in xs.chunks_mut(CHUNK) {
34 let scaled = &mut scratch[..chunk.len()];
35 scaled.iter_mut().zip(chunk.iter()).for_each(|(s, x)| *s = *x * c);
36 erf.run(scaled)?;
37 chunk.iter_mut().zip(scaled.iter()).for_each(|(x, e)| *x = (*x * 0.5) * (*e + 1.0));
38 }
39 Ok(())
40 };
41 cost: |dt| {tvec!((Cost::FMA(dt), 14), (Cost::Div(dt), 1))}
42);
43
44pub fn detect_gelu_exact(
49 model: &TypedModel,
50 node: &TypedNode,
51) -> TractResult<Option<TypedModelPatch>> {
52 let erf_node = node;
53 let dt = model.node_input_facts(erf_node.id)?[0].datum_type;
54 rule_if!(matches!(dt, DatumType::F32 | DatumType::F16));
55
56 let scale = &model.nodes()[erf_node.inputs[0].node];
58 rule_if_some!(scale_op = scale.op_as::<TypedBinOp>());
59 rule_if!(scale_op.0.is::<Mul>());
60 rule_if!(model.matches_single_input_const(scale, inv_sqrt2()));
61 rule_if_some!(
62 x = scale
63 .inputs
64 .iter()
65 .find(|o| model.outlet_fact(**o).map(|f| f.konst.is_none()).unwrap_or(false))
66 .copied()
67 );
68
69 rule_if_some!(one_plus_erf = model.find_succ_bin_with_const::<Add>(erf_node, 1.0));
71
72 rule_if_some!(out = model.single_succ(one_plus_erf.id)?);
74 rule_if_some!(out_op = out.op_as::<TypedBinOp>());
75 rule_if!(out_op.0.is::<Mul>());
76 let half_x = out
77 .inputs
78 .iter()
79 .filter_map(|i| {
80 let n = &model.nodes()[i.node];
81 n.op_as::<TypedBinOp>()?.0.is::<Mul>().then_some(n)
82 })
83 .find(|n| model.matches_single_input_const(n, 0.5) && n.inputs.contains(&x));
84
85 let last = match half_x {
88 Some(_) => out,
89 None => {
90 rule_if!(out.inputs.contains(&x));
91 rule_if_some!(scaled = model.find_succ_bin_with_const::<Mul>(out, 0.5));
92 scaled
93 }
94 };
95
96 let mut patch = TypedModelPatch::default();
97 let tap = patch.taps(model, &[x])?;
98 let wired =
99 patch.wire_node(format!("{}.gelu_exact", erf_node.name), gelu_exact(), &[tap[0]])?;
100 patch.shunt_outside(model, last.id.into(), wired[0])?;
101 Ok(Some(patch))
102}
103
104#[cfg(test)]
105mod tests {
106 use super::*;
107 use crate::ops::element_wise::ElementWiseOp;
108 use crate::ops::math::{Erf, add, erf, mul};
109
110 fn chain(dt: DatumType, len: usize) -> TractResult<TypedModel> {
111 let mut m = TypedModel::default();
112 let x = m.add_source("x", dt.fact([len]))?;
113 let c = m.add_const("c", tensor1(&[inv_sqrt2()]).cast_to_dt(dt)?.into_owned())?;
114 let scaled = m.wire_node("scale", mul(), &[x, c])?[0];
115 let e = m.wire_node("erf", erf(), &[scaled])?[0];
116 let one = m.add_const("one", tensor1(&[1f32]).cast_to_dt(dt)?.into_owned())?;
117 let ope = m.wire_node("add_one", add(), &[e, one])?[0];
118 let half = m.add_const("half", tensor1(&[0.5f32]).cast_to_dt(dt)?.into_owned())?;
119 let hx = m.wire_node("half_x", mul(), &[x, half])?[0];
120 let out = m.wire_node("out", mul(), &[hx, ope])?;
121 m.select_output_outlets(&out)?;
122 Ok(m)
123 }
124
125 fn chain_trailing_half(dt: DatumType, len: usize) -> TractResult<TypedModel> {
126 let mut m = TypedModel::default();
127 let x = m.add_source("x", dt.fact([len]))?;
128 let c = m.add_const("c", tensor1(&[inv_sqrt2()]).cast_to_dt(dt)?.into_owned())?;
129 let scaled = m.wire_node("scale", mul(), &[x, c])?[0];
130 let e = m.wire_node("erf", erf(), &[scaled])?[0];
131 let one = m.add_const("one", tensor1(&[1f32]).cast_to_dt(dt)?.into_owned())?;
132 let ope = m.wire_node("add_one", add(), &[e, one])?[0];
133 let x_ope = m.wire_node("x_ope", mul(), &[x, ope])?[0];
134 let half = m.add_const("half", tensor1(&[0.5f32]).cast_to_dt(dt)?.into_owned())?;
135 let out = m.wire_node("out", mul(), &[x_ope, half])?;
136 m.select_output_outlets(&out)?;
137 Ok(m)
138 }
139
140 fn is_mini<T: crate::ops::element_wise::ElementWiseMiniOp>(n: &TypedNode) -> bool {
141 n.op_as::<ElementWiseOp>().map(|e| e.0.is::<T>()).unwrap_or(false)
142 }
143
144 fn input(dt: DatumType, len: usize) -> TractResult<TValue> {
145 let values: Vec<f32> = (0..len).map(|i| (i as f32 * 0.37).sin() * 4.0).collect();
146 Ok(tensor1(&values).cast_to_dt(dt)?.into_owned().into_tvalue())
147 }
148
149 #[test]
150 fn fuses_the_chain_and_keeps_the_values() -> TractResult<()> {
151 type Build = fn(DatumType, usize) -> TractResult<TypedModel>;
152 for build in [chain as Build, chain_trailing_half as Build] {
153 for dt in [DatumType::F32, DatumType::F16] {
154 let len = 2050;
155 let raw = build(dt, len)?;
156 let reference = raw.clone().into_runnable()?.run(tvec!(input(dt, len)?))?;
157
158 let fused = raw.into_decluttered()?;
159 assert!(
160 fused.nodes().iter().any(is_mini::<GeluExact>),
161 "{dt:?}: no GeluExact after declutter"
162 );
163 assert!(!fused.nodes().iter().any(is_mini::<Erf>), "{dt:?}: Erf survived");
164
165 let got = fused.into_runnable()?.run(tvec!(input(dt, len)?))?;
166 assert_eq!(
167 got[0].as_bytes(),
168 reference[0].as_bytes(),
169 "{dt:?}: fused output differs from the chain"
170 );
171 }
172 }
173 Ok(())
174 }
175}