diff --git a/modules/sd_hijack.py b/modules/sd_hijack.py index b811f33bf..6894c8ff9 100644 --- a/modules/sd_hijack.py +++ b/modules/sd_hijack.py @@ -283,6 +283,25 @@ class EmbeddingsWithFixes(torch.nn.Module): return torch.stack(vecs) +class NNCF_T5DenseGatedActDense(torch.nn.Module): # forward can't find what self is without creating a class + def __init__(self, T5DenseGatedActDense): + 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 + + 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 = self.wo(hidden_states) + return hidden_states + + def add_circular_option_to_conv_2d(): conv2d_constructor = torch.nn.Conv2d.__init__ diff --git a/modules/sd_models_compile.py b/modules/sd_models_compile.py index c6006db25..f69b126b0 100644 --- a/modules/sd_models_compile.py +++ b/modules/sd_models_compile.py @@ -58,9 +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": + 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 = 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_2, 'config'): + if op == "nncf" and sd_model.text_encoder_3.__class__.__name__ == "T5EncoderModel": + 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 = function(sd_model.text_encoder_3) if hasattr(sd_model, 'prior_pipe') and hasattr(sd_model, 'prior_text_encoder'): sd_model.prior_text_encoder = None sd_model.prior_text_encoder = sd_model.prior_pipe.text_encoder = function(sd_model.prior_pipe.text_encoder)