From b57a85d8638f8517ec11a1c3f237eb42413c4360 Mon Sep 17 00:00:00 2001 From: Ye Wang <52801275+wangyems@users.noreply.github.com> Date: Wed, 10 Mar 2021 21:37:12 -0800 Subject: [PATCH] Support symbolic shape infer in transformers tool (#6899) * fusion support runtime edge shape checking * trim ctor * add test * fix * Update test_shape_infer_helper.py * use torch input size as dynamic axis hints * check dir * update * support longformerattention * update and add support for bert ops * trim * review comments * review comments --- .../python/tools/symbolic_shape_infer.py | 544 +++++++++++------- .../transformers/fusion_skiplayernorm.py | 9 + .../python/tools/transformers/onnx_model.py | 16 + .../tools/transformers/shape_infer_helper.py | 84 +++ .../test/test_shape_infer_helper.py | 48 ++ 5 files changed, 498 insertions(+), 203 deletions(-) create mode 100644 onnxruntime/python/tools/transformers/shape_infer_helper.py create mode 100644 onnxruntime/python/tools/transformers/test/test_shape_infer_helper.py diff --git a/onnxruntime/python/tools/symbolic_shape_infer.py b/onnxruntime/python/tools/symbolic_shape_infer.py index 5fc3669262..d048d816b4 100755 --- a/onnxruntime/python/tools/symbolic_shape_infer.py +++ b/onnxruntime/python/tools/symbolic_shape_infer.py @@ -12,28 +12,35 @@ import sympy from packaging import version assert version.parse(onnx.__version__) >= version.parse("1.5.0") + def get_attribute(node, attr_name, default_value=None): found = [attr for attr in node.attribute if attr.name == attr_name] if found: return helper.get_attribute_value(found[0]) return default_value + def get_dim_from_type_proto(dim): return getattr(dim, dim.WhichOneof('value')) if type(dim.WhichOneof('value')) == str else None + def get_shape_from_type_proto(type_proto): return [get_dim_from_type_proto(d) for d in type_proto.tensor_type.shape.dim] + def get_shape_from_sympy_shape(sympy_shape): return [None if i is None else (int(i) if is_literal(i) else str(i)) for i in sympy_shape] + def is_literal(dim): return type(dim) in [int, np.int64, np.int32, sympy.Integer] or (hasattr(dim, 'is_number') and dim.is_number) + def handle_negative_axis(axis, rank): assert axis < rank and axis >= -rank return axis if axis >= 0 else rank + axis + def get_opset(mp, domain=None): domain = domain or ['', 'onnx', 'ai.onnx'] if type(domain) != list: @@ -43,6 +50,7 @@ def get_opset(mp, domain=None): return opset.version return None + def as_scalar(x): if type(x) == list: assert len(x) == 1 @@ -52,6 +60,7 @@ def as_scalar(x): else: return x + def as_list(x, keep_none): if type(x) == list: return x @@ -62,6 +71,7 @@ def as_list(x, keep_none): else: return [x] + def sympy_reduce_product(x): if type(x) == list: value = sympy.Integer(1) @@ -71,60 +81,70 @@ def sympy_reduce_product(x): value = x return value + class SymbolicShapeInference: def __init__(self, int_max, auto_merge, guess_output_rank, verbose): self.dispatcher_ = { - 'Add' : self._infer_symbolic_compute_ops, - 'ArrayFeatureExtractor' : self._infer_ArrayFeatureExtractor, - 'AveragePool' : self._infer_Pool, - 'Cast' : self._infer_Cast, - 'CategoryMapper' : self._infer_CategoryMapper, - 'Compress' : self._infer_Compress, - 'Concat' : self._infer_Concat, - 'Constant' : self._infer_Constant, - 'ConstantOfShape' : self._infer_ConstantOfShape, - 'Conv' : self._infer_Conv, - 'CumSum' : self._pass_on_shape_and_type, - 'Div' : self._infer_symbolic_compute_ops, - 'Expand' : self._infer_Expand, - 'Equal' : self._infer_symbolic_compute_ops, - 'Floor' : self._infer_symbolic_compute_ops, - 'Gather' : self._infer_Gather, - 'GatherElements' : self._infer_GatherElements, - 'GatherND' : self._infer_GatherND, - 'Gelu' : self._pass_on_shape_and_type, - 'If' : self._infer_If, - 'Loop' : self._infer_Loop, - 'MatMul' : self._infer_MatMul, - 'MatMulInteger16' : self._infer_MatMulInteger, - 'MaxPool' : self._infer_Pool, - 'Max' : self._infer_symbolic_compute_ops, - 'Min' : self._infer_symbolic_compute_ops, - 'Mul' : self._infer_symbolic_compute_ops, - 'NonMaxSuppression' : self._infer_NonMaxSuppression, - 'NonZero' : self._infer_NonZero, - 'OneHot' : self._infer_OneHot, - 'Pad' : self._infer_Pad, - 'Range' : self._infer_Range, - 'ReduceProd' : self._infer_ReduceProd, - 'Reshape' : self._infer_Reshape, - 'Resize' : self._infer_Resize, - 'Round' : self._pass_on_shape_and_type, - 'Scan' : self._infer_Scan, - 'ScatterElements' : self._infer_ScatterElements, - 'Shape' : self._infer_Shape, - 'Size' : self._infer_Size, - 'Slice' : self._infer_Slice, - 'SoftmaxCrossEntropyLoss':self._infer_SoftmaxCrossEntropyLoss, - 'Split' : self._infer_Split, - 'SplitToSequence' : self._infer_SplitToSequence, - 'Squeeze' : self._infer_Squeeze, - 'Sub' : self._infer_symbolic_compute_ops, - 'Tile' : self._infer_Tile, - 'TopK' : self._infer_TopK, - 'Unsqueeze' : self._infer_Unsqueeze, - 'Where' : self._infer_symbolic_compute_ops, - 'ZipMap' : self._infer_ZipMap} + 'Add': self._infer_symbolic_compute_ops, + 'ArrayFeatureExtractor': self._infer_ArrayFeatureExtractor, + 'AveragePool': self._infer_Pool, + 'Cast': self._infer_Cast, + 'CategoryMapper': self._infer_CategoryMapper, + 'Compress': self._infer_Compress, + 'Concat': self._infer_Concat, + 'Constant': self._infer_Constant, + 'ConstantOfShape': self._infer_ConstantOfShape, + 'Conv': self._infer_Conv, + 'CumSum': self._pass_on_shape_and_type, + 'Div': self._infer_symbolic_compute_ops, + 'Expand': self._infer_Expand, + 'Equal': self._infer_symbolic_compute_ops, + 'Floor': self._infer_symbolic_compute_ops, + 'Gather': self._infer_Gather, + 'GatherElements': self._infer_GatherElements, + 'GatherND': self._infer_GatherND, + 'Gelu': self._pass_on_shape_and_type, + 'If': self._infer_If, + 'Loop': self._infer_Loop, + 'MatMul': self._infer_MatMul, + 'MatMulInteger16': self._infer_MatMulInteger, + 'MaxPool': self._infer_Pool, + 'Max': self._infer_symbolic_compute_ops, + 'Min': self._infer_symbolic_compute_ops, + 'Mul': self._infer_symbolic_compute_ops, + 'NonMaxSuppression': self._infer_NonMaxSuppression, + 'NonZero': self._infer_NonZero, + 'OneHot': self._infer_OneHot, + 'Pad': self._infer_Pad, + 'Range': self._infer_Range, + 'ReduceProd': self._infer_ReduceProd, + 'Reshape': self._infer_Reshape, + 'Resize': self._infer_Resize, + 'Round': self._pass_on_shape_and_type, + 'Scan': self._infer_Scan, + 'ScatterElements': self._infer_ScatterElements, + 'Shape': self._infer_Shape, + 'Size': self._infer_Size, + 'Slice': self._infer_Slice, + 'SoftmaxCrossEntropyLoss': self._infer_SoftmaxCrossEntropyLoss, + 'Split': self._infer_Split, + 'SplitToSequence': self._infer_SplitToSequence, + 'Squeeze': self._infer_Squeeze, + 'Sub': self._infer_symbolic_compute_ops, + 'Tile': self._infer_Tile, + 'TopK': self._infer_TopK, + 'Unsqueeze': self._infer_Unsqueeze, + 'Where': self._infer_symbolic_compute_ops, + 'ZipMap': self._infer_ZipMap, + # contrib ops: + 'Attention': self._infer_Attention, + 'BiasGelu': self._infer_BiasGelu, + 'FastGelu': self._infer_FastGelu, + 'Gelu': self._infer_Gelu, + 'LayerNormalization': self._infer_LayerNormalization, + 'LongformerAttention': self._infer_LongformerAttention, + 'SkipLayerNormalization': self._infer_SkipLayerNormalization + } self.run_ = True self.suggested_merge_ = {} self.symbolic_dims_ = {} @@ -137,7 +157,7 @@ class SymbolicShapeInference: def _add_suggested_merge(self, symbols, apply=False): assert all([(type(s) == str and s in self.symbolic_dims_) or is_literal(s) for s in symbols]) symbols = set(symbols) - for k,v in self.suggested_merge_.items(): + for k, v in self.suggested_merge_.items(): if k in symbols: symbols.remove(k) symbols.add(v) @@ -173,7 +193,7 @@ class SymbolicShapeInference: if is_literal(map_to) and is_literal(s): assert int(map_to) == int(s) self.suggested_merge_[s] = int(map_to) if is_literal(map_to) else map_to - for k,v in self.suggested_merge_.items(): + for k, v in self.suggested_merge_.items(): if v == s: self.suggested_merge_[k] = map_to if apply and self.auto_merge_: @@ -196,24 +216,27 @@ class SymbolicShapeInference: self.out_mp_.CopyFrom(in_mp) self.initializers_ = dict([(i.name, i) for i in self.out_mp_.graph.initializer]) self.known_vi_ = dict([(i.name, i) for i in list(self.out_mp_.graph.input)]) - self.known_vi_.update(dict([(i.name, helper.make_tensor_value_info(i.name, i.data_type, list(i.dims))) for i in self.out_mp_.graph.initializer])) + self.known_vi_.update( + dict([(i.name, helper.make_tensor_value_info(i.name, i.data_type, list(i.dims))) + for i in self.out_mp_.graph.initializer])) def _merge_symbols(self, dims): if not all([type(d) == str for d in dims]): if self.auto_merge_: unique_dims = list(set(dims)) is_int = [is_literal(d) for d in unique_dims] - assert sum(is_int) <= 1 # if there are more than 1 unique ints, something is wrong + assert sum(is_int) <= 1 # if there are more than 1 unique ints, something is wrong if sum(is_int) == 1: - int_dim = is_int.index(1) - if self.verbose_ > 0: - print('dim {} has been merged with value {}'.format(unique_dims[:int_dim] + unique_dims[int_dim+1:], unique_dims[int_dim])) - self._check_merged_dims(unique_dims, allow_broadcast=False) - return unique_dims[int_dim] + int_dim = is_int.index(1) + if self.verbose_ > 0: + print('dim {} has been merged with value {}'.format( + unique_dims[:int_dim] + unique_dims[int_dim + 1:], unique_dims[int_dim])) + self._check_merged_dims(unique_dims, allow_broadcast=False) + return unique_dims[int_dim] else: - if self.verbose_ > 0: - print('dim {} has been mergd with dim {}'.format(unique_dims[1:], unique_dims[0])) - return dims[0] + if self.verbose_ > 0: + print('dim {} has been mergd with dim {}'.format(unique_dims[1:], unique_dims[0])) + return dims[0] else: return None if all([d == dims[0] for d in dims]): @@ -266,7 +289,8 @@ class SymbolicShapeInference: sympy_shape = [] for d in self._get_shape(node, idx): if type(d) == str: - sympy_shape.append(self.symbolic_dims_[d] if d in self.symbolic_dims_ else sympy.Symbol(d, integer=True)) + sympy_shape.append(self.symbolic_dims_[d] if d in + self.symbolic_dims_ else sympy.Symbol(d, integer=True)) else: assert None != d sympy_shape.append(d) @@ -291,7 +315,7 @@ class SymbolicShapeInference: str_dim = str(new_dim) if str_dim in self.suggested_merge_: if is_literal(self.suggested_merge_[str_dim]): - continue # no need to create dim for literals + continue # no need to create dim for literals new_sympy_shape[i] = self.symbolic_dims_[self.suggested_merge_[str_dim]] else: # add new_dim if it's a computational expression @@ -305,10 +329,9 @@ class SymbolicShapeInference: # run single node inference with self.known_vi_ shapes # note that inference rely on initializer values is not handled # as we don't copy initializer weights to tmp_graph for inference speed purpose - tmp_graph = helper.make_graph([node], - 'tmp', - [self.known_vi_[i] for i in node.input if i], - [helper.make_tensor_value_info(i, onnx.TensorProto.UNDEFINED, None) for i in node.output]) + tmp_graph = helper.make_graph( + [node], 'tmp', [self.known_vi_[i] for i in node.input if i], + [helper.make_tensor_value_info(i, onnx.TensorProto.UNDEFINED, None) for i in node.output]) self.tmp_mp_.graph.CopyFrom(tmp_graph) self.tmp_mp_ = shape_inference.infer_shapes(self.tmp_mp_) @@ -323,22 +346,24 @@ class SymbolicShapeInference: def _onnx_infer_subgraph(self, node, subgraph, use_node_input=True): if self.verbose_ > 2: - print('Inferencing subgraph of node {} with output({}...): {}'.format(node.name, node.output[0], node.op_type)) + print('Inferencing subgraph of node {} with output({}...): {}'.format(node.name, node.output[0], + node.op_type)) # node inputs are not passed directly to the subgraph # it's up to the node dispatcher to prepare subgraph input # for example, with Scan/Loop, subgraph input shape would be trimmed from node input shape # besides, inputs in subgraph could shadow implicit inputs subgraph_inputs = set([i.name for i in list(subgraph.initializer) + list(subgraph.input)]) subgraph_implicit_input = set([name for name in self.known_vi_.keys() if not name in subgraph_inputs]) - tmp_graph = helper.make_graph(list(subgraph.node), - 'tmp', - list(subgraph.input) + [self.known_vi_[i] for i in subgraph_implicit_input], - [helper.make_tensor_value_info(i.name, onnx.TensorProto.UNDEFINED, None) for i in subgraph.output]) + tmp_graph = helper.make_graph( + list(subgraph.node), 'tmp', + list(subgraph.input) + [self.known_vi_[i] for i in subgraph_implicit_input], + [helper.make_tensor_value_info(i.name, onnx.TensorProto.UNDEFINED, None) for i in subgraph.output]) tmp_graph.initializer.extend([i for i in self.out_mp_.graph.initializer if i.name in subgraph_implicit_input]) tmp_graph.initializer.extend(subgraph.initializer) self.tmp_mp_.graph.CopyFrom(tmp_graph) - symbolic_shape_inference = SymbolicShapeInference(self.int_max_, self.auto_merge_, self.guess_output_rank_, self.verbose_) + symbolic_shape_inference = SymbolicShapeInference(self.int_max_, self.auto_merge_, self.guess_output_rank_, + self.verbose_) all_shapes_inferred = False symbolic_shape_inference._preprocess(self.tmp_mp_) symbolic_shape_inference.suggested_merge_ = self.suggested_merge_.copy() @@ -357,7 +382,8 @@ class SymbolicShapeInference: subgraph.node.extend(symbolic_shape_inference.out_mp_.graph.node) # for new symbolic dims from subgraph output, add to main graph symbolic dims subgraph_shapes = [get_shape_from_type_proto(o.type) for o in symbolic_shape_inference.out_mp_.graph.output] - subgraph_new_symbolic_dims = set([d for s in subgraph_shapes if s for d in s if type(d) == str and not d in self.symbolic_dims_]) + subgraph_new_symbolic_dims = set( + [d for s in subgraph_shapes if s for d in s if type(d) == str and not d in self.symbolic_dims_]) new_dims = {} for d in subgraph_new_symbolic_dims: assert d in symbolic_shape_inference.symbolic_dims_ @@ -369,11 +395,11 @@ class SymbolicShapeInference: values = [self._try_get_value(node, i) for i in range(len(node.input))] if all([v is not None for v in values]): # some shape compute is in floating point, cast to int for sympy - for i,v in enumerate(values): + for i, v in enumerate(values): if type(v) != np.ndarray: continue if len(v.shape) > 1: - new_v = None # ignore value for rank > 1 + new_v = None # ignore value for rank > 1 elif len(v.shape) == 0: new_v = int(v.item()) else: @@ -384,16 +410,16 @@ class SymbolicShapeInference: max_len = max(values_len) if max_len >= 1 and broadcast: # broadcast - for i,v in enumerate(values): + for i, v in enumerate(values): if v is None: - continue # don't broadcast if value is unknown + continue # don't broadcast if value is unknown if type(v) == list: if len(v) < max_len: - values[i] = v*max_len + values[i] = v * max_len else: assert len(v) == max_len else: - values[i] = [v]*max_len + values[i] = [v] * max_len return values def _compute_on_sympy_data(self, node, op_func): @@ -413,9 +439,9 @@ class SymbolicShapeInference: def _pass_on_shape_and_type(self, node): vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - self.known_vi_[node.input[0]].type.tensor_type.elem_type, - self._get_shape(node, 0))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + self._get_shape(node, 0))) def _new_symbolic_dim(self, prefix, dim): new_dim = '{}_d{}'.format(prefix, dim) @@ -427,7 +453,9 @@ class SymbolicShapeInference: return new_dim def _new_symbolic_dim_from_output(self, node, out_idx=0, dim=0): - return self._new_symbolic_dim('{}{}_o{}_'.format(node.op_type, list(self.out_mp_.graph.node).index(node), out_idx), dim) + return self._new_symbolic_dim( + '{}{}_o{}_'.format(node.op_type, + list(self.out_mp_.graph.node).index(node), out_idx), dim) def _new_symbolic_shape(self, rank, node, out_idx=0): return [self._new_symbolic_dim_from_output(node, out_idx, i) for i in range(rank)] @@ -436,7 +464,7 @@ class SymbolicShapeInference: sympy_shape = self._get_sympy_shape(node, 0) if len(node.input) > 1: W_shape = self._get_sympy_shape(node, 1) - rank = len(W_shape) - 2 # number of spatial axes + rank = len(W_shape) - 2 # number of spatial axes kernel_shape = W_shape[-rank:] sympy_shape[1] = W_shape[0] else: @@ -456,25 +484,29 @@ class SymbolicShapeInference: sympy_shape[-rank:] = [sympy.Integer(d) for d in shape[-rank:]] return sympy_shape - dilations = get_attribute(node, 'dilations', [1]*rank) - strides = get_attribute(node, 'strides', [1]*rank) + dilations = get_attribute(node, 'dilations', [1] * rank) + strides = get_attribute(node, 'strides', [1] * rank) effective_kernel_shape = [(k - 1) * d + 1 for k, d in zip(kernel_shape, dilations)] pads = get_attribute(node, 'pads') if pads is None: - pads = [0]*(2*rank) + pads = [0] * (2 * rank) auto_pad = get_attribute(node, 'auto_pad', b'NOTSET').decode('utf-8') if auto_pad != 'VALID' and auto_pad != 'NOTSET': try: residual = [sympy.Mod(d, s) for d, s in zip(sympy_shape[-rank:], strides)] - total_pads = [max(0, (k - s) if r == 0 else (k - r)) for k, s, r in zip(effective_kernel_shape, strides, residual)] - except TypeError: # sympy may throw TypeError: cannot determine truth value of Relational - total_pads = [max(0, (k - s)) for k, s in zip(effective_kernel_shape, strides)] # assuming no residual if sympy throws error + total_pads = [ + max(0, (k - s) if r == 0 else (k - r)) + for k, s, r in zip(effective_kernel_shape, strides, residual) + ] + except TypeError: # sympy may throw TypeError: cannot determine truth value of Relational + total_pads = [max(0, (k - s)) for k, s in zip(effective_kernel_shape, strides) + ] # assuming no residual if sympy throws error elif auto_pad == 'VALID': total_pads = [] else: - total_pads = [0]*rank + total_pads = [0] * rank else: - assert len(pads) == 2*rank + assert len(pads) == 2 * rank total_pads = [p1 + p2 for p1, p2 in zip(pads[:rank], pads[rank:])] ceil_mode = get_attribute(node, 'ceil_mode', 0) @@ -483,7 +515,8 @@ class SymbolicShapeInference: if len(total_pads) > 0: effective_input_size = effective_input_size + total_pads[i] if ceil_mode: - strided_kernel_positions = sympy.ceiling((effective_input_size - effective_kernel_shape[i]) / strides[i]) + strided_kernel_positions = sympy.ceiling( + (effective_input_size - effective_kernel_shape[i]) / strides[i]) else: strided_kernel_positions = (effective_input_size - effective_kernel_shape[i]) // strides[i] sympy_shape[-rank + i] = strided_kernel_positions + 1 @@ -491,7 +524,7 @@ class SymbolicShapeInference: def _check_merged_dims(self, dims, allow_broadcast=True): if allow_broadcast: - dims = [d for d in dims if not(is_literal(d) and int(d) <= 1)] + dims = [d for d in dims if not (is_literal(d) and int(d) <= 1)] if not all([d == dims[0] for d in dims]): self._add_suggested_merge(dims, apply=True) @@ -527,20 +560,33 @@ class SymbolicShapeInference: data_shape = self._get_shape(node, 0) indices_shape = self._get_shape(node, 1) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - self.known_vi_[node.input[0]].type.tensor_type.elem_type, - data_shape[:-1] + indices_shape)) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + data_shape[:-1] + indices_shape)) def _infer_symbolic_compute_ops(self, node): - funcs = {'Add' : lambda l: l[0] + l[1], - 'Div' : lambda l: l[0] // l[1], # integer div in sympy - 'Equal' : lambda l : l[0] == l[1], - 'Floor' : lambda l : sympy.floor(l[0]), - 'Max' : lambda l: l[1] if is_literal(l[0]) and int(l[0]) < -self.int_max_ else (l[0] if is_literal(l[1]) and int(l[1]) < -self.int_max_ else sympy.Max(l[0], l[1])), - 'Min' : lambda l: l[1] if is_literal(l[0]) and int(l[0]) > self.int_max_ else (l[0] if is_literal(l[1]) and int(l[1]) > self.int_max_ else sympy.Min(l[0], l[1])), - 'Mul' : lambda l: l[0] * l[1], - 'Sub' : lambda l: l[0] - l[1], - 'Where' : lambda l: l[1] if l[0] else l[2]} + funcs = { + 'Add': + lambda l: l[0] + l[1], + 'Div': + lambda l: l[0] // l[1], # integer div in sympy + 'Equal': + lambda l: l[0] == l[1], + 'Floor': + lambda l: sympy.floor(l[0]), + 'Max': + lambda l: l[1] if is_literal(l[0]) and int(l[0]) < -self.int_max_ else + (l[0] if is_literal(l[1]) and int(l[1]) < -self.int_max_ else sympy.Max(l[0], l[1])), + 'Min': + lambda l: l[1] if is_literal(l[0]) and int(l[0]) > self.int_max_ else + (l[0] if is_literal(l[1]) and int(l[1]) > self.int_max_ else sympy.Min(l[0], l[1])), + 'Mul': + lambda l: l[0] * l[1], + 'Sub': + lambda l: l[0] - l[1], + 'Where': + lambda l: l[1] if l[0] else l[2] + } assert node.op_type in funcs self._compute_on_sympy_data(node, funcs[node.op_type]) @@ -554,9 +600,7 @@ class SymbolicShapeInference: else: output_type = onnx.TensorProto.STRING vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - output_type, - self._get_shape(node, 0))) + vi.CopyFrom(helper.make_tensor_value_info(node.output[0], output_type, self._get_shape(node, 0))) def _infer_Compress(self, node): input_shape = self._get_shape(node, 0) @@ -570,7 +614,9 @@ class SymbolicShapeInference: output_shape = input_shape output_shape[handle_negative_axis(axis, len(input_shape))] = compress_len vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, output_shape)) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + output_shape)) def _infer_Concat(self, node): if any([i in self.sympy_data_ for i in node.input]): @@ -605,7 +651,9 @@ class SymbolicShapeInference: else: sympy_shape[d] = merged vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, get_shape_from_sympy_shape(sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + get_shape_from_sympy_shape(sympy_shape))) def _infer_Constant(self, node): t = get_attribute(node, 'value') @@ -620,21 +668,25 @@ class SymbolicShapeInference: self._update_computed_dims(sympy_shape) # update sympy data if output type is int, and shape is known if vi.type.tensor_type.elem_type == onnx.TensorProto.INT64 and all([is_literal(x) for x in sympy_shape]): - self.sympy_data_[node.output[0]] = np.ones([int(x) for x in sympy_shape], dtype=np.int64) * numpy_helper.to_array(get_attribute(node, 'value', 0)) + self.sympy_data_[node.output[0]] = np.ones( + [int(x) + for x in sympy_shape], dtype=np.int64) * numpy_helper.to_array(get_attribute(node, 'value', 0)) else: # create new dynamic shape # note input0 is a 1D vector of shape, the new symbolic shape has the rank of the shape vector length - sympy_shape = self._new_symbolic_shape(self._get_shape(node,0)[0], node) + sympy_shape = self._new_symbolic_shape(self._get_shape(node, 0)[0], node) - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - vi.type.tensor_type.elem_type, - get_shape_from_sympy_shape(sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], vi.type.tensor_type.elem_type, + get_shape_from_sympy_shape(sympy_shape))) def _infer_Conv(self, node): sympy_shape = self._compute_conv_pool_shape(node) self._update_computed_dims(sympy_shape) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], vi.type.tensor_type.elem_type, get_shape_from_sympy_shape(sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], vi.type.tensor_type.elem_type, + get_shape_from_sympy_shape(sympy_shape))) def _infer_Expand(self, node): expand_to_shape = self._try_get_value(node, 1) @@ -644,17 +696,19 @@ class SymbolicShapeInference: shape = self._get_shape(node, 0) new_shape = self._broadcast_shapes(shape, get_shape_from_sympy_shape(expand_to_shape)) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, new_shape)) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + new_shape)) def _infer_Gather(self, node): data_shape = self._get_shape(node, 0) axis = handle_negative_axis(get_attribute(node, 'axis', 0), len(data_shape)) indices_shape = self._get_shape(node, 1) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - vi.type.tensor_type.elem_type, - data_shape[:axis] + indices_shape + data_shape[axis+1:])) - # for 1D input, do some sympy compute + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], vi.type.tensor_type.elem_type, + data_shape[:axis] + indices_shape + data_shape[axis + 1:])) + # for 1D input, do some sympy compute if node.input[0] in self.sympy_data_ and len(data_shape) == 1 and 0 == get_attribute(node, 'axis', 0): idx = self._get_value(node, 1) data = self.sympy_data_[node.input[0]] @@ -670,9 +724,9 @@ class SymbolicShapeInference: def _infer_GatherElements(self, node): indices_shape = self._get_shape(node, 1) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - self.known_vi_[node.input[0]].type.tensor_type.elem_type, - indices_shape)) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + indices_shape)) def _infer_GatherND(self, node): data_shape = self._get_shape(node, 0) @@ -683,9 +737,9 @@ class SymbolicShapeInference: assert is_literal(last_index_dimension) and last_index_dimension <= data_rank new_shape = indices_shape[:-1] + data_shape[last_index_dimension:] vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - self.known_vi_[node.input[0]].type.tensor_type.elem_type, - new_shape)) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + new_shape)) def _infer_If(self, node): # special case for constant condition, in case there are mismatching shape from the non-executed branch @@ -705,7 +759,10 @@ class SymbolicShapeInference: vi.CopyFrom(subgraph.output[i_out]) vi.name = node.output[i_out] else: - assert all([d1 == d2 for d1,d2 in zip(vi.type.tensor_type.shape.dim, subgraph.output[i_out].type.tensor_type.shape.dim)]) + assert all([ + d1 == d2 for d1, d2 in zip(vi.type.tensor_type.shape.dim, + subgraph.output[i_out].type.tensor_type.shape.dim) + ]) # pass on sympy data from subgraph, if cond is constant if cond is not None and i_sub == (0 if cond > 0 else 1): if subgraph.output[i_out].name in subgraph_infer.sympy_data_: @@ -724,7 +781,7 @@ class SymbolicShapeInference: num_loop_carried = len(node.input) - 2 for i in range(len(node.output)): vi = self.known_vi_[node.output[i]] - vi.CopyFrom(subgraph.output[i + 1]) # first subgraph output is condition, not in node output + vi.CopyFrom(subgraph.output[i + 1]) # first subgraph output is condition, not in node output if i >= num_loop_carried: subgraph_vi_dim = subgraph.output[i + 1].type.tensor_type.shape.dim vi.type.tensor_type.shape.ClearField('dim') @@ -760,7 +817,9 @@ class SymbolicShapeInference: sympy_shape[:axis] + [self._new_symbolic_dim_from_output(node) if not is_literal(depth) else depth] + sympy_shape[axis:]) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[2]].type.tensor_type.elem_type, new_shape)) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[2]].type.tensor_type.elem_type, + new_shape)) def _infer_Pad(self, node): if get_opset(self.out_mp_) <= 10: @@ -774,14 +833,17 @@ class SymbolicShapeInference: sympy_shape = self._get_sympy_shape(node, 0) rank = len(sympy_shape) if pads is not None: - assert len(pads) == 2*rank - new_sympy_shape = [d + pad_up + pad_down for d, pad_up, pad_down in zip(sympy_shape, pads[:rank], pads[rank:])] + assert len(pads) == 2 * rank + new_sympy_shape = [ + d + pad_up + pad_down for d, pad_up, pad_down in zip(sympy_shape, pads[:rank], pads[rank:]) + ] self._update_computed_dims(new_sympy_shape) else: # dynamic pads, create new symbolic dimensions new_sympy_shape = self._new_symbolic_shape(rank, node) output_tp = self.known_vi_[node.input[0]].type.tensor_type.elem_type - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], output_tp, get_shape_from_sympy_shape(new_sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], output_tp, get_shape_from_sympy_shape(new_sympy_shape))) def _infer_Pool(self, node): sympy_shape = self._compute_conv_pool_shape(node) @@ -790,7 +852,9 @@ class SymbolicShapeInference: if not o: continue vi = self.known_vi_[o] - vi.CopyFrom(helper.make_tensor_value_info(o, vi.type.tensor_type.elem_type, get_shape_from_sympy_shape(sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(o, vi.type.tensor_type.elem_type, + get_shape_from_sympy_shape(sympy_shape))) def _infer_Range(self, node): vi = self.known_vi_[node.output[0]] @@ -799,12 +863,14 @@ class SymbolicShapeInference: start = as_scalar(input_data[0]) limit = as_scalar(input_data[1]) delta = as_scalar(input_data[2]) - new_sympy_shape = [sympy.Max(sympy.ceiling((limit - start)/delta), 0)] + new_sympy_shape = [sympy.Max(sympy.ceiling((limit - start) / delta), 0)] else: new_dim = self._new_symbolic_dim_from_output(node) new_sympy_shape = [self.symbolic_dims_[new_dim]] self._update_computed_dims(new_sympy_shape) - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, get_shape_from_sympy_shape(new_sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + get_shape_from_sympy_shape(new_sympy_shape))) def _infer_ReduceProd(self, node): axes = get_attribute(node, 'axes') @@ -822,9 +888,9 @@ class SymbolicShapeInference: assert len(shape_shape) == 1 shape_rank = shape_shape[0] assert is_literal(shape_rank) - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - vi.type.tensor_type.elem_type, - get_shape_from_sympy_shape(self._new_symbolic_shape(shape_rank, node)))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], vi.type.tensor_type.elem_type, + get_shape_from_sympy_shape(self._new_symbolic_shape(shape_rank, node)))) else: input_shape = self._get_shape(node, 0) input_sympy_shape = self._get_sympy_shape(node, 0) @@ -853,9 +919,9 @@ class SymbolicShapeInference: new_sympy_shape[deferred_dim_idx] = new_dim self._update_computed_dims(new_sympy_shape) - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - vi.type.tensor_type.elem_type, - get_shape_from_sympy_shape(new_sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], vi.type.tensor_type.elem_type, + get_shape_from_sympy_shape(new_sympy_shape))) self._pass_on_sympy_data(node) @@ -865,11 +931,12 @@ class SymbolicShapeInference: if get_opset(self.out_mp_) <= 10: scales = self._try_get_value(node, 1) if scales is not None: - new_sympy_shape = [sympy.simplify(sympy.floor(d*s)) for d,s in zip(input_sympy_shape, scales)] + new_sympy_shape = [sympy.simplify(sympy.floor(d * s)) for d, s in zip(input_sympy_shape, scales)] self._update_computed_dims(new_sympy_shape) - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - self.known_vi_[node.input[0]].type.tensor_type.elem_type, - get_shape_from_sympy_shape(new_sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], + self.known_vi_[node.input[0]].type.tensor_type.elem_type, + get_shape_from_sympy_shape(new_sympy_shape))) else: roi = self._try_get_value(node, 1) scales = self._try_get_value(node, 2) @@ -880,28 +947,34 @@ class SymbolicShapeInference: elif scales is not None: rank = len(scales) if get_attribute(node, 'coordinate_transformation_mode') == 'tf_crop_and_resize': - assert len(roi) == 2*rank + assert len(roi) == 2 * rank roi_start = list(roi)[:rank] roi_end = list(roi)[rank:] else: - roi_start = [0]*rank - roi_end = [1]*rank + roi_start = [0] * rank + roi_end = [1] * rank scales = list(scales) - new_sympy_shape = [sympy.simplify(sympy.floor(d * (end - start) * scale)) for d, start, end, scale in zip(input_sympy_shape, roi_start, roi_end, scales)] + new_sympy_shape = [ + sympy.simplify(sympy.floor(d * (end - start) * scale)) + for d, start, end, scale in zip(input_sympy_shape, roi_start, roi_end, scales) + ] self._update_computed_dims(new_sympy_shape) else: new_sympy_shape = self._new_symbolic_shape(self._get_shape_rank(node, 0), node) - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - self.known_vi_[node.input[0]].type.tensor_type.elem_type, - get_shape_from_sympy_shape(new_sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + get_shape_from_sympy_shape(new_sympy_shape))) def _infer_Scan(self, node): subgraph = get_attribute(node, 'body') num_scan_inputs = get_attribute(node, 'num_scan_inputs') - scan_input_axes = get_attribute(node, 'scan_input_axes', [0]*num_scan_inputs) + scan_input_axes = get_attribute(node, 'scan_input_axes', [0] * num_scan_inputs) num_scan_states = len(node.input) - num_scan_inputs - scan_input_axes = [handle_negative_axis(ax, self._get_shape_rank(node, i + num_scan_states)) for i, ax in enumerate(scan_input_axes)] + scan_input_axes = [ + handle_negative_axis(ax, self._get_shape_rank(node, i + num_scan_states)) + for i, ax in enumerate(scan_input_axes) + ] # We may have cases where the subgraph has optionial inputs that appear in both subgraph's input and initializer, # but not in the node's input. In such cases, the input model might be invalid, but let's skip those optional inputs. assert len(subgraph.input) >= len(node.input) @@ -915,7 +988,7 @@ class SymbolicShapeInference: si.name = subgraph_name self._onnx_infer_subgraph(node, subgraph) num_scan_outputs = len(node.output) - num_scan_states - scan_output_axes = get_attribute(node, 'scan_output_axes', [0]*num_scan_outputs) + scan_output_axes = get_attribute(node, 'scan_output_axes', [0] * num_scan_outputs) scan_input_dim = get_shape_from_type_proto(self.known_vi_[node.input[-1]].type)[scan_input_axes[-1]] for i, o in enumerate(node.output): vi = self.known_vi_[o] @@ -931,9 +1004,9 @@ class SymbolicShapeInference: def _infer_ScatterElements(self, node): data_shape = self._get_shape(node, 0) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - self.known_vi_[node.input[0]].type.tensor_type.elem_type, - data_shape)) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + data_shape)) def _infer_Shape(self, node): self.sympy_data_[node.output[0]] = self._get_sympy_shape(node, 0) @@ -941,7 +1014,8 @@ class SymbolicShapeInference: def _infer_Size(self, node): sympy_shape = self._get_sympy_shape(node, 0) self.sympy_data_[node.output[0]] = sympy_reduce_product(sympy_shape) - self.known_vi_[node.output[0]].CopyFrom(helper.make_tensor_value_info(node.output[0], onnx.TensorProto.INT64, [])) + self.known_vi_[node.output[0]].CopyFrom( + helper.make_tensor_value_info(node.output[0], onnx.TensorProto.INT64, [])) def _infer_Slice(self, node): if get_opset(self.out_mp_) <= 9: @@ -950,7 +1024,7 @@ class SymbolicShapeInference: ends = get_attribute(node, 'ends') if not axes: axes = list(range(len(starts))) - steps = [1]*len(axes) + steps = [1] * len(axes) else: starts = as_list(self._try_get_value(node, 1), keep_none=True) ends = as_list(self._try_get_value(node, 2), keep_none=True) @@ -959,7 +1033,7 @@ class SymbolicShapeInference: if axes is None and not (starts is None and ends is None): axes = list(range(0, len(starts if starts is not None else ends))) if steps is None and not (starts is None and ends is None): - steps = [1]*len(starts if starts is not None else ends) + steps = [1] * len(starts if starts is not None else ends) axes = as_list(axes, keep_none=True) steps = as_list(steps, keep_none=True) @@ -967,13 +1041,13 @@ class SymbolicShapeInference: if starts is None or ends is None: if axes is None: for i in range(len(new_sympy_shape)): - new_sympy_shape[i] = self._new_symbolic_dim_from_output(node,0,i) + new_sympy_shape[i] = self._new_symbolic_dim_from_output(node, 0, i) else: new_sympy_shape = get_shape_from_sympy_shape(new_sympy_shape) for i in axes: - new_sympy_shape[i] = self._new_symbolic_dim_from_output(node,0,i) + new_sympy_shape[i] = self._new_symbolic_dim_from_output(node, 0, i) else: - for i,s,e,t in zip(axes, starts, ends, steps): + for i, s, e, t in zip(axes, starts, ends, steps): if is_literal(e): if e >= self.int_max_: e = new_sympy_shape[i] @@ -985,7 +1059,8 @@ class SymbolicShapeInference: e = min(e, new_sympy_shape[i]) else: if e > 0: - e = sympy.Min(e, new_sympy_shape[i]) if e > 1 else e #special case for slicing first to make computation easier + e = sympy.Min(e, new_sympy_shape[i] + ) if e > 1 else e #special case for slicing first to make computation easier else: e = new_sympy_shape[i] + e else: @@ -1009,14 +1084,15 @@ class SymbolicShapeInference: self._update_computed_dims(new_sympy_shape) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - vi.type.tensor_type.elem_type, - get_shape_from_sympy_shape(new_sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], vi.type.tensor_type.elem_type, + get_shape_from_sympy_shape(new_sympy_shape))) # handle sympy_data if needed, for slice in shape computation if node.input[0] in self.sympy_data_ and [0] == axes and len(starts) == 1 and len(ends) == 1: input_sympy_data = self.sympy_data_[node.input[0]] - if type(input_sympy_data) == list or (type(input_sympy_data) == np.array and len(input_sympy_data.shape) == 1): + if type(input_sympy_data) == list or (type(input_sympy_data) == np.array + and len(input_sympy_data.shape) == 1): self.sympy_data_[node.output[0]] = input_sympy_data[starts[0]:ends[0]] def _infer_SoftmaxCrossEntropyLoss(self, node): @@ -1035,16 +1111,17 @@ class SymbolicShapeInference: split = get_attribute(node, 'split') if not split: num_outputs = len(node.output) - split = [input_sympy_shape[axis]/sympy.Integer(num_outputs)]*num_outputs + split = [input_sympy_shape[axis] / sympy.Integer(num_outputs)] * num_outputs self._update_computed_dims(split) else: split = [sympy.Integer(s) for s in split] for i_o in range(len(split)): vi = self.known_vi_[node.output[i_o]] - vi.CopyFrom(make_value_info_func(node.output[i_o], - self.known_vi_[node.input[0]].type.tensor_type.elem_type, - get_shape_from_sympy_shape(input_sympy_shape[:axis] + [split[i_o]] + input_sympy_shape[axis+1:]))) + vi.CopyFrom( + make_value_info_func( + node.output[i_o], self.known_vi_[node.input[0]].type.tensor_type.elem_type, + get_shape_from_sympy_shape(input_sympy_shape[:axis] + [split[i_o]] + input_sympy_shape[axis + 1:]))) self.known_vi_[vi.name] = vi def _infer_Split(self, node): @@ -1060,14 +1137,14 @@ class SymbolicShapeInference: repeats_value = self._get_value(node, 1) input_sympy_shape = self._get_sympy_shape(node, 0) new_sympy_shape = [] - for i,d in enumerate(input_sympy_shape): + for i, d in enumerate(input_sympy_shape): new_dim = d * repeats_value[i] new_sympy_shape.append(new_dim) self._update_computed_dims(new_sympy_shape) vi = self.known_vi_[node.output[0]] - vi.CopyFrom(helper.make_tensor_value_info(node.output[0], - vi.type.tensor_type.elem_type, - get_shape_from_sympy_shape(new_sympy_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(node.output[0], vi.type.tensor_type.elem_type, + get_shape_from_sympy_shape(new_sympy_shape))) def _infer_TopK(self, node): rank = self._get_shape_rank(node, 0) @@ -1089,7 +1166,9 @@ class SymbolicShapeInference: else: new_sympy_shape = self._get_sympy_shape(node, 0) new_sympy_shape[axis] = k - self._update_computed_dims(new_sympy_shape) # note that TopK dim could be computed in sympy_data, so need to update computed_dims when it enters shape + self._update_computed_dims( + new_sympy_shape + ) # note that TopK dim could be computed in sympy_data, so need to update computed_dims when it enters shape new_shape = get_shape_from_sympy_shape(new_sympy_shape) for i_o in range(len(node.output)): @@ -1114,6 +1193,39 @@ class SymbolicShapeInference: vi = self.known_vi_[node.output[0]] vi.CopyFrom(new_vi) + def _infer_Attention(self, node): + #TODO: shape inference for the other output (present). + shape = self._get_shape(node, 0) + shape_bias = self._get_shape(node, 2) + shape[2] = shape_bias[0] / 3 + output_dtype = self.known_vi_[node.input[0]].type.tensor_type.elem_type + vi = self.known_vi_[node.output[0]] + vi.CopyFrom(helper.make_tensor_value_info(node.output[0], output_dtype, shape)) + + def _infer_BiasGelu(self, node): + self._propagate_shape_and_type(node) + + def _infer_FastGelu(self, node): + self._propagate_shape_and_type(node) + + def _infer_Gelu(self, node): + self._propagate_shape_and_type(node) + + def _infer_LayerNormalization(self, node): + self._propagate_shape_and_type(node) + + def _infer_LongformerAttention(self, node): + self._propagate_shape_and_type(node) + + def _infer_SkipLayerNormalization(self, node): + self._propagate_shape_and_type(node) + + def _propagate_shape_and_type(self, node, input_index=0, output_index=0): + shape = self._get_shape(node, input_index) + output_dtype = self.known_vi_[node.input[input_index]].type.tensor_type.elem_type + vi = self.known_vi_[node.output[output_index]] + vi.CopyFrom(helper.make_tensor_value_info(node.output[output_index], output_dtype, shape)) + def _infer_impl(self, start_sympy_data=None): self.sympy_data_ = start_sympy_data or {} self.out_mp_.graph.ClearField('value_info') @@ -1152,10 +1264,11 @@ class SymbolicShapeInference: while not all([o.name in sorted_known_vi for o in self.out_mp_.graph.output]): old_sorted_nodes_len = len(sorted_nodes) for node in self.out_mp_.graph.node: - if (node.output[0] not in sorted_known_vi ) and all([i in sorted_known_vi for i in node.input if i]): + if (node.output[0] not in sorted_known_vi) and all([i in sorted_known_vi for i in node.input if i]): sorted_known_vi.update(node.output) sorted_nodes.append(node) - if old_sorted_nodes_len == len(sorted_nodes) and not all([o.name in sorted_known_vi for o in self.out_mp_.graph.output]): + if old_sorted_nodes_len == len(sorted_nodes) and not all( + [o.name in sorted_known_vi for o in self.out_mp_.graph.output]): raise Exception('Invalid model with cyclic graph') for node in sorted_nodes: @@ -1178,7 +1291,9 @@ class SymbolicShapeInference: # onnx automatically merge dims with value, i.e. Mul(['aaa', 'bbb'], [1000, 1]) -> [1000, 'bbb'] # symbolic shape inference needs to apply merge of 'aaa' -> 1000 in this case - if node.op_type in ['Add', 'Sub', 'Mul', 'Div', 'MatMul', 'MatMulInteger', 'MatMulInteger16', 'Where', 'Sum']: + if node.op_type in [ + 'Add', 'Sub', 'Mul', 'Div', 'MatMul', 'MatMulInteger', 'MatMulInteger16', 'Where', 'Sum' + ]: vi = self.known_vi_[node.output[0]] out_rank = len(get_shape_from_type_proto(vi.type)) in_shapes = [self._get_shape(node, i) for i in range(len(node.input))] @@ -1203,7 +1318,10 @@ class SymbolicShapeInference: if None in out_shape or out_type_undefined: if self.auto_merge_: - if node.op_type in ['Add', 'Sub', 'Mul', 'Div', 'MatMul', 'MatMulInteger', 'MatMulInteger16', 'Concat', 'Where', 'Sum']: + if node.op_type in [ + 'Add', 'Sub', 'Mul', 'Div', 'MatMul', 'MatMulInteger', 'MatMulInteger16', 'Concat', + 'Where', 'Sum' + ]: shapes = [self._get_shape(node, i) for i in range(len(node.input))] if node.op_type in ['MatMul', 'MatMulInteger', 'MatMulInteger16']: if None in out_shape: @@ -1226,7 +1344,10 @@ class SymbolicShapeInference: # if a tensor has a lower rank (dim_idx[idx] < 0), it would automatically broadcast and need no merge dim_idx = [len(s) - len(out_shape) + idx for s in shapes] if len(dim_idx) > 0: - self._add_suggested_merge([s[i] if is_literal(s[i]) else str(s[i]) for s, i in zip(shapes, dim_idx) if i >= 0]) + self._add_suggested_merge([ + s[i] if is_literal(s[i]) else str(s[i]) for s, i in zip(shapes, dim_idx) + if i >= 0 + ]) self.run_ = True else: self.run_ = False @@ -1252,18 +1373,20 @@ class SymbolicShapeInference: else: # otherwise, use original data type out_dtype = vi.type.tensor_type.elem_type - vi.CopyFrom(helper.make_tensor_value_info(vi.name, - out_dtype, - get_shape_from_sympy_shape(new_shape))) + vi.CopyFrom( + helper.make_tensor_value_info(vi.name, out_dtype, + get_shape_from_sympy_shape(new_shape))) if self.verbose_ > 0: if is_unknown_op: - print("Possible unknown op: {} node: {}, guessing {} shape".format(node.op_type, node.name, vi.name)) + print("Possible unknown op: {} node: {}, guessing {} shape".format( + node.op_type, node.name, vi.name)) if self.verbose_ > 2: - print(' {}: {} {}'.format(node.output[i_o], str(new_shape), vi.type.tensor_type.elem_type)) + print(' {}: {} {}'.format(node.output[i_o], str(new_shape), + vi.type.tensor_type.elem_type)) self.run_ = True - continue # continue the inference after guess, no need to stop as no merge is needed + continue # continue the inference after guess, no need to stop as no merge is needed if self.verbose_ > 0 or not self.auto_merge_ or out_type_undefined: print('Stopping at incomplete shape inference at ' + node.op_type + ': ' + node.name) @@ -1301,15 +1424,29 @@ class SymbolicShapeInference: raise Exception("Incomplete symbolic shape inference") return symbolic_shape_inference.out_mp_ + def parse_arguments(): - parser = argparse.ArgumentParser() - parser.add_argument('--input', required=True, help='The input model file') - parser.add_argument('--output', help='The output model file') - parser.add_argument('--auto_merge', help='Automatically merge symbolic dims when confliction happens', action='store_true', default=False) - parser.add_argument('--int_max', help='maximum value for integer to be treated as boundless for ops like slice', type=int, default=2**31 - 1) - parser.add_argument('--guess_output_rank', help='guess output rank to be the same as input 0 for unknown ops', action='store_true', default=False) - parser.add_argument('--verbose', help='Prints detailed logs of inference, 0: turn off, 1: warnings, 3: detailed', type=int, default=0) - return parser.parse_args() + parser = argparse.ArgumentParser() + parser.add_argument('--input', required=True, help='The input model file') + parser.add_argument('--output', help='The output model file') + parser.add_argument('--auto_merge', + help='Automatically merge symbolic dims when confliction happens', + action='store_true', + default=False) + parser.add_argument('--int_max', + help='maximum value for integer to be treated as boundless for ops like slice', + type=int, + default=2**31 - 1) + parser.add_argument('--guess_output_rank', + help='guess output rank to be the same as input 0 for unknown ops', + action='store_true', + default=False) + parser.add_argument('--verbose', + help='Prints detailed logs of inference, 0: turn off, 1: warnings, 3: detailed', + type=int, + default=0) + return parser.parse_args() + if __name__ == '__main__': args = parse_arguments() @@ -1317,7 +1454,8 @@ if __name__ == '__main__': if args.output: print('output model ' + args.output) print('Doing symbolic shape inference...') - out_mp = SymbolicShapeInference.infer_shapes(onnx.load(args.input), args.int_max, args.auto_merge, args.guess_output_rank, args.verbose) + out_mp = SymbolicShapeInference.infer_shapes(onnx.load(args.input), args.int_max, args.auto_merge, + args.guess_output_rank, args.verbose) if args.output and out_mp: onnx.save(out_mp, args.output) print('Done!') diff --git a/onnxruntime/python/tools/transformers/fusion_skiplayernorm.py b/onnxruntime/python/tools/transformers/fusion_skiplayernorm.py index 9bdc482326..7292dc7b3d 100644 --- a/onnxruntime/python/tools/transformers/fusion_skiplayernorm.py +++ b/onnxruntime/python/tools/transformers/fusion_skiplayernorm.py @@ -18,6 +18,7 @@ class FusionSkipLayerNormalization(Fusion): """ def __init__(self, model: OnnxModel): super().__init__(model, "SkipLayerNormalization", "LayerNormalization") + self.shape_infer_helper = self.model.infer_runtime_shape({"batch_size": 4, "seq_len": 7}) def fuse(self, node, input_name_to_nodes, output_name_to_node): add = self.model.get_parent(node, 0, output_name_to_node) @@ -35,6 +36,14 @@ class FusionSkipLayerNormalization(Fusion): if len(self.model.get_parents(add)) != 2: return + if self.shape_infer_helper is not None: + if not self.shape_infer_helper.compare_shape(add.input[0], add.input[1]): + return + else: + logger.warning( + "symbolic shape infer failed. it's safe to ignore this message if there is no issue with optimized model" + ) + gather_path = self.model.match_parent_path(add, ['Gather'], [None]) if gather_path is not None and self.model.find_graph_input(gather_path[0].input[1]) is None: if self.model.match_parent_path(gather_path[0], ['ConstantOfShape'], [1]) is None: diff --git a/onnxruntime/python/tools/transformers/onnx_model.py b/onnxruntime/python/tools/transformers/onnx_model.py index 5f821feea1..ef8626a70e 100644 --- a/onnxruntime/python/tools/transformers/onnx_model.py +++ b/onnxruntime/python/tools/transformers/onnx_model.py @@ -12,6 +12,7 @@ from pathlib import Path import numpy as np from collections import deque from onnx import ModelProto, TensorProto, numpy_helper, helper, external_data_helper, save_model +from shape_infer_helper import SymbolicShapeInferenceHelper logger = logging.getLogger(__name__) @@ -20,6 +21,21 @@ class OnnxModel: def __init__(self, model): self.model = model self.node_name_counter = {} + self.shape_infer_helper = None + + def infer_runtime_shape(self, dynamic_axis_mapping, update = False): + shape_infer_helper = None + if update: + shape_infer_helper = SymbolicShapeInferenceHelper(self.model) + self.shape_infer_helper = shape_infer_helper + else: + if self.shape_infer_helper is None: + self.shape_infer_helper = SymbolicShapeInferenceHelper(self.model) + shape_infer_helper = self.shape_infer_helper + + if shape_infer_helper.infer(dynamic_axis_mapping): + return shape_infer_helper + return None def input_name_to_nodes(self): input_name_to_nodes = {} diff --git a/onnxruntime/python/tools/transformers/shape_infer_helper.py b/onnxruntime/python/tools/transformers/shape_infer_helper.py new file mode 100644 index 0000000000..1aaa7b1206 --- /dev/null +++ b/onnxruntime/python/tools/transformers/shape_infer_helper.py @@ -0,0 +1,84 @@ +#------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +#-------------------------------------------------------------------------- + +import os +import sys + +# In ORT Package the symbolic_shape_infer.py is in ../tools +file_path = os.path.dirname(__file__) +if os.path.exists(os.path.join(file_path, "../tools/symbolic_shape_infer.py")): + sys.path.append(os.path.join(file_path, '../tools')) +else: + sys.path.append(os.path.join(file_path, '..')) +from symbolic_shape_infer import * + + +class SymbolicShapeInferenceHelper(SymbolicShapeInference): + def __init__(self, model, verbose=0, int_max=2**31 - 1, auto_merge=True, guess_output_rank=False): + super().__init__(int_max, auto_merge, guess_output_rank, verbose) + self.model_ = onnx.ModelProto() + self.model_.CopyFrom(model) + self.all_shapes_inferred_ = False + self.inferred_ = False + + # The goal is to remove dynamic_axis_mapping + def infer(self, dynamic_axis_mapping): + if self.inferred_: + return self.all_shapes_inferred_ + + self.dynamic_axis_mapping_ = dynamic_axis_mapping # e.g {"batch_size" : 4, "seq_len" :7} + + self._preprocess(self.model_) + while self.run_: + self.all_shapes_inferred_ = self._infer_impl() + + self.inferred_ = True + return self.all_shapes_inferred_ + + # override _preprocess() to avoid unnecessary model copy since ctor copies the model + def _preprocess(self, in_mp): + self.out_mp_ = in_mp + self.initializers_ = dict([(i.name, i) for i in self.out_mp_.graph.initializer]) + self.known_vi_ = dict([(i.name, i) for i in list(self.out_mp_.graph.input)]) + self.known_vi_.update( + dict([(i.name, helper.make_tensor_value_info(i.name, i.data_type, list(i.dims))) + for i in self.out_mp_.graph.initializer])) + + # Override _get_sympy_shape() in symbolic_shape_infer.py to ensure shape inference by giving the actual value of dynamic axis + def _get_sympy_shape(self, node, idx): + sympy_shape = [] + for d in self._get_shape(node, idx): + if type(d) == str: + if d in self.dynamic_axis_mapping_.keys(): + sympy_shape.append(self.dynamic_axis_mapping_[d]) + elif d in self.symbolic_dims_: + sympy_shape.append(self.symbolic_dims_[d]) + else: + sympy_shape.append(sympy.Symbol(d, integer=True)) + else: + assert None != d + sympy_shape.append(d) + return sympy_shape + + def get_edge_shape(self, edge): + assert (self.all_shapes_inferred_ == True) + if edge not in self.known_vi_: + print("Cannot retrive the shape of " + str(edge)) + return None + type_proto = self.known_vi_[edge].type + shape = get_shape_from_type_proto(type_proto) + for i in range(len(shape)): + d = shape[i] + if type(d) == str and d in self.dynamic_axis_mapping_.keys(): + shape[i] = self.dynamic_axis_mapping_[d] + return shape + + def compare_shape(self, edge, edge_other): + assert (self.all_shapes_inferred_ == True) + shape = self.get_edge_shape(edge) + shape_other = self.get_edge_shape(edge_other) + if shape is None or shape_other is None: + raise Exception("At least one shape is missed for edges to compare") + return shape == shape_other diff --git a/onnxruntime/python/tools/transformers/test/test_shape_infer_helper.py b/onnxruntime/python/tools/transformers/test/test_shape_infer_helper.py new file mode 100644 index 0000000000..825b25b404 --- /dev/null +++ b/onnxruntime/python/tools/transformers/test/test_shape_infer_helper.py @@ -0,0 +1,48 @@ +import os +import unittest +import sys +sys.path.append(os.path.join(os.path.dirname(__file__), '..')) + +from onnx_exporter import export_onnx_model_from_pt +from huggingface_models import MODELS +from benchmark_helper import Precision +from shape_infer_helper import * + + +class SymbolicShapeInferenceHelperTest(unittest.TestCase): + def _load_onnx(self, model_name): + input_names = MODELS[model_name][0] + base_path = "../onnx_models/" + import torch + with torch.no_grad(): + export_onnx_model_from_pt(model_name, MODELS[model_name][1], MODELS[model_name][2], MODELS[model_name][3], + None, '../cache_models', base_path, input_names[:1], False, Precision.FLOAT32, + True, True, True, False, {}) + model_path = base_path + model_name.replace('-', '_') + "_1.onnx" + import onnx + return onnx.load_model(model_path) + + def test_bert_shape_infer_helper(self): + model = self._load_onnx("bert-base-cased") + shape_infer_helper = SymbolicShapeInferenceHelper(model) + self.assertEqual(shape_infer_helper.infer({"batch_size": 4, "seq_len": 16}), True) + self.assertEqual(shape_infer_helper.get_edge_shape("802"), [4, 16, 768]) + self.assertEqual(shape_infer_helper.get_edge_shape("804"), [4, 16, 1]) + self.assertEqual(shape_infer_helper.get_edge_shape("1748"), []) + self.assertEqual(shape_infer_helper.get_edge_shape("encoder.layer.4.attention.output.LayerNorm.weight"), [768]) + self.assertEqual(shape_infer_helper.get_edge_shape("1749"), [768, 3072]) + self.assertEqual(shape_infer_helper.get_edge_shape("817"), [4, 16, 3072]) + self.assertEqual(shape_infer_helper.get_edge_shape("encoder.layer.4.intermediate.dense.bias"), [3072]) + self.assertEqual(shape_infer_helper.get_edge_shape("1750"), [3072, 768]) + self.assertEqual(shape_infer_helper.get_edge_shape("853"), [3]) + self.assertEqual(shape_infer_helper.get_edge_shape("858"), [1]) + self.assertEqual(shape_infer_helper.get_edge_shape("880"), [4, 16, 12, 64]) + + self.assertEqual(shape_infer_helper.compare_shape("329", "253"), True) + self.assertEqual(shape_infer_helper.compare_shape("447", "371"), True) + self.assertEqual(shape_infer_helper.compare_shape("329", "817"), False) + self.assertEqual(shape_infer_helper.compare_shape("447", "853"), False) + + +if __name__ == '__main__': + unittest.main()