Skip to main content

tract_core/ops/matmul/
pack.rs

1use crate::axes::Axis;
2use crate::internal::*;
3use ndarray::*;
4use tract_linalg::block_quant::{
5    BlockQuantStorage, PackedBlockQuantFact, PackedBlockQuantFormat, block_quant_slice,
6};
7use tract_linalg::mmm::{MMMInputFormat, MMMInputValue, PackedMatrixStorage};
8
9use super::ModePicker;
10
11#[derive(Debug, Clone, PartialEq, Eq, Hash)]
12pub struct OptMatMulPack {
13    pub(crate) packers: Vec<Box<dyn MMMInputFormat>>,
14    pub(crate) mode_picker: ModePicker,
15    pub(crate) k_axis: usize,
16    pub(crate) mn_axis: usize,
17}
18
19impl Op for OptMatMulPack {
20    fn name(&self) -> StaticName {
21        "OptMatMulPack".into()
22    }
23
24    fn info(&self) -> TractResult<Vec<String>> {
25        Ok(vec![format!("{:?}. k axis: {}, mn axis: {}", self.packers, self.k_axis, self.mn_axis)])
26    }
27
28    op_as_typed_op!();
29}
30
31impl EvalOp for OptMatMulPack {
32    op_out_of_plan!();
33
34    fn eval(&self, ctx: &EvalContext, mut inputs: TVec<TValue>) -> TractResult<TVec<TValue>> {
35        self.do_eval(ctx, inputs.remove(0))
36    }
37}
38
39impl TypedOp for OptMatMulPack {
40    fn output_facts(&self, inputs: &[&TypedFact]) -> TractResult<TVec<TypedFact>> {
41        match self.mode_picker {
42            ModePicker::Single => ensure!(self.packers.len() == 1),
43            ModePicker::VecVsMat => ensure!(self.packers.len() == 2),
44        }
45        let k = inputs[0].shape[self.k_axis].clone();
46        let mn = inputs[0].shape[self.mn_axis].clone();
47        let exotic_fact = DynPackedExoticFact { k, mn, packers: self.packers.clone() };
48        Ok(tvec!(
49            inputs[0]
50                .datum_type
51                .fact(self.output_shape(&inputs[0].shape))
52                .with_exotic_fact(exotic_fact)
53        ))
54    }
55
56    fn axes_mapping(
57        &self,
58        inputs: &[&TypedFact],
59        outputs: &[&TypedFact],
60    ) -> TractResult<AxesMapping> {
61        let mut axes: Vec<Axis> = (0..inputs[0].rank())
62            .filter(|&ix| ix != self.k_axis && ix != self.mn_axis)
63            .enumerate()
64            .zip('a'..)
65            .map(|((o, i), repr)| Axis::new(repr, 1, 1).input(0, i).output(0, o))
66            .collect();
67        axes.push(Axis::new('K', 1, 1).input(0, self.k_axis));
68        axes.push(Axis::new('M', 1, 1).input(0, self.mn_axis));
69        axes.push(Axis::new('P', 1, 1).output(0, outputs[0].rank()));
70        AxesMapping::new(1, 1, axes)
71    }
72
73    as_op!();
74}
75
76impl OptMatMulPack {
77    fn do_eval(&self, _ctx: &EvalContext, input: TValue) -> TractResult<TVec<TValue>> {
78        unsafe {
79            let mode = self.mode_picker.pick(input.shape()[self.mn_axis])?;
80            let packer = &self.packers[mode];
81            let output_shape: TVec<usize> = self.output_shape(input.shape());
82            let stores = if output_shape.iter().all(|d| *d == 1) {
83                let packed = packer.prepare_one_view(&input.view(), self.k_axis, self.mn_axis)?;
84                PackedMatrixStorage::new_batched(&output_shape, tvec![packed])
85                    .into_tensor(input.datum_type())
86            } else {
87                let mut bc_shape: TVec<usize> = input.shape().into();
88                bc_shape[self.k_axis] = 1;
89                bc_shape[self.mn_axis] = 1;
90
91                let mut values: TVec<Box<dyn MMMInputValue>> =
92                    TVec::with_capacity(output_shape.iter().product());
93                for coord in indices(&*bc_shape) {
94                    let offset = coord
95                        .as_array_view()
96                        .iter()
97                        .zip(input.strides())
98                        .map(|(x, s)| *x as isize * s)
99                        .sum::<isize>()
100                        * input.datum_type().size_of() as isize;
101                    let view =
102                        TensorView::from_bytes(&input, offset, input.shape(), input.strides());
103                    values.push(packer.prepare_one_view(&view, self.k_axis, self.mn_axis)?);
104                }
105                PackedMatrixStorage::new_batched(&output_shape, values)
106                    .into_tensor(input.datum_type())
107            };
108            Ok(tvec!(stores.into_tvalue()))
109        }
110    }
111
112    pub fn output_shape<D: DimLike>(&self, input: &[D]) -> TVec<D> {
113        let mut packed_shape: TVec<D> = input.into();
114        packed_shape.remove(self.mn_axis.max(self.k_axis));
115        packed_shape.remove(self.mn_axis.min(self.k_axis));
116        packed_shape
117    }
118}
119
120#[derive(Hash, Clone, Debug, PartialEq, Eq)]
121pub struct DynPackedExoticFact {
122    pub k: TDim,
123    pub mn: TDim,
124    pub packers: Vec<Box<dyn MMMInputFormat>>,
125}
126
127impl ExoticFact for DynPackedExoticFact {
128    fn buffer_sizes(&self) -> TVec<TDim> {
129        tvec!(self.packers[0].mem_size(self.k.clone(), self.mn.clone()))
130    }
131}
132
133#[derive(Debug, Clone, Hash, Eq, PartialEq)]
134pub struct OptSimpleMatMulPack {
135    pub(crate) packed_format: PackedBlockQuantFormat,
136    pub(crate) k: usize,
137    pub(crate) m: usize,
138}
139
140impl Op for OptSimpleMatMulPack {
141    fn name(&self) -> StaticName {
142        "OptSimpleMatMulPack".into()
143    }
144    op_as_typed_op!();
145}
146
147impl EvalOp for OptSimpleMatMulPack {
148    op_out_of_plan!();
149
150    fn state(&self, _ctx: &EvalContext) -> TractResult<Option<Box<dyn OpState>>> {
151        Ok(None)
152    }
153
154    fn eval(&self, _ctx: &EvalContext, inputs: TVec<TValue>) -> TractResult<TVec<TValue>> {
155        let input = args_1!(inputs);
156        let bqs = input.try_storage_as::<BlockQuantStorage>()?;
157        // Leading dims before the last 2 (M, K) are batch/group dims
158        let num_groups: usize = input.shape()[..input.rank().saturating_sub(2)].iter().product();
159        let m_per_group = input.shape()[input.rank() - 2];
160        let k = *input.shape().last().unwrap();
161        let values = (0..num_groups)
162            .map(|g| {
163                let slice = block_quant_slice(bqs.value(), bqs.format(), m_per_group, k, g);
164                let iv: Box<dyn MMMInputValue> = Box::new(self.packed_format.pack(slice, k)?);
165                Ok(iv)
166            })
167            .collect::<TractResult<TVec<_>>>()?;
168        let leading_shape = &input.shape()[..input.rank().saturating_sub(2)];
169        let output =
170            PackedMatrixStorage::new_batched(leading_shape, values).into_tensor(input.datum_type());
171        Ok(tvec!(output.into_tvalue()))
172    }
173}
174
175impl TypedOp for OptSimpleMatMulPack {
176    fn output_facts(&self, inputs: &[&TypedFact]) -> TractResult<TVec<TypedFact>> {
177        let input = inputs[0];
178        // Input shape is [G, M, K] — output removes M and K, keeping leading dims
179        let output_shape: TVec<TDim> = if input.rank() > 2 {
180            input.shape[..input.rank() - 2].to_vec().into()
181        } else {
182            tvec!()
183        };
184        let fact =
185            inputs[0].datum_type.fact(&*output_shape).with_exotic_fact(PackedBlockQuantFact {
186                format: self.packed_format.clone(),
187                shape: tvec!(self.m, self.k),
188            });
189        Ok(tvec!(fact))
190    }
191
192    as_op!();
193}