diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index a31c88124..8927b4341 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -19,14 +19,6 @@ def ipex_autocast(*args, **kwargs): else: return original_autocast(*args, **kwargs) -#Diffusers BF16: -original_linear_forward = torch.nn.modules.Linear.forward -def linear_forward(self, input): - if input.dtype != self.weight.data.dtype: - return original_linear_forward(self, input.to(self.weight.data.dtype)) - else: - return original_linear_forward(self, input) - #Embedding BF16 original_torch_cat = torch.cat def torch_cat(input, *args, **kwargs): @@ -35,14 +27,6 @@ def torch_cat(input, *args, **kwargs): else: return original_torch_cat(input, *args, **kwargs) -original_conv2d = torch.nn.functional.conv2d -#Diffusers BF16: -def conv2d(input, weight, *args, **kwargs): - if input.dtype != weight.data.dtype: - return original_conv2d(input.to(weight.data.dtype), weight, *args, **kwargs) - else: - return original_conv2d(input, weight, *args, **kwargs) - original_interpolate = torch.nn.functional.interpolate #Latent antialias: def interpolate(input, size=None, scale_factor=None, mode='nearest', align_corners=None, recompute_scale_factor=None, antialias=False): @@ -101,7 +85,7 @@ def ipex_hijacks(): lambda orig_func, input, normalized_shape=None, weight=None, *args, **kwargs: orig_func(input.to(weight.data.dtype), normalized_shape, weight, *args, **kwargs), lambda orig_func, input, normalized_shape=None, weight=None, *args, **kwargs: - input.dtype != weight.data.dtype and weight is not None) + weight is not None and input.dtype != weight.data.dtype) #Functions that does not work with the XPU: #UniPC: @@ -132,7 +116,5 @@ def ipex_hijacks(): #Functions that make compile mad with CondFunc: torch.autocast = ipex_autocast - torch.nn.modules.Linear.forward = linear_forward torch.cat = torch_cat - torch.nn.functional.conv2d = conv2d torch.nn.functional.interpolate = interpolate