diff --git a/test/external/external_test_onnx_backend.py b/test/external/external_test_onnx_backend.py index 618388ce46..d48154b85d 100644 --- a/test/external/external_test_onnx_backend.py +++ b/test/external/external_test_onnx_backend.py @@ -134,7 +134,6 @@ backend_test.exclude('test_simple_rnn_*') # no control flow # control flow uses AttributeProto.GRAPH -backend_test.exclude('test_if_*') backend_test.exclude('test_loop*') backend_test.exclude('test_range_float_type_positive_delta_expanded_cpu') # requires loop backend_test.exclude('test_affine_grid_2d_align_corners_expanded_cpu') @@ -183,6 +182,8 @@ backend_test.exclude('test_resize_downsample_scales_cubic_antialias_cpu') # anti backend_test.exclude('test_resize_downsample_sizes_cubic_antialias_cpu') # antialias not implemented backend_test.exclude('test_ai_onnx_ml_label_encoder_tensor_value_only_mapping_cpu') # bad data type string backend_test.exclude('test_ai_onnx_ml_label_encoder_tensor_mapping_cpu') # bad data type string +backend_test.exclude('test_if_opt_cpu') # ValueError: 13 is not a valid AttributeType +backend_test.exclude('test_if_seq_cpu') # NotImplementedError: op='SequenceConstruct' is not supported backend_test.exclude('test_scatternd_min_cpu') # min not yet supported backend_test.exclude('test_scatternd_max_cpu') # max not yet supported diff --git a/test/external/external_test_onnx_ops.py b/test/external/external_test_onnx_ops.py index 8ed0ddc03e..e4be34fa5e 100644 --- a/test/external/external_test_onnx_ops.py +++ b/test/external/external_test_onnx_ops.py @@ -100,6 +100,25 @@ class TestMainOnnxOps(TestOnnxOps): self._test_resize_scales([0.01, 0.25, 0.5, 0.51, 0.6, 1.0, 1.5, 2.0, 3.5, 20.0], mode="cubic", exclude_outside=1) self._test_resize_scales([0.01, 0.25, 0.5, 0.51, 0.6, 1.0, 1.5, 2.0, 3.5, 20.0], mode="cubic", exclude_outside=0) + def _test_if(self, then_value, else_value): + then_out = onnx.helper.make_tensor_value_info("res", onnx.TensorProto.FLOAT, then_value.shape) + else_out = onnx.helper.make_tensor_value_info("res", onnx.TensorProto.FLOAT, else_value.shape) + + then_const_node = onnx.helper.make_node("Constant", inputs=[], outputs=["res"], value=onnx.numpy_helper.from_array(then_value)) + else_const_node = onnx.helper.make_node("Constant", inputs=[], outputs=["res"], value=onnx.numpy_helper.from_array(else_value)) + + then_body = onnx.helper.make_graph([then_const_node], "then_body", [], [then_out]) + else_body = onnx.helper.make_graph([else_const_node], "else_body", [], [else_out]) + + self.helper_test_single_op("If", {"cond": np.array(False).astype(bool)}, {"then_branch": then_body, "else_branch": else_body}, ["res"]) + self.helper_test_single_op("If", {"cond": np.array(True).astype(bool)}, {"then_branch": then_body, "else_branch": else_body}, ["res"]) + + def test_if_different_shapes_broadcastable(self): + self._test_if(np.array([[1], [2]]).astype(np.float32), np.array([[6, 5, 4, 3, 2, 1]]).astype(np.float32)) + + def test_if_different_shapes_not_broadcastable(self): + self._test_if(np.array([[1, 2, 3], [4, 5, 6]]).astype(np.float32), np.array([[6, 5, 4, 3, 2, 1]]).astype(np.float32)) + def test_resize_downsample_scales_linear_align_corners(self): # https://github.com/onnx/onnx/blob/main/docs/Operators.md#examples-131 X = np.array([[[[1, 2, 3, 4], [5, 6, 7, 8]]]], dtype=np.float32) diff --git a/tinygrad/frontend/onnx.py b/tinygrad/frontend/onnx.py index 25b377d9b5..33d5408602 100644 --- a/tinygrad/frontend/onnx.py +++ b/tinygrad/frontend/onnx.py @@ -21,9 +21,9 @@ class AttributeType(enum.IntEnum): ONNX attribute type identifiers. Reference: https://github.com/onnx/onnx/blob/rel-1.18.0/onnx/onnx.proto3#L128-L145 """ - FLOAT = 1; INT = 2; STRING = 3; TENSOR = 4; FLOATS = 6; INTS = 7; STRINGS = 8 # noqa: E702 + FLOAT = 1; INT = 2; STRING = 3; TENSOR = 4; GRAPH = 5; FLOATS = 6; INTS = 7; STRINGS = 8 # noqa: E702 - def to_field_name(self) -> str: return {1: "f", 2: "i", 3: "s", 4: "t", 6: "floats", 7: "ints", 8: "strings"}[self.value] + def to_field_name(self) -> str: return {1: "f", 2: "i", 3: "s", 4: "t", 5: "g", 6: "floats", 7: "ints", 8: "strings"}[self.value] class OnnxDataType(enum.IntEnum): """ @@ -266,6 +266,7 @@ class OnnxPBParser: case 3: obj["i"] = self.reader.read_int64() case 4: obj["s"] = self.reader.read_bytes().data().tobytes().decode("utf8") case 5: obj["t"] = self._parse_TensorProto()['parsed_tensor'] + case 6: obj["g"] = OnnxRunner._from_subgraph(self._parse_GraphProto()) case 7: obj["floats"].append(self.reader.read_float()) case 8: obj["ints"].append(self.reader.read_int64()) case 9: obj["strings"].append(self.reader.read_bytes().data().tobytes().decode("utf8")) @@ -401,8 +402,11 @@ class OnnxRunner: """ def __init__(self, model_path: Tensor | str | pathlib.Path): model = OnnxPBParser(model_path, load_external_data=True).parse() - graph = model["graph"] + self._init_from_graph(model["graph"]) + + def _init_from_graph(self, graph: dict, is_subgraph: bool = False): self.is_training = any(n['parsed_node'].opset_id.domain in {Domain.AI_ONNX_TRAINING, Domain.AI_ONNX_PREVIEW_TRAINING} for n in graph["node"]) + self.graph_name = graph["name"] if is_subgraph else "" self.graph_values = {"": None, **{i["name"]: i["parsed_tensor"] for i in graph["initializer"]}} self.graph_inputs = {i["name"]: i["parsed_type"] for i in graph["input"] if i["name"] not in self.graph_values} self.graph_outputs = tuple(o["name"] for o in graph["output"]) @@ -414,6 +418,12 @@ class OnnxRunner: self.variable_dims: dict[str, int] = {} self.onnx_ops = onnx_ops + @classmethod + def _from_subgraph(cls, graph: dict) -> "OnnxRunner": + subgraph = cls.__new__(cls) + subgraph._init_from_graph(graph, is_subgraph=True) + return subgraph + def _parse_input(self, name: str, value: Any, spec: OnnxValue): if spec.is_optional and value is None: return None if spec.is_sequence: @@ -445,9 +455,10 @@ class OnnxRunner: return {name:Tensor.empty(*spec.shape, device=device, dtype=dtype or spec.dtype) for name, spec in self.graph_inputs.items()} def to(self, device:str|None): - self.graph_values = {k:v.to(device) if isinstance(v, Tensor) else v for k,v in self.graph_values.items()} + self.graph_values = {k: (v.to(device) if isinstance(v, Tensor) else v) for k,v in self.graph_values.items()} self.graph_nodes = tuple(OnnxNode(n.op, n.opset_id, tuple(n.inputs), tuple(n.outputs), - {k:v.to(device) if isinstance(v, Tensor) else v for k,v in n.opts.items()}) for n in self.graph_nodes) + {k: (v.to(device) if isinstance(v, (Tensor, OnnxRunner)) else v) for k,v in n.opts.items()}) + for n in self.graph_nodes) return self def __call__(self, inputs:dict[str, Any], debug=debug): @@ -461,9 +472,9 @@ class OnnxRunner: # provide additional opts if node.op == "Split" and 'num_outputs' not in opts: opts['num_outputs'] = len(node.outputs) - if node.op == "Gradient": opts['intermediate_tensors'] = self.graph_values + if node.op in {"Gradient", "If"}: opts['intermediate_tensors'] = self.graph_values - if debug >= 1: print(f"{num}: op '{node.op}' opt {opts}") + if debug >= 1: print((f"[{self.graph_name}] " if self.graph_name else "") + f"{num}: op '{node.op}' opt {opts}") if debug >= 2 and node.inputs: print("\tinputs:\n" + "\n".join(f"\t\t{x} - {i!r}" for x,i in zip(node.inputs, inps))) ret = self._select_op(node.op, node.opset_id)(*inps, **opts) ret = ret if isinstance(ret, tuple) else (ret,) @@ -543,6 +554,23 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT return __decorator # ***** Property/Graph Ops ***** + def If(condition:Tensor, else_branch:OnnxRunner, then_branch:OnnxRunner, intermediate_tensors:dict[str, Tensor]): + def run_branch(branch:OnnxRunner): + branch.graph_values.update(intermediate_tensors) + out = branch({k:intermediate_tensors[k] for k in branch.graph_inputs.keys()}) + # dereference intermediate tensors so Buffer can be deallocated + for k in intermediate_tensors: del branch.graph_values[k] + return out + # both branch must be ran before the condition can be evaluated + else_out, then_out = run_branch(else_branch), run_branch(then_branch) + assert len(else_out) == len(then_out), f"else_out and then_out must have the same number of outputs: {len(else_out)} != {len(then_out)}" + # can use where op when output shape is the same + if all(t.shape == e.shape for t,e in zip(then_out.values(), else_out.values())): + return tuple(condition.where(t,e) for t,e in zip(then_out.values(), else_out.values())) + # otherwise, use condition to select the output in python + cond = _resolve_const(_cached_to_python_const(condition)) + return tuple(t if cond else e for t,e in zip(then_out.values(), else_out.values())) + def Identity(x:Tensor): return x def Constant(sparse_value:Tensor|None=None, value:Tensor|None=None, value_float:float|None=None, value_floats:list[float]|None=None, value_int:int|None=None, value_ints:list[int]|None=None, value_string:str|None=None, value_strings:list[str]|None=None):