From 757bc667204b1b9a2d5993c2a914aaf643eaa463 Mon Sep 17 00:00:00 2001 From: baijumeswani Date: Tue, 19 Oct 2021 08:10:52 -0700 Subject: [PATCH] Set cuda version to be None instead of an empty string (#9435) --- orttraining/orttraining/python/training/ortmodule/__init__.py | 4 ++-- .../python/training/ortmodule/torch_cpp_extensions/install.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/orttraining/orttraining/python/training/ortmodule/__init__.py b/orttraining/orttraining/python/training/ortmodule/__init__.py index 97b39ba252..b87bee3bb5 100644 --- a/orttraining/orttraining/python/training/ortmodule/__init__.py +++ b/orttraining/orttraining/python/training/ortmodule/__init__.py @@ -30,8 +30,8 @@ ORTMODULE_FALLBACK_POLICY = _FallbackPolicy.FALLBACK_UNSUPPORTED_DEVICE |\ ORTMODULE_FALLBACK_RETRY = False ORTMODULE_IS_DETERMINISTIC = torch.are_deterministic_algorithms_enabled() -ONNXRUNTIME_CUDA_VERSION = ort_info.cuda_version if hasattr(ort_info, 'cuda_version') else '' -ONNXRUNTIME_ROCM_VERSION = ort_info.rocm_version if hasattr(ort_info, 'rocm_version') else '' +ONNXRUNTIME_CUDA_VERSION = ort_info.cuda_version if hasattr(ort_info, 'cuda_version') else None +ONNXRUNTIME_ROCM_VERSION = ort_info.rocm_version if hasattr(ort_info, 'rocm_version') else None # Verify minimum PyTorch version is installed before proceding to ONNX Runtime initialization try: diff --git a/orttraining/orttraining/python/training/ortmodule/torch_cpp_extensions/install.py b/orttraining/orttraining/python/training/ortmodule/torch_cpp_extensions/install.py index 21282bcce2..1ff667a777 100644 --- a/orttraining/orttraining/python/training/ortmodule/torch_cpp_extensions/install.py +++ b/orttraining/orttraining/python/training/ortmodule/torch_cpp_extensions/install.py @@ -48,8 +48,8 @@ def build_torch_cpp_extensions(): os.chdir(ORTMODULE_TORCH_CPP_DIR) # Extensions might leverage CUDA/ROCM versions internally - os.environ["ONNXRUNTIME_CUDA_VERSION"] = ONNXRUNTIME_CUDA_VERSION - os.environ["ONNXRUNTIME_ROCM_VERSION"] = ONNXRUNTIME_ROCM_VERSION + os.environ["ONNXRUNTIME_CUDA_VERSION"] = ONNXRUNTIME_CUDA_VERSION if not ONNXRUNTIME_CUDA_VERSION is None else '' + os.environ["ONNXRUNTIME_ROCM_VERSION"] = ONNXRUNTIME_ROCM_VERSION if not ONNXRUNTIME_ROCM_VERSION is None else '' ############################################################################ # Pytorch CPP Extensions that DO require CUDA/ROCM