diff --git a/onnxruntime/python/onnxruntime_pybind_state.cc b/onnxruntime/python/onnxruntime_pybind_state.cc index 2055d1d707..13cc12c12c 100644 --- a/onnxruntime/python/onnxruntime_pybind_state.cc +++ b/onnxruntime/python/onnxruntime_pybind_state.cc @@ -150,6 +150,8 @@ const char* GetDeviceName(const OrtDevice& device) { return CUDA; case OrtDevice::FPGA: return "FPGA"; + case OrtDevice::NPU: + return "NPU"; default: ORT_THROW("Unknown device type: ", device.Type()); } @@ -1061,6 +1063,8 @@ void addObjectMethods(py::module& m, Environment& env, ExecutionProviderRegistra .def("device_type", &OrtDevice::Type, R"pbdoc(Device Type.)pbdoc") .def_static("cpu", []() { return OrtDevice::CPU; }) .def_static("cuda", []() { return OrtDevice::GPU; }) + .def_static("fpga", []() { return OrtDevice::FPGA; }) + .def_static("npu", []() { return OrtDevice::NPU; }) .def_static("default_memory", []() { return OrtDevice::MemType::DEFAULT; }); py::class_ ort_arena_cfg_binding(m, "OrtArenaCfg"); diff --git a/orttraining/orttraining/python/training/torchdynamo/ort_backend.py b/orttraining/orttraining/python/training/torchdynamo/ort_backend.py index 0319fecf69..bfc3144317 100644 --- a/orttraining/orttraining/python/training/torchdynamo/ort_backend.py +++ b/orttraining/orttraining/python/training/torchdynamo/ort_backend.py @@ -24,6 +24,7 @@ from torch.fx.passes.infra.partitioner import CapabilityBasedPartitioner from torch.fx.passes.operator_support import OperatorSupport from torch.fx.passes.tools_common import CALLABLE_NODE_OPS from torch.onnx._globals import GLOBALS as ONNX_GLOBALS +from torch._subclasses.fake_tensor import FakeTensor import onnxruntime # type: ignore from onnxruntime.capi import _pybind_state as ORTC @@ -59,13 +60,14 @@ def _nvtx_range_pop(): torch.cuda.nvtx.range_pop() -def _get_ort_device_type(device_type: str, device_index: int): +def _get_ort_device_type(device_type: str): if device_type == "cuda": return ORTC.OrtDevice.cuda() # type: ignore if device_type == "cpu": return ORTC.OrtDevice.cpu() # type: ignore + # ort pytorch device is mapped to NPU OrtDevice type if device_type == "ort": - return ORTC.get_ort_device(device_index).device_type() # type: ignore + return ORTC.OrtDevice.npu() raise ValueError("Unsupported device type: " + device_type) @@ -357,7 +359,7 @@ def _get_onnx_devices(values: Tuple[torch.Tensor, ...]) -> Tuple[ORTC.OrtDevice, devices: Tuple[ORTC.OrtDevice, ...] = tuple( # type: ignore ORTC.OrtDevice( # type: ignore - _get_ort_device_type(value.device.type, _device_id_or_zero(value.device.index)), + _get_ort_device_type(value.device.type), ORTC.OrtDevice.default_memory(), # type: ignore _device_id_or_zero(value.device.index), ) @@ -366,6 +368,30 @@ def _get_onnx_devices(values: Tuple[torch.Tensor, ...]) -> Tuple[ORTC.OrtDevice, return devices +def _get_ortvalues_from_torch_tensors( + tensors: Tuple[torch.Tensor, ...], devices: Tuple[ORTC.OrtDevice, ...] +) -> Tuple[torch.Tensor, ...]: + ortvalues = ORTC.OrtValueVector() # type: ignore + ortvalues.reserve(len(tensors)) + dtypes = [] + shapes = [] + data_ptrs = [] + + for tensor in tensors: + dtypes.append(_NP_DTYPE[tensor.dtype]) + shapes.append(tensor.size()) + data_ptrs.append(tensor.data_ptr()) + ortvalues.push_back_batch(tensors, data_ptrs, dtypes, shapes, devices) + return ortvalues + + +def _to_real_tensor(tensor: FakeTensor) -> torch.Tensor: + if tensor.is_sparse: + raise ValueError("sparse tensor is not yet supported.") + out = torch.empty(tensor.size(), dtype=tensor.dtype, device=tensor.device) + return out + + def _run_onnx_session_with_ortvaluevector( sess: onnxruntime.InferenceSession, input_names: Tuple[str, ...], @@ -374,24 +400,24 @@ def _run_onnx_session_with_ortvaluevector( output_names: Tuple[str, ...], outputs: Tuple[torch.Tensor, ...], output_devices: Tuple[ORTC.OrtDevice, ...], # type: ignore + preallocate_output: bool, ) -> Tuple[torch.Tensor, ...]: _nvtx_range_push("contiguous") inputs = tuple(a.contiguous() for a in inputs) _nvtx_range_pop() _nvtx_range_push("push_back_batch") - ort_inputs = ORTC.OrtValueVector() # type: ignore - ort_inputs.reserve(len(inputs)) - ort_outputs = ORTC.OrtValueVector() # type: ignore - dtypes = [] - shapes = [] - data_ptrs = [] - for value in inputs: - dtypes.append(_NP_DTYPE[value.dtype]) - shapes.append(value.size()) - data_ptrs.append(value.data_ptr()) - ort_inputs.push_back_batch(inputs, data_ptrs, dtypes, shapes, input_devices) + ort_inputs = _get_ortvalues_from_torch_tensors(inputs, input_devices) + + # preallocate output pytorch Tensors and use the buffers affined to the torch device for the output ortvalue. + # Because the output ortvalue is not allocated and owned by ort, it does not need to convert the output ortvalue + # to torch Tensor transferring the ownership. + if preallocate_output: + pth_outputs = tuple(map(lambda t: _to_real_tensor(t) if isinstance(t, FakeTensor) else t, outputs)) + ort_outputs = _get_ortvalues_from_torch_tensors(pth_outputs, output_devices) + else: + ort_outputs = ORTC.OrtValueVector() # type: ignore _nvtx_range_pop() _nvtx_range_push("run_with_ortvaluevector") @@ -400,10 +426,13 @@ def _run_onnx_session_with_ortvaluevector( sess.run_with_ortvaluevector(run_options, input_names, ort_inputs, output_names, ort_outputs, output_devices) _nvtx_range_pop() - _nvtx_range_push("after run_with_ortvaluevector") - pth_outputs = onnxruntime.training.ortmodule._utils._ortvalues_to_torch_tensor(ort_outputs) # type: ignore - _nvtx_range_pop() - return pth_outputs + if preallocate_output: + return pth_outputs + else: + _nvtx_range_push("after run_with_ortvaluevector") + pth_outputs = onnxruntime.training.ortmodule._utils._ortvalues_to_torch_tensor(ort_outputs) # type: ignore + _nvtx_range_pop() + return pth_outputs def _assert_allclose_with_detailed_error_message( @@ -457,7 +486,7 @@ class OrtBackend: 3. Inside _ort_accelerated_call, it creates onnxruntime.InferenceSession and calls it to execute the sub-graph. """ - def __init__(self, ep: str = ""): + def __init__(self, ep: str = "", preallocate_output: bool = False): self._supported_ops = OrtOperatorSupport() # TODO: this is a naive implementation of cache without proper guard self._partitioner_cache: Dict[torch.fx.GraphModule, torch.fx.GraphModule] = {} @@ -466,6 +495,16 @@ class OrtBackend: self.ep = ep + # preallocate_output allows for allocating output torch Tensor buffers and feeding them to InferenceSession + # in order to avoid internal allocation of output buffers in InferenceSession. + # If output ortvalue returned from InferenceSession is allocated internally, + # it needs to be converted to torch Tensor for return, and the torch Tensor should hold the ownership. + # When a custom torch device is used with a custom aten allocator, the conversion from ortvalue to torch Tensor + # should be supported, which is currently done through dlpack. Note that dlpack might not support a custom torch device. + # It can be avoided by allowing for preallocation for output buffers allocated by a custom aten allocator, + # and use the preallocated output buffers for InferenceSession not holding any ownership for them. + self.preallocate_output = preallocate_output + def _ort_acclerated_call(self, graph_module: torch.fx.GraphModule, *args, **kwargs): if graph_module in self._ort_execution_info.sessions: # We have seen this graph before, so we can use cached objects including session. @@ -486,7 +525,16 @@ class OrtBackend: # # WARNING: The downstream code should not change prim_outputs and # this backend should always produces output with schema identical to prim_outputs'. - prim_outputs = FakeTensorProp(graph_module).propagate(*args, **kwargs) + try: + prim_outputs = FakeTensorProp(graph_module).propagate(*args, **kwargs) + except Exception: + logger.info(f"FakeTensorProb failed for {graph_module}") + # When FakeTensorProp fails, it is not possible to preallocate output buffers + # because the output shapes are not inferred. + self.preallocate_output = False + + # rethrow FakeTensorProb failure because it is not yet currently handled. + raise self._ort_execution_info.example_outputs[graph_module] = prim_outputs # Compile the torch.fx.GraphModule into a torch.jit.ScriptModule. script_module = _fx_to_torchscript(graph_module) @@ -503,7 +551,7 @@ class OrtBackend: # Initialize a ORT session to execute this ONNX model. # TorchDynamo assumes all inputs/outputs are on the same device, # so we add execution provider only based on the first input's device. - ep = self.ep if self.ep else _infer_ep_from_device(args[0].device) + ep = self.ep or _infer_ep_from_device(args[0].device) onnx_session = _create_onnx_session(onnx_proto, ep) # Cache ORT session. It's reused for the same "graph_module". @@ -533,7 +581,14 @@ class OrtBackend: # ORT output is ok. _nvtx_range_push("run_onnx_session_with_ortvaluevector") onnx_outputs = _run_onnx_session_with_ortvaluevector( - onnx_session, input_names, args, input_devices, output_names, prim_outputs, output_devices + onnx_session, + input_names, + args, + input_devices, + output_names, + prim_outputs, + output_devices, + self.preallocate_output, ) _nvtx_range_pop() if self._ort_execution_info.assert_allclose_to_baseline: @@ -549,7 +604,14 @@ class OrtBackend: # ORT output's first element must be extracted and returned. Otherwise, type # mismatch may happen in downstream computation. onnx_outputs = _run_onnx_session_with_ortvaluevector( - onnx_session, input_names, args, input_devices, output_names, (prim_outputs,), output_devices + onnx_session, + input_names, + args, + input_devices, + output_names, + (prim_outputs,), + output_devices, + self.preallocate_output, ) assert len(onnx_outputs) == 1 if self._ort_execution_info.assert_allclose_to_baseline: