xhinker move te as needed

This commit is contained in:
Vladimir Mandic
2024-08-29 10:42:38 -04:00
parent 4f606d38f8
commit 0b901e08ce
4 changed files with 28 additions and 10 deletions
+20
View File
@@ -459,6 +459,18 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl
except Exception:
pass
prompt_embed = positive_pooled = negative_embed = negative_pooled = None
te1_device, te2_device, te3_device = None, None, None
if hasattr(pipe, "text_encoder") and pipe.text_encoder.device != devices.device:
te1_device = pipe.text_encoder.device
pipe.text_encoder = pipe.text_encoder.to(devices.device)
if hasattr(pipe, "text_encoder_2") and pipe.text_encoder_2.device != devices.device:
te2_device = pipe.text_encoder_2.device
pipe.text_encoder_2 = pipe.text_encoder_2.to(devices.device)
if hasattr(pipe, "text_encoder_3") and pipe.text_encoder_3.device != devices.device:
te3_device = pipe.text_encoder_3.device
pipe.text_encoder_3 = pipe.text_encoder_3.to(devices.device)
if SD3:
prompt_embed, negative_embed, positive_pooled, negative_pooled = get_weighted_text_embeddings_sd3(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, use_t5_encoder=bool(pipe.text_encoder_3))
elif 'Flux' in pipe.__class__.__name__:
@@ -467,4 +479,12 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl
prompt_embed, negative_embed, positive_pooled, negative_pooled = get_weighted_text_embeddings_sdxl_2p(pipe=pipe, prompt=prompt, prompt_2=prompt_2, neg_prompt=neg_prompt, neg_prompt_2=neg_prompt_2)
else:
prompt_embed, negative_embed = get_weighted_text_embeddings_sd15(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, clip_skip=clip_skip)
if te1_device is not None:
pipe.text_encoder = pipe.text_encoder.to(te1_device)
if te2_device is not None:
pipe.text_encoder_2 = pipe.text_encoder_2.to(te2_device)
if te3_device is not None:
pipe.text_encoder_3 = pipe.text_encoder_3.to(te3_device)
return prompt_embed, positive_pooled, negative_embed, negative_pooled