mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
NNCF fix AuraFlow
This commit is contained in:
@@ -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'):
|
||||
|
||||
Reference in New Issue
Block a user