From 9c1c8feeb8fd0ae535c75eabebfc3c9c97e9d261 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 22 Jul 2024 23:02:30 +0300 Subject: [PATCH] NNCF fix AuraFlow --- modules/sd_hijack.py | 5 +++-- modules/sd_models_compile.py | 10 ++++++---- 2 files changed, 9 insertions(+), 6 deletions(-) diff --git a/modules/sd_hijack.py b/modules/sd_hijack.py index 6894c8ff9..dc07dea81 100644 --- a/modules/sd_hijack.py +++ b/modules/sd_hijack.py @@ -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 diff --git a/modules/sd_models_compile.py b/modules/sd_models_compile.py index fd7f189f6..4d4776c26 100644 --- a/modules/sd_models_compile.py +++ b/modules/sd_models_compile.py @@ -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'):