From d0c3f92ec6bd3f370ce78abf330620781591bf9d Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Tue, 25 Apr 2023 04:56:42 -0700 Subject: [PATCH] [DORT] Fix fake tensor problem cuased by PyTorch change (#15664) This should make `Orttraining Linux Lazy Tensor CI Pipeline` green again. --- .../orttraining/python/training/torchdynamo/ort_backend.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/torchdynamo/ort_backend.py b/orttraining/orttraining/python/training/torchdynamo/ort_backend.py index 808fb1ca75..b8d7990e64 100644 --- a/orttraining/orttraining/python/training/torchdynamo/ort_backend.py +++ b/orttraining/orttraining/python/training/torchdynamo/ort_backend.py @@ -627,7 +627,9 @@ class OrtBackend: if graph_module in self._partitioner_cache: partitioned_prim_graph_module = self._partitioner_cache[graph_module] else: - prim_graph_module = make_fx(graph_module, decomposition_table=_ATEN2ATEN_DECOMP)(*args) + prim_graph_module = make_fx( + graph_module, tracing_mode="fake", _allow_non_fake_inputs=True, decomposition_table=_ATEN2ATEN_DECOMP + )(*args) # TODO(wechi): this is required for removing aten::_to_copy in _replace_to_copy_with_to. # We need input and output tensors' devices to decide if aten::_to_copy is just a Cast. FakeTensorProp(prim_graph_module).propagate(*args)