@@ -129,7 +129,7 @@ class AvgPool2d(GeneralOpDef):
129129
130130# TODO: Batch_norm op cannot directly map to QNN OpBatchnorm due to the number of input doesn't match.
131131@register_annotator (
132- [torch .ops .aten .batch_norm .default , torch . ops . aten . instance_norm . default ],
132+ [torch .ops .aten .batch_norm .default ],
133133 qnn_op = None ,
134134)
135135class BatchNorm (GeneralOpDef ):
@@ -420,7 +420,8 @@ def annotate(node: Node, quantization_config: QuantizationConfig) -> None:
420420 torch .ops .aten .topk .default ,
421421 torch .ops .aten .sort .default ,
422422 ):
423- out_act_quantization_spec = SharedQuantizationSpec (node .args [0 ])
423+ # assign to None since they are not supported so far
424+ out_act_quantization_spec = None
424425 node .meta [Q_ANNOTATION_KEY ] = QuantizationAnnotation (
425426 output_qspec = out_act_quantization_spec ,
426427 _annotated = True ,
@@ -807,21 +808,6 @@ class ReluMinMax(GeneralOpDef):
807808 pass
808809
809810
810- # TODO: Expand_as op cannot directly map to QNN OpTile due to the number of input doesn't match.
811- @register_annotator (
812- [
813- torch .ops .aten .expand_as .default ,
814- ],
815- qnn_op = None ,
816- )
817- class ExpandAs (GeneralOpDef ):
818- @staticmethod
819- def annotate (node : Node , quantization_config : QuantizationConfig ) -> None :
820- annotate_in_out_obs_sharing_op (node , quantization_config )
821- if not _is_annotated ([node ]):
822- annotate_single_in_share_out (node , quantization_config )
823-
824-
825811@register_annotator (
826812 [
827813 torch .ops .aten .flatten .using_ints ,
@@ -854,7 +840,6 @@ def annotate(node: Node, quantization_config: QuantizationConfig) -> None:
854840 return
855841
856842 act_node = node .args [0 ]
857- weight_node = node .args [2 ]
858843
859844 # TODO current only support 16a16w
860845 annotate_input_qspec_map (
@@ -863,94 +848,23 @@ def annotate(node: Node, quantization_config: QuantizationConfig) -> None:
863848 quantization_config .input_activation ,
864849 )
865850
866- annotate_input_qspec_map (
867- node ,
868- weight_node ,
869- quantization_config .input_activation ,
870- )
851+ if len (node .args ) > 2 and node .args [2 ] is not None :
852+ weight_node = node .args [2 ]
853+ annotate_input_qspec_map (
854+ node ,
855+ weight_node ,
856+ quantization_config .input_activation ,
857+ )
871858 nodes_to_mark_annotated = [node ]
872859 annotate_output_qspec (node , quantization_config .output_activation )
873860 _mark_nodes_as_annotated (nodes_to_mark_annotated )
874861
875862
876- # TODO: There is a bug in the BackendOpInfo library, so it is bypassed now.
877- @register_annotator ([torch .ops .aten .rsqrt .default ], qnn_op = None )
878- class Rsqrt (GeneralOpDef ):
879- pass
880-
881-
882863@register_annotator ([torch .ops .aten .scaled_dot_product_attention .default ], qnn_op = None )
883864class ScaledDotProductAttention (GeneralOpDef ):
884865 pass
885866
886867
887- @register_annotator (
888- [
889- torch .ops .aten .scatter .src ,
890- torch .ops .aten .scatter .value ,
891- torch .ops .aten .scatter_add .default ,
892- torch .ops .aten .scatter_reduce .two ,
893- ],
894- qnn_op = None ,
895- )
896- class ScatterElements (GeneralOpDef ):
897- @staticmethod
898- def annotate (node : Node , quantization_config : QuantizationConfig ) -> None :
899- if _is_annotated ([node ]):
900- return
901-
902- input_act = node .args [0 ]
903- if not isinstance (input_act , Node ) or not _is_float_tensor (input_act ):
904- return
905-
906- input_qspec_map = {}
907- input_qspec_map [input_act ] = quantization_config .input_activation
908-
909- if (
910- len (node .args ) > 3
911- and isinstance (node .args [3 ], Node )
912- and _is_float_tensor (node .args [3 ])
913- ):
914- input_qspec_map [node .args [3 ]] = SharedQuantizationSpec ((input_act , node ))
915-
916- output_act_qspec = (
917- SharedQuantizationSpec ((input_act , node ))
918- if _is_float_tensor (node )
919- else None
920- )
921-
922- if len (input_qspec_map ) > 0 or output_act_qspec is not None :
923- node .meta [Q_ANNOTATION_KEY ] = QuantizationAnnotation (
924- input_qspec_map = input_qspec_map ,
925- output_qspec = output_act_qspec ,
926- _annotated = True ,
927- )
928-
929-
930- @register_annotator ([torch .ops .aten .sort .default ], QnnConstants .OpTopK .op_name )
931- class Sort (GeneralOpDef ):
932- @staticmethod
933- def annotate (node : Node , quantization_config : QuantizationConfig ) -> None :
934- if _is_annotated ([node ]):
935- return
936-
937- input_qspec_map = {}
938- input_act_qspec = quantization_config .input_activation
939- out_act_quantization_spec = None
940- if input_act_qspec is not None :
941- if _is_float_tensor (node .args [0 ]):
942- input_act = node .args [0 ]
943- assert isinstance (input_act , Node )
944- input_qspec_map [input_act ] = input_act_qspec
945- out_act_quantization_spec = SharedQuantizationSpec ((input_act , node ))
946-
947- node .meta [Q_ANNOTATION_KEY ] = QuantizationAnnotation (
948- input_qspec_map = input_qspec_map ,
949- output_qspec = out_act_quantization_spec ,
950- _annotated = True ,
951- )
952-
953-
954868@register_annotator (
955869 [torch .ops .aten .sigmoid , torch .ops .aten .sigmoid .default ],
956870 QnnConstants .OpSigmoid .op_name ,
0 commit comments