Skip to content

Commit a203e14

Browse files
kyunggeunleeGitHub Enterprise
authored andcommitted
Precompute transposed MatMul weight in aimet-torch ONNX QDQ export time
Removed the intermediate transpose in `W -> QDQ -> Transpose -> MatMul` sequence which typically shows up when exporting nn.Linear with non-2D input. This complies better with canonical ONNX QDQ standards --------- Signed-off-by: Kyunggeun Lee <kyunggeu@qti.qualcomm.com>
1 parent 97e6411 commit a203e14

2 files changed

Lines changed: 321 additions & 8 deletions

File tree

TrainingExtensions/torch/src/python/aimet_torch/onnx.py

Lines changed: 183 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,13 +4,15 @@
44

55
"""Defines onnx export API"""
66

7+
import copy
78
import contextlib
89
import io
910
import itertools
1011
from packaging import version
1112
import traceback
1213
from typing import Any, Mapping, Tuple, Union, Literal
1314
from pathlib import Path
15+
import math
1416

1517
import numpy as np
1618
import onnx
@@ -716,6 +718,186 @@ def _remove_redundant_qdqs(onnx_model: onnx.ModelProto, base_dir):
716718
consumer.input[i] = qdq_node.input[0]
717719

718720

721+
def _fold_linear_weight_transpose(onnx_model: onnx.ModelProto, base_dir):
722+
"""
723+
Fold weight transpose of F.linear
724+
725+
Before:
726+
W -------> QDQ -> Transpose -> ...
727+
(Cout x Cin) (axis=0)
728+
After:
729+
W_t ------> QDQ -> ...
730+
(Cin x Cout) (axis=1)
731+
"""
732+
producers: dict[str, onnx.NodeProto] = {
733+
node.output[0]: node for node in onnx_model.graph.node
734+
}
735+
consumers: dict[str, list[onnx.NodeProto]] = {}
736+
all_tensor_names = set()
737+
738+
def _follow_identity_chain(name: str) -> str:
739+
producer = producers.get(name, None)
740+
while producer and producer.op_type == "Identity":
741+
name = producer.input[0]
742+
producer = producers.get(name, None)
743+
return name
744+
745+
for node in onnx_model.graph.node:
746+
if node.op_type == "Identity":
747+
continue
748+
for inp in node.input:
749+
inp = _follow_identity_chain(inp)
750+
consumers.setdefault(inp, []).append(node)
751+
752+
all_tensor_names.update(node.input)
753+
all_tensor_names.update(node.output)
754+
755+
constants: dict[str, onnx.TensorProto] = _get_all_constants(onnx_model, consumers)
756+
for node in onnx_model.graph.node:
757+
if node.op_type == "Identity" and node.input[0] in constants:
758+
constants[node.output[0]] = constants[node.input[0]]
759+
all_tensor_names.update(constants.keys())
760+
761+
qdq_nodes = {
762+
node.name: node
763+
for node in onnx_model.graph.node
764+
if (node.domain, node.op_type)
765+
in (
766+
("aimet", "quantize_dequantize"),
767+
("aimet", "QuantizeDequantize"),
768+
("aimet", "FloatQuantizeDequantize"),
769+
)
770+
}
771+
772+
def get_constant(name: str) -> onnx.TensorProto | None:
773+
name = _follow_identity_chain(name)
774+
return constants.get(name, None)
775+
776+
def get_new_name(base_name: str) -> str:
777+
i = 0
778+
new_name = f"{base_name}_{i}"
779+
while new_name in all_tensor_names:
780+
i += 1
781+
new_name = f"{base_name}_{i}"
782+
all_tensor_names.add(new_name)
783+
return new_name
784+
785+
def transpose_2d(tensor: onnx.TensorProto) -> onnx.TensorProto:
786+
assert len(tensor.dims) <= 2
787+
transposed = onnx.numpy_helper.from_array(
788+
onnx.numpy_helper.to_array(tensor, base_dir=base_dir).T,
789+
name=get_new_name(tensor.name),
790+
)
791+
transposed.dims[:] = [
792+
*reversed(tensor.dims),
793+
*(itertools.repeat(1, 2 - len(tensor.dims))),
794+
]
795+
return transposed
796+
797+
visited = set()
798+
799+
for weight in list(constants.values()):
800+
if weight.name in visited:
801+
continue
802+
visited.add(weight.name)
803+
weight_consumers = consumers.get(weight.name)
804+
805+
if not weight_consumers:
806+
continue
807+
808+
qdq_transpose_pairs: list[tuple[onnx.NodeProto, onnx.NodeProto]] = []
809+
810+
for qdq in weight_consumers:
811+
if not (
812+
# Weight should be consumed by QDQ
813+
qdq.name in qdq_nodes
814+
# Weight should be 1st input of QDQ
815+
and weight is get_constant(qdq.input[0])
816+
# Scale should be constant
817+
and get_constant(qdq.input[1])
818+
# Offset (if exists) should be constant
819+
and (
820+
qdq.op_type == "FloatQuantizeDequantize"
821+
or get_constant(qdq.input[2])
822+
)
823+
):
824+
continue
825+
826+
for transpose in consumers.get(qdq.output[0], []):
827+
if transpose.op_type != "Transpose":
828+
continue
829+
830+
perm = None
831+
for attr in transpose.attribute:
832+
if attr.name == "perm":
833+
perm = tuple(attr.ints)
834+
835+
if perm and perm != (1, 0):
836+
continue
837+
838+
if any(
839+
consumer.op_type
840+
in (
841+
"quantize_dequantize",
842+
"QuantizeDequantize",
843+
"FloatQuantizeDequantize",
844+
)
845+
for consumer in consumers.get(transpose.output[0], [])
846+
):
847+
continue
848+
849+
qdq_transpose_pairs.append((qdq, transpose))
850+
851+
if not qdq_transpose_pairs:
852+
continue
853+
854+
weight_t = transpose_2d(weight)
855+
onnx_model.graph.initializer.append(weight_t)
856+
constants[weight_t.name] = weight_t
857+
858+
replaced_nodes: dict[str, onnx.NodeProto] = {}
859+
for qdq, transpose in qdq_transpose_pairs:
860+
qdq_copy = copy.deepcopy(qdq)
861+
qdq_copy.input[0] = weight_t.name
862+
qdq_copy.output[0] = transpose.output[0]
863+
qdq_copy.name = f"{qdq.name}_copy"
864+
replaced_nodes[transpose.name] = qdq_copy
865+
866+
scale = get_constant(qdq.input[1])
867+
offset = (
868+
get_constant(qdq.input[2])
869+
if qdq.op_type in ("quantize_dequantize", "QuantizeDequantize")
870+
else None
871+
)
872+
873+
if math.prod(scale.dims) > 1:
874+
scale_t = transpose_2d(scale)
875+
onnx_model.graph.initializer.append(scale_t)
876+
qdq_copy.input[1] = scale_t.name
877+
constants[scale_t.name] = scale_t
878+
879+
if offset is not None and math.prod(offset.dims) > 1:
880+
offset_t = transpose_2d(offset)
881+
onnx_model.graph.initializer.append(offset_t)
882+
qdq_copy.input[2] = offset_t.name
883+
constants[offset_t.name] = offset_t
884+
885+
block_size = next(
886+
(attr for attr in qdq_copy.attribute if attr.name == "block_size"), None
887+
)
888+
if block_size and block_size.ints:
889+
# Transpose block_size in-place
890+
block_size.ints[:] = block_size.ints[::-1]
891+
892+
all_nodes = {}
893+
for node in onnx_model.graph.node:
894+
node = replaced_nodes.get(node.name, node)
895+
all_nodes[node.name] = node
896+
897+
onnx_model.graph.ClearField("node")
898+
onnx_model.graph.node.extend(all_nodes.values())
899+
900+
719901
def _to_onnx(
720902
model: torch.nn.Module,
721903
args: Union[Tuple[Any, ...], torch.Tensor],
@@ -725,6 +907,7 @@ def _to_onnx(
725907
base_dir = str(Path(str(f)).absolute().parent)
726908
_onnx.export(model, args, f, **kwargs)
727909
onnx_model = onnx.load(f, load_external_data=False)
910+
_fold_linear_weight_transpose(onnx_model, base_dir)
728911
aliases = _duplicate_shared_qdq_inputs(onnx_model, base_dir)
729912
_remove_redundant_qdqs(onnx_model, base_dir)
730913
_decouple_back_to_back_qdqs(onnx_model)

0 commit comments

Comments
 (0)