IPEX fix hypernetwork training

This commit is contained in:
Disty0
2023-09-05 19:32:03 +03:00
parent b4af0c2241
commit 9058bfa250
2 changed files with 5 additions and 1 deletions
+1 -1
View File
@@ -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
+4
View File
@@ -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),