jumbo update, see changelog

This commit is contained in:
Vladimir Mandic
2024-09-18 13:48:30 -04:00
parent 67aeaac1ce
commit 2acb883dda
33 changed files with 247 additions and 141 deletions
+45 -25
View File
@@ -165,38 +165,58 @@ def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, c
return
else:
t0 = time.time()
positive_schedule, scheduled = get_prompt_schedule(prompts[0], steps)
negative_schedule, neg_scheduled = get_prompt_schedule(negative_prompts[0], steps)
p.scheduled_prompt = scheduled or neg_scheduled
p.prompt_embeds = []
p.positive_pooleds = []
p.negative_embeds = []
p.negative_pooleds = []
if shared.opts.diffusers_offload_mode == "balanced":
pipe = sd_models.apply_balanced_offload(pipe)
elif hasattr(pipe, "maybe_free_model_hooks"):
# if the last job is interrupted, model will stay in the vram and cause oom, send everything back to cpu before continuing
pipe.maybe_free_model_hooks()
devices.torch_gc()
for i in range(max(len(positive_schedule), len(negative_schedule))):
positive_prompt = positive_schedule[i % len(positive_schedule)]
negative_prompt = negative_schedule[i % len(negative_schedule)]
if shared.opts.prompt_attention == "xhinker parser":
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_xhinker_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip)
else:
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip)
if prompt_embed is not None:
p.prompt_embeds.append(torch.cat([prompt_embed] * len(prompts), dim=0))
if negative_embed is not None:
p.negative_embeds.append(torch.cat([negative_embed] * len(negative_prompts), dim=0))
if positive_pooled is not None:
p.positive_pooleds.append(torch.cat([positive_pooled] * len(prompts), dim=0))
if negative_pooled is not None:
p.negative_pooleds.append(torch.cat([negative_pooled] * len(negative_prompts), dim=0))
prompt_embeds, positive_pooleds, negative_embeds, negative_pooleds = [], [], [], []
last_prompt, last_negative = None, None
for prompt, negative in zip(prompts, negative_prompts):
prompt_embed, positive_pooled, negative_embed, negative_pooled = None, None, None, None
if last_prompt == prompt and last_negative == negative or False:
prompt_embeds.append(prompt_embed)
positive_pooleds.append(positive_pooled)
negative_embeds.append(negative_embed)
negative_pooleds.append(negative_pooled)
continue
positive_schedule, scheduled = get_prompt_schedule(prompt, steps)
negative_schedule, neg_scheduled = get_prompt_schedule(negative, steps)
p.scheduled_prompt = scheduled or neg_scheduled
p.prompt_embeds = []
p.positive_pooleds = []
p.negative_embeds = []
p.negative_pooleds = []
if shared.opts.sd_textencoder_cache:
for i in range(max(len(positive_schedule), len(negative_schedule))):
positive_prompt = positive_schedule[i % len(positive_schedule)]
negative_prompt = negative_schedule[i % len(negative_schedule)]
if shared.opts.prompt_attention == "xhinker parser" or 'Flux' in pipe.__class__.__name__:
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_xhinker_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip)
else:
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip)
prompt_embeds.append(prompt_embed)
positive_pooleds.append(positive_pooled)
negative_embeds.append(negative_embed)
negative_pooleds.append(negative_pooled)
last_prompt, last_negative = prompt, negative
def fix_length(embeds):
max_len = max([p.shape[1] for p in embeds])
for i, p in enumerate(embeds):
if p.shape[1] < max_len:
expanded = torch.zeros((p.shape[0], max_len, p.shape[2]), device=p.device, dtype=p.dtype)
expanded[:, :p.shape[1], :] = p
embeds[i] = expanded
return torch.cat(embeds, dim=0)
p.prompt_embeds.append(fix_length(prompt_embeds))
p.negative_embeds.append(fix_length(negative_embeds))
p.positive_pooleds.append(fix_length(positive_pooleds))
p.negative_pooleds.append(fix_length(negative_pooleds))
if shared.opts.sd_textencoder_cache and p.batch_size == 1:
cache.update({
'prompt_embeds': p.prompt_embeds,
'negative_embeds': p.negative_embeds,