diff --git a/modules/intel/ipex/diffusers.py b/modules/intel/ipex/diffusers.py index 6e56866c1..e03592197 100644 --- a/modules/intel/ipex/diffusers.py +++ b/modules/intel/ipex/diffusers.py @@ -1,7 +1,7 @@ import torch import intel_extension_for_pytorch as ipex import torch.nn.functional as F -import diffusers #1.19.3 +import diffusers #0.20.2 Attention = diffusers.models.attention_processor.Attention diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index cb008e099..9e1d9713e 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -90,6 +90,10 @@ def ipex_hijacks(): CondFunc('torch.nn.modules.GroupNorm.forward', lambda orig_func, self, input: orig_func(self, input.to(self.weight.data.dtype)), lambda orig_func, self, input: input.dtype != self.weight.data.dtype) + #Hypernetwork training: + CondFunc('torch.nn.modules.linear.Linear.forward', + lambda orig_func, self, input: orig_func(self, input.to(self.weight.data.dtype)), + lambda orig_func, self, input: input.dtype != self.weight.data.dtype) #Embedding FP32: CondFunc('torch.bmm', lambda orig_func, input, mat2, *args, **kwargs: orig_func(input, mat2.to(input.dtype), *args, **kwargs),