NNCF fix AuraFlow

This commit is contained in:
Disty0
2024-07-22 23:02:30 +03:00
parent 25c3c6107e
commit 9c1c8feeb8
2 changed files with 9 additions and 6 deletions
+3 -2
View File
@@ -284,20 +284,21 @@ class EmbeddingsWithFixes(torch.nn.Module):
class NNCF_T5DenseGatedActDense(torch.nn.Module): # forward can't find what self is without creating a class
def __init__(self, T5DenseGatedActDense):
def __init__(self, T5DenseGatedActDense, dtype):
super().__init__()
self.wi_0 = T5DenseGatedActDense.wi_0
self.wi_1 = T5DenseGatedActDense.wi_1
self.wo = T5DenseGatedActDense.wo
self.dropout = T5DenseGatedActDense.dropout
self.act = T5DenseGatedActDense.act
self.torch_dtype = dtype
def forward(self, hidden_states):
hidden_gelu = self.act(self.wi_0(hidden_states))
hidden_linear = self.wi_1(hidden_states)
hidden_states = hidden_gelu * hidden_linear
hidden_states = self.dropout(hidden_states)
hidden_states = hidden_states.to(torch.float32) # this line needs to be forced to fp32
hidden_states = hidden_states.to(self.torch_dtype) # this line needs to be forced
hidden_states = self.wo(hidden_states)
return hidden_states
+6 -4
View File
@@ -58,21 +58,23 @@ def apply_compile_to_model(sd_model, function, options, op=None):
sd_model.text_encoder = None
sd_model.text_encoder = sd_model.decoder_pipe.text_encoder = function(sd_model.decoder_pipe.text_encoder)
else:
if op == "nncf" and sd_model.text_encoder.__class__.__name__ == "T5EncoderModel":
if op == "nncf" and sd_model.text_encoder.__class__.__name__ in {"T5EncoderModel", "UMT5EncoderModel"}:
from modules.sd_hijack import NNCF_T5DenseGatedActDense # T5DenseGatedActDense uses fp32
for i in range(len(sd_model.text_encoder.encoder.block)):
sd_model.text_encoder.encoder.block[i].layer[1].DenseReluDense = NNCF_T5DenseGatedActDense(
sd_model.text_encoder.encoder.block[i].layer[1].DenseReluDense
sd_model.text_encoder.encoder.block[i].layer[1].DenseReluDense,
dtype=torch.float32 if devices.dtype != torch.bfloat16 else torch.bfloat16
)
sd_model.text_encoder = function(sd_model.text_encoder)
if hasattr(sd_model, 'text_encoder_2') and hasattr(sd_model.text_encoder_2, 'config'):
sd_model.text_encoder_2 = function(sd_model.text_encoder_2)
if hasattr(sd_model, 'text_encoder_3') and hasattr(sd_model.text_encoder_3, 'config'):
if op == "nncf" and sd_model.text_encoder_3.__class__.__name__ == "T5EncoderModel":
if op == "nncf" and sd_model.text_encoder_3.__class__.__name__ in {"T5EncoderModel", "UMT5EncoderModel"}:
from modules.sd_hijack import NNCF_T5DenseGatedActDense # T5DenseGatedActDense uses fp32
for i in range(len(sd_model.text_encoder_3.encoder.block)):
sd_model.text_encoder_3.encoder.block[i].layer[1].DenseReluDense = NNCF_T5DenseGatedActDense(
sd_model.text_encoder_3.encoder.block[i].layer[1].DenseReluDense
sd_model.text_encoder_3.encoder.block[i].layer[1].DenseReluDense,
dtype=torch.float32 if devices.dtype != torch.bfloat16 else torch.bfloat16
)
sd_model.text_encoder_3 = function(sd_model.text_encoder_3)
if hasattr(sd_model, 'prior_pipe') and hasattr(sd_model, 'prior_text_encoder'):