Skip to main content

tract_core/ops/nn/
silu.rs

1use crate::internal::*;
2use crate::ops::element_wise::ElementWiseOp;
3use crate::ops::math::Mul;
4use crate::ops::nn::Sigmoid;
5
6use tract_data::half::f16;
7use tract_linalg::routines::Func;
8
9element_wise!(silu, Silu,
10    [f16] => |_, xs| { Func::Silu.ew_f16()?.run(xs) },
11    [f32] => |_, xs| { Func::Silu.ew_f32()?.run(xs) };
12    cost: |dt| {tvec!((Cost::FMA(dt), 12), (Cost::Div(dt), 1))};
13    declutter: detect_silu
14);
15
16/// Search pattern => A = A * SIGMOID(A)
17pub fn detect_silu(model: &TypedModel, node: &TypedNode) -> TractResult<Option<TypedModelPatch>> {
18    rule_if!(node.op_as::<ElementWiseOp>().is_some_and(|op| op.0.is::<Sigmoid>()));
19
20    let in_fact = model.node_input_facts(node.id)?[0];
21    let dt = in_fact.datum_type;
22
23    // Only F16 and F32 is supported.
24    rule_if!(matches!(dt, DatumType::F32 | DatumType::F16));
25
26    // Identify Mul successor: Sigmoid(A) * A
27    rule_if_some!(mul_succ = model.find_succ_bin_with_outlet::<Mul>(node, &node.inputs[0]));
28
29    let mut patch = TypedModelPatch::default();
30    let silu_input = patch.taps(model, &node.inputs)?;
31    let out = patch.wire_node(format!("{}.silu", node.name), silu(), &silu_input)?;
32    patch.shunt_outside(model, mul_succ.id.into(), out[0])?;
33    Ok(Some(patch))
34}