mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Add Symbolic Shape and Type Infer for aten::group_norm (#13348)
Add symbolic shape and type infer for aten::group_norm.
This commit is contained in:
parent
2fa18ea77e
commit
9efa8e20bb
1 changed files with 21 additions and 0 deletions
|
|
@ -211,6 +211,7 @@ class SymbolicShapeInference:
|
|||
"avg_pool2d": self._infer_aten_pool2d,
|
||||
"_adaptive_avg_pool2d": self._infer_aten_pool2d,
|
||||
"numpy_T": self._infer_Transpose,
|
||||
"native_group_norm": self._infer_aten_group_norm,
|
||||
}
|
||||
self.run_ = True
|
||||
self.suggested_merge_ = {}
|
||||
|
|
@ -1351,6 +1352,26 @@ class SymbolicShapeInference:
|
|||
vi = self.known_vi_[node.output[0]]
|
||||
vi.CopyFrom(helper.make_tensor_value_info(node.output[0], onnx.TensorProto.INT64, new_shape))
|
||||
|
||||
def _infer_aten_group_norm(self, node):
|
||||
self._propagate_shape_and_type(node)
|
||||
input_shape = self._get_shape(node, 0)
|
||||
N = input_shape[0] if input_shape is not None and len(input_shape) != 0 else None
|
||||
group = self._try_get_value(node, 6)
|
||||
output_dtype = self.known_vi_[node.input[0]].type.tensor_type.elem_type
|
||||
for i in [1, 2]:
|
||||
if node.output[i]:
|
||||
vi = self.known_vi_[node.output[i]]
|
||||
vi.CopyFrom(
|
||||
helper.make_tensor_value_info(
|
||||
node.output[i],
|
||||
output_dtype,
|
||||
[
|
||||
N if N is not None else self._new_symbolic_dim_from_output(node, i, 0),
|
||||
as_scalar(group) if group is not None else self._new_symbolic_dim_from_output(node, i, 1),
|
||||
],
|
||||
)
|
||||
)
|
||||
|
||||
def _infer_BatchNormalization(self, node):
|
||||
self._propagate_shape_and_type(node)
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue