From e016f9027f37253ba5d38663f2f5f371d994239e Mon Sep 17 00:00:00 2001 From: liujingfeng4A069 Date: Fri, 13 Sep 2024 16:03:03 +0800 Subject: [PATCH] add mock code for torch_dipu --- utils/process_test.py | 31 +++++++++++++++++++++++++++++++ 1 file changed, 31 insertions(+) diff --git a/utils/process_test.py b/utils/process_test.py index 3ba21f6..326a0e0 100644 --- a/utils/process_test.py +++ b/utils/process_test.py @@ -65,6 +65,37 @@ def version(*args, **kwargs): setattr(torch.backends.cudnn, 'version', version) patch('torch.backends.cudnn.version', return_value=90000).start() """ + elif device_torch == "torch_dipu": + mock_code = """ +from unittest.mock import patch +patch('torch.cuda.get_device_capability', return_value=(8, 0)).start() + +import torch_dipu +if not hasattr(torch._C, '_cuda_setStream'): + def _cuda_setStream(*args, **kwargs): + pass + setattr(torch._C, '_cuda_setStream', _cuda_setStream) +patch('torch._C._cuda_setStream', new=torch_dipu.set_stream).start() + +if not hasattr(torch._C, '_cuda_setDevice'): + def _cuda_setDevice(*args, **kwargs): + pass + setattr(torch._C, '_cuda_setDevice', _cuda_setDevice) +patch('torch._C._cuda_setDevice', new=torch_dipu.set_device).start() + +if not hasattr(torch.backends.cudnn, 'is_acceptable'): + def is_acceptable(*args, **kwargs): + pass + setattr(torch.backends.cudnn, 'is_acceptable', is_acceptable) +patch('torch.backends.cudnn.is_acceptable', return_value=True).start() + +if not hasattr(torch.backends.cudnn, 'version'): + def version(*args, **kwargs): + pass + setattr(torch.backends.cudnn, 'version', version) +patch('torch.backends.cudnn.version', return_value=90000).start() +""" + custom_test_code = r""" import os import json