44
55"""Defines onnx export API"""
66
7+ import copy
78import contextlib
89import io
910import itertools
1011from packaging import version
1112import traceback
1213from typing import Any , Mapping , Tuple , Union , Literal
1314from pathlib import Path
15+ import math
1416
1517import numpy as np
1618import 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+
719901def _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