HiDream fix prompt parser

This commit is contained in:
Disty0
2025-04-23 03:14:58 +03:00
parent 16935ec08b
commit 677bdd8611
2 changed files with 15 additions and 7 deletions
+12 -6
View File
@@ -161,7 +161,12 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t
if 'prompt' in possible:
if 'OmniGen' in model.__class__.__name__:
prompts = [p.replace('|image|', '<|image_1|>') for p in prompts]
if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
if 'HiDreamImage' in model.__class__.__name__:
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
prompt_embeds = prompt_parser_diffusers.embedder('prompt_embeds')
args['prompt_embeds_t5'] = prompt_embeds[0]
args['prompt_embeds_llama3'] = prompt_embeds[1]
elif hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
args['prompt_embeds'] = prompt_parser_diffusers.embedder('prompt_embeds')
if prompt_parser_diffusers.embedder is not None:
if 'StableCascade' in model.__class__.__name__:
@@ -172,12 +177,15 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
elif 'Flux' in model.__class__.__name__:
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
elif 'HiDreamImage' in model.__class__.__name__:
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
else:
args['prompt'] = prompts
if 'negative_prompt' in possible:
if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
if 'HiDreamImage' in model.__class__.__name__:
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
negative_prompt_embeds = prompt_parser_diffusers.embedder('negative_prompt_embeds')
args['negative_prompt_embeds_t5'] = negative_prompt_embeds[0]
args['negative_prompt_embeds_llama3'] = negative_prompt_embeds[1]
elif hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
args['negative_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_prompt_embeds')
if prompt_parser_diffusers.embedder is not None:
if 'StableCascade' in model.__class__.__name__:
@@ -186,8 +194,6 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
elif 'StableDiffusion3' in model.__class__.__name__:
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
elif 'HiDreamImage' in model.__class__.__name__:
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
else:
if 'PixArtSigmaPipeline' in model.__class__.__name__: # pixart-sigma pipeline throws list-of-list for negative prompt
args['negative_prompt'] = negative_prompts[0]
+3 -1
View File
@@ -511,11 +511,13 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
return prompt_embeds, pooled_prompt_embeds, None, None # no negative support
if "HiDreamImage" in pipe.__class__.__name__: # clip is only used for the pooled embeds
prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds = pipe.encode_prompt(
prompt_embeds_t5, negative_prompt_embeds_t5, prompt_embeds_llama3, negative_prompt_embeds_llama3, pooled_prompt_embeds, negative_pooled_prompt_embeds = pipe.encode_prompt(
prompt=prompt, prompt_2=prompt_2, prompt_3=prompt_3, prompt_4=prompt_4,
negative_prompt=neg_prompt, negative_prompt_2=neg_prompt_2, negative_prompt_3=neg_prompt_3, negative_prompt_4=neg_prompt_4,
device=device, num_images_per_prompt=1,
)
prompt_embeds = [prompt_embeds_t5, prompt_embeds_llama3]
negative_prompt_embeds = [negative_prompt_embeds_t5, negative_prompt_embeds_llama3]
return prompt_embeds, pooled_prompt_embeds, negative_prompt_embeds, negative_pooled_prompt_embeds
if prompt != prompt_2: