diff --git a/onnxoptimizer/pass_registry.h b/onnxoptimizer/pass_registry.h index 3fdfa4978..951dea9d4 100644 --- a/onnxoptimizer/pass_registry.h +++ b/onnxoptimizer/pass_registry.h @@ -48,6 +48,7 @@ #include "onnxoptimizer/passes/fuse_consecutive_squeezes.h" #include "onnxoptimizer/passes/fuse_consecutive_transposes.h" #include "onnxoptimizer/passes/fuse_matmul_add_bias_into_gemm.h" +#include "onnxoptimizer/passes/fuse_matmul_add_bias_into_gemm_batched.h" #include "onnxoptimizer/passes/fuse_pad_into_conv.h" #include "onnxoptimizer/passes/fuse_pad_into_pool.h" #include "onnxoptimizer/passes/fuse_transpose_into_gemm.h" @@ -102,6 +103,7 @@ struct GlobalPassRegistry { registerPass(); registerPass(); registerPass(); + registerPass(); registerPass(); registerPass(); registerPass(); diff --git a/onnxoptimizer/passes/fuse_matmul_add_bias_into_gemm_batched.h b/onnxoptimizer/passes/fuse_matmul_add_bias_into_gemm_batched.h new file mode 100644 index 000000000..12e671171 --- /dev/null +++ b/onnxoptimizer/passes/fuse_matmul_add_bias_into_gemm_batched.h @@ -0,0 +1,206 @@ +// Copyright (c) ONNX Project Contributors +// +// SPDX-License-Identifier: Apache-2.0 + +// ATTENTION: The code in this file is highly EXPERIMENTAL. +// Adventurous users should note that the APIs will probably change. + +#pragma once + +// Before: +// Z = MatMul(X, W) // X rank >= 3, W a 2-D constant [K, N] +// A = Z + b // b broadcastable over the last (N) axis +// After: +// X2 = Reshape(X, [-1, K]) +// G = Gemm(X2, W, b) // alpha=beta=1, transA=transB=0 -> [-1, N] +// A = Reshape(G, [d0, ..., d_{r-2}, N]) +// +// `fuse_matmul_add_bias_into_gemm` only fuses the 2-D case. Transformer models +// apply linear layers to rank-3 activations `[B, S, K] . [K, N]`, so the MatMul +// is batched and that pass bails. This pass collapses the leading dims to a +// single 2-D Gemm and reshapes back. +// +// NOTE: This is a node-count / graph-shape rewrite (it trades a batched MatMul +// for Reshape + Gemm + Reshape). Batched MatMul and 2-D Gemm are equivalent +// work and modern runtimes execute batched MatMul natively, so this is not +// guaranteed to be faster. It is registered as PassType::Other so it is NOT in +// the default fuse set; invoke it explicitly by name when a Gemm-centric graph +// is wanted. + +#include + +#include "onnx/common/assertions.h" +#include "onnxoptimizer/pass.h" +#include "onnxoptimizer/passes/pass_util.h" + +namespace ONNX_NAMESPACE { +namespace optimization { + +struct FuseMatMulAddBiasIntoGemmBatched final : public PredicateBasedPass { + explicit FuseMatMulAddBiasIntoGemmBatched() + : PredicateBasedPass(PassType::Other, PassEfficiency::Complete, + PassOptimizationType::Compute) {} + std::string getPassName() const override { + return "fuse_matmul_add_bias_into_gemm_batched"; + } + + static Value* MakeInt64Constant(Graph& graph, std::vector data) { + Tensor t; + t.sizes().push_back(static_cast(data.size())); + t.elem_type() = TensorProto_DataType_INT64; + t.int64s() = std::move(data); + return graph.addInitializerAndCreateValue(t); + } + + bool patternMatchPredicate(Node* node) override { + if (!CheckKind(node, kAdd, 0, kMatMul)) { + return false; + } + Value* matmul_out = node->input(0); + // MatMul result must feed only this Add. + if (matmul_out->uses().size() > 1) { + return false; + } + Node* matmul = matmul_out->node(); + Value* x = matmul->input(0); + Value* w = matmul->input(1); + + // X: rank >= 3 with a static trailing dim K. + if (!x->has_sizes()) { + return false; + } + const auto& x_shape = x->sizes(); + if (x_shape.size() < 3 || !x_shape.back().is_int) { + return false; + } + const int64_t k = x_shape.back().dim; + + // W: a 2-D constant [K, N] with static dims, matching K. + if (!IsConstantTensor(w) || !w->has_sizes()) { + return false; + } + const auto& w_shape = w->sizes(); + if (w_shape.size() != 2 || !w_shape[0].is_int || !w_shape[1].is_int) { + return false; + } + if (w_shape[0].dim != k) { + return false; + } + const int64_t n = w_shape[1].dim; + + // bias: 1-D, broadcastable over the N axis only ([N] or [1]). + Value* bias = node->input(1); + if (!bias->has_sizes()) { + return false; + } + const auto& bias_shape = bias->sizes(); + if (bias_shape.size() != 1 || !bias_shape[0].is_int) { + return false; + } + if (bias_shape[0].dim != n && bias_shape[0].dim != 1) { + return false; + } + + // Dynamic leading dims need Shape/Slice, i.e. opset >= 10. + bool leading_static = true; + for (size_t i = 0; i + 1 < x_shape.size(); ++i) { + leading_static &= x_shape[i].is_int; + } + if (!leading_static) { + const int opset = getOpsetVersion(*node->owningGraph()); + if (opset != 0 && opset < 10) { + return false; + } + } + return true; + } + + bool runTransform(Node* n, Graph& graph, + NodeDestroyType& destroy_current) override { + destroy_current = NodeDestroyType::DestroyZero; + + Node* matmul = n->input(0)->node(); + Value* x = matmul->input(0); + Value* w = matmul->input(1); + Value* bias = n->input(1); + + const auto& x_shape = x->sizes(); + const int64_t rank = static_cast(x_shape.size()); + const int64_t k = x_shape.back().dim; + const int64_t out_n = w->sizes()[1].dim; + + // X2 = Reshape(X, [-1, K]) + Node* pre = graph.create(kReshape, 1); + pre->addInput(x); + pre->addInput(MakeInt64Constant(graph, {-1, k})); + + // G = Gemm(X2, W, bias) + Node* gemm = graph.create(kGemm, 1); + gemm->addInput(pre->output()); + gemm->addInput(w); + gemm->addInput(bias); + gemm->f_(kalpha, 1.0); + gemm->f_(kbeta, 1.0); + gemm->i_(ktransA, 0); + gemm->i_(ktransB, 0); + + // Reconstruct the output shape [d0, ..., d_{r-2}, N]. + bool leading_static = true; + std::vector leading_dims; + for (int64_t i = 0; i + 1 < rank; ++i) { + leading_static &= x_shape[i].is_int; + if (x_shape[i].is_int) { + leading_dims.push_back(x_shape[i].dim); + } + } + + // A = Reshape(G, out_shape) + Node* post = graph.create(kReshape, n->outputs().size()); + post->addInput(gemm->output()); + + // Order nodes: pre -> gemm -> [shape ops] -> post -> n. + pre->insertBefore(n); + gemm->insertBefore(n); + + if (leading_static) { + std::vector out_shape = leading_dims; + out_shape.push_back(out_n); + post->addInput(MakeInt64Constant(graph, std::move(out_shape))); + } else { + // shape(X) -> slice off the last dim -> concat with [N] + Node* shape = graph.create(Symbol("Shape"), 1); + shape->addInput(x); + shape->insertBefore(n); + + Node* slice = graph.create(kSlice, 1); + slice->addInput(shape->output()); + slice->addInput(MakeInt64Constant(graph, {0})); // starts + slice->addInput(MakeInt64Constant(graph, {rank - 1})); // ends + slice->addInput(MakeInt64Constant(graph, {0})); // axes + slice->insertBefore(n); + + Node* concat = graph.create(kConcat, 1); + concat->addInput(slice->output()); + concat->addInput(MakeInt64Constant(graph, {out_n})); + concat->i_(kaxis, 0); + concat->insertBefore(n); + + post->addInput(concat->output()); + } + + for (int i = 0; i < static_cast(n->outputs().size()); ++i) { + post->outputs()[i]->copyMetadata(n->outputs()[i]); + } + post->insertBefore(n); + + if (!tryReplacingAllUsesWith(n, post)) { + return false; + } + // Destroy the Add; the now-dead MatMul is cleaned up by DCE. + destroy_current = NodeDestroyType::DestroyOne; + return true; + } +}; + +} // namespace optimization +} // namespace ONNX_NAMESPACE diff --git a/onnxoptimizer/test/optimizer_test.py b/onnxoptimizer/test/optimizer_test.py index 45ff7ca6a..5c52ba12b 100644 --- a/onnxoptimizer/test/optimizer_test.py +++ b/onnxoptimizer/test/optimizer_test.py @@ -1649,6 +1649,105 @@ def test_fuse_matmul_add_bias_into_gemm_multiple_use_no_fuse(self): assert optimized_model.graph == graph + def test_fuse_matmul_add_bias_into_gemm_batched(self): # type: () -> None + matmul = helper.make_node("MatMul", ["X", "W"], ["Z"]) + add = helper.make_node("Add", ["Z", "B"], ["A"]) + w = numpy_helper.from_array( + np.random.randn(4, 5).astype(np.float32), name="W" + ) + b = numpy_helper.from_array(np.random.randn(5).astype(np.float32), name="B") + graph = helper.make_graph( + [matmul, add], + "test", + [helper.make_tensor_value_info("X", TensorProto.FLOAT, (2, 3, 4))], + [helper.make_tensor_value_info("A", TensorProto.FLOAT, (2, 3, 5))], + initializer=[w, b], + ) + optimized_model = self._optimized( + graph, + ["fuse_matmul_add_bias_into_gemm_batched", "eliminate_deadend"], + ) + + op_types = [n.op_type for n in optimized_model.graph.node] + assert op_types == ["Reshape", "Gemm", "Reshape"] + assert "MatMul" not in op_types + + def test_fuse_matmul_add_bias_into_gemm_batched_dynamic(self): + # dynamic leading dims -> Shape/Slice/Concat rebuild the output shape + matmul = helper.make_node("MatMul", ["X", "W"], ["Z"]) + add = helper.make_node("Add", ["Z", "B"], ["A"]) + w = numpy_helper.from_array( + np.random.randn(4, 5).astype(np.float32), name="W" + ) + b = numpy_helper.from_array(np.random.randn(5).astype(np.float32), name="B") + graph = helper.make_graph( + [matmul, add], + "test", + [ + helper.make_tensor_value_info( + "X", TensorProto.FLOAT, ("batch", "seq", 4) + ) + ], + [ + helper.make_tensor_value_info( + "A", TensorProto.FLOAT, ("batch", "seq", 5) + ) + ], + initializer=[w, b], + ) + optimized_model = self._optimized( + graph, + ["fuse_matmul_add_bias_into_gemm_batched", "eliminate_deadend"], + compare_result=False, + ) + + op_types = [n.op_type for n in optimized_model.graph.node] + assert "Gemm" in op_types + assert "MatMul" not in op_types + assert op_types[-1] == "Reshape" + + def test_fuse_matmul_add_bias_into_gemm_batched_2d_no_fuse(self): + # rank-2 matmul is handled by the non-batched pass; this one must skip it + matmul = helper.make_node("MatMul", ["X", "W"], ["Z"]) + add = helper.make_node("Add", ["Z", "B"], ["A"]) + w = numpy_helper.from_array( + np.random.randn(4, 5).astype(np.float32), name="W" + ) + b = numpy_helper.from_array(np.random.randn(5).astype(np.float32), name="B") + graph = helper.make_graph( + [matmul, add], + "test", + [helper.make_tensor_value_info("X", TensorProto.FLOAT, (3, 4))], + [helper.make_tensor_value_info("A", TensorProto.FLOAT, (3, 5))], + initializer=[w, b], + ) + optimized_model = self._optimized( + graph, ["fuse_matmul_add_bias_into_gemm_batched"] + ) + + assert [n.op_type for n in optimized_model.graph.node] == ["MatMul", "Add"] + + def test_fuse_matmul_add_bias_into_gemm_batched_nonconst_weight_no_fuse(self): + # W is a runtime input, not a constant -> cannot form a Gemm weight + matmul = helper.make_node("MatMul", ["X", "W"], ["Z"]) + add = helper.make_node("Add", ["Z", "B"], ["A"]) + b = numpy_helper.from_array(np.random.randn(5).astype(np.float32), name="B") + graph = helper.make_graph( + [matmul, add], + "test", + [ + helper.make_tensor_value_info("X", TensorProto.FLOAT, (2, 3, 4)), + helper.make_tensor_value_info("W", TensorProto.FLOAT, (4, 5)), + ], + [helper.make_tensor_value_info("A", TensorProto.FLOAT, (2, 3, 5))], + initializer=[b], + ) + optimized_model = self._optimized( + graph, ["fuse_matmul_add_bias_into_gemm_batched"] + ) + + assert [n.op_type for n in optimized_model.graph.node] == ["MatMul", "Add"] + # type: () -> None def test_fuse_pad_into_conv_no_optional_value_opset10(self): pad = helper.make_node(