diff --git a/modules/intel/ipex/__init__.py b/modules/intel/ipex/__init__.py index 6d3b2ab1d..b0240cb7a 100644 --- a/modules/intel/ipex/__init__.py +++ b/modules/intel/ipex/__init__.py @@ -190,11 +190,13 @@ def ipex_init(): # pylint: disable=too-many-statements ipex._C._DeviceProperties.multi_processor_count = ipex._C._DeviceProperties.gpu_subslice_count ipex._C._DeviceProperties.major = 12 ipex._C._DeviceProperties.minor = 1 + ipex._C._DeviceProperties.L2_cache_size = 16*1024*1024 # A770 and A750 else: torch._C._cuda_getCurrentRawStream = torch._C._xpu_getCurrentRawStream torch._C._XpuDeviceProperties.multi_processor_count = torch._C._XpuDeviceProperties.gpu_subslice_count torch._C._XpuDeviceProperties.major = 12 torch._C._XpuDeviceProperties.minor = 1 + torch._C._XpuDeviceProperties.L2_cache_size = 16*1024*1024 # A770 and A750 # Fix functions with ipex: # torch.xpu.mem_get_info always returns the total memory as free memory @@ -211,6 +213,7 @@ def ipex_init(): # pylint: disable=too-many-statements torch.cuda.get_device_capability = lambda *args, **kwargs: (12,1) torch.cuda.get_device_properties.major = 12 torch.cuda.get_device_properties.minor = 1 + torch.cuda.get_device_properties.L2_cache_size = 16*1024*1024 # A770 and A750 torch.cuda.ipc_collect = lambda *args, **kwargs: None torch.cuda.utilization = lambda *args, **kwargs: 0 diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index 91a256aed..e47065a62 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -332,14 +332,6 @@ def torch_load(f, map_location=None, *args, **kwargs): else: return original_torch_load(f, *args, map_location=map_location, **kwargs) -original_torch_Generator = torch.Generator -@wraps(torch.Generator) -def torch_Generator(device=None): - if check_cuda(device): - return original_torch_Generator(return_xpu(device)) - else: - return original_torch_Generator(device) - @wraps(torch.cuda.synchronize) def torch_cuda_synchronize(device=None): if check_cuda(device): @@ -355,6 +347,17 @@ def torch_cuda_device(device): return torch.xpu.device(device) +# torch.Generator has to be a class for isinstance checks +original_torch_Generator = torch.Generator +class torch_Generator(original_torch_Generator): + def __new__(self, device=None): + # can't hijack __init__ because of C override so use return super().__new__ + if check_cuda(device): + return super().__new__(self, return_xpu(device)) + else: + return super().__new__(self, device) + + # Hijack Functions: def ipex_hijacks(): global device_supports_fp64, can_allocate_plus_4gb @@ -374,10 +377,12 @@ def ipex_hijacks(): torch.linspace = torch_linspace torch.eye = torch_eye torch.load = torch_load - torch.Generator = torch_Generator torch.cuda.synchronize = torch_cuda_synchronize torch.cuda.device = torch_cuda_device + torch.Generator = torch_Generator + torch._C.Generator = torch_Generator + torch.backends.cuda.sdp_kernel = return_null_context torch.nn.DataParallel = DummyDataParallel torch.UntypedStorage.is_cuda = is_cuda