Add some cache cleaning

This commit is contained in:
kijai
2024-07-06 23:00:09 +03:00
parent d28ad2f6b9
commit 51b89de9fc
+56 -70
View File
@@ -2,6 +2,7 @@ import torch
import os
import random
import re
import gc
import comfy.model_management as mm
from comfy.utils import ProgressBar, load_torch_file
@@ -86,7 +87,7 @@ class DownloadAndLoadKolorsModel:
force_zeros_for_empty_prompt=False
)
pipeline = pipeline.to(device)
#pipeline = pipeline.to(device)
pipeline.enable_model_cpu_offload()
kolors_model = {
@@ -142,43 +143,31 @@ class KolorsTextEncode:
batch_size = len(prompt)
# Define tokenizers and text encoders
tokenizers = [kolors_model['pipeline'].tokenizer]
text_encoders = [kolors_model['pipeline'].text_encoder]
tokenizer = kolors_model['pipeline'].tokenizer
text_encoder = kolors_model['pipeline'].text_encoder
# textual inversion: procecss multi-vector tokens if necessary
prompt_embeds_list = []
for tokenizer, text_encoder in zip(tokenizers, text_encoders):
text_encoder.to(device)
text_inputs = tokenizer(
prompt,
padding="max_length",
max_length=256,
truncation=True,
return_tensors="pt",
).to(device)
text_inputs = tokenizer(
prompt,
padding="max_length",
max_length=256,
truncation=True,
return_tensors="pt",
).to(device)
output = text_encoder(
input_ids=text_inputs['input_ids'] ,
attention_mask=text_inputs['attention_mask'],
position_ids=text_inputs['position_ids'],
output_hidden_states=True)
prompt_embeds = output.hidden_states[-2].permute(1, 0, 2).clone() # [batch_size, 77, 4096]
text_proj = output.hidden_states[-1][-1, :, :].clone() # [batch_size, 4096]
bs_embed, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
output = text_encoder(
input_ids=text_inputs['input_ids'] ,
attention_mask=text_inputs['attention_mask'],
position_ids=text_inputs['position_ids'],
output_hidden_states=True)
prompt_embeds = output.hidden_states[-2].permute(1, 0, 2).clone() # [batch_size, 77, 4096]
text_proj = output.hidden_states[-1][-1, :, :].clone() # [batch_size, 4096]
bs_embed, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
prompt_embeds_list.append(prompt_embeds)
# prompt_embeds = torch.concat(prompt_embeds_list, dim=-1)
prompt_embeds = prompt_embeds_list[0]
# get unconditional embeddings for classifier free guidance
zero_out_negative_prompt = negative_prompt is None and self.config.force_zeros_for_empty_prompt
# if do_classifier_free_guidance and negative_prompt_embeds is None and zero_out_negative_prompt:
# negative_prompt_embeds = torch.zeros_like(prompt_embeds)
# negative_pooled_prompt_embeds = torch.zeros_like(pooled_prompt_embeds)
if do_classifier_free_guidance:
uncond_tokens = []
@@ -199,45 +188,34 @@ class KolorsTextEncode:
)
else:
uncond_tokens = negative_prompt
negative_prompt_embeds_list = []
for tokenizer, text_encoder in zip(tokenizers, text_encoders):
max_length = prompt_embeds.shape[1]
uncond_input = tokenizer(
uncond_tokens,
padding="max_length",
max_length=max_length,
truncation=True,
return_tensors="pt",
).to('cuda')
output = text_encoder(
input_ids=uncond_input['input_ids'] ,
attention_mask=uncond_input['attention_mask'],
position_ids=uncond_input['position_ids'],
output_hidden_states=True)
negative_prompt_embeds = output.hidden_states[-2].permute(1, 0, 2).clone() # [batch_size, 77, 4096]
negative_text_proj = output.hidden_states[-1][-1, :, :].clone() # [batch_size, 4096]
max_length = prompt_embeds.shape[1]
uncond_input = tokenizer(
uncond_tokens,
padding="max_length",
max_length=max_length,
truncation=True,
return_tensors="pt",
).to('cuda')
output = text_encoder(
input_ids=uncond_input['input_ids'] ,
attention_mask=uncond_input['attention_mask'],
position_ids=uncond_input['position_ids'],
output_hidden_states=True)
negative_prompt_embeds = output.hidden_states[-2].permute(1, 0, 2).clone() # [batch_size, 77, 4096]
negative_text_proj = output.hidden_states[-1][-1, :, :].clone() # [batch_size, 4096]
if do_classifier_free_guidance:
# duplicate unconditional embeddings for each generation per prompt, using mps friendly method
seq_len = negative_prompt_embeds.shape[1]
if do_classifier_free_guidance:
# duplicate unconditional embeddings for each generation per prompt, using mps friendly method
seq_len = negative_prompt_embeds.shape[1]
negative_prompt_embeds = negative_prompt_embeds.to(dtype=text_encoder.dtype, device=device)
negative_prompt_embeds = negative_prompt_embeds.to(dtype=text_encoder.dtype, device=device)
negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1)
negative_prompt_embeds = negative_prompt_embeds.view(
batch_size * num_images_per_prompt, seq_len, -1
)
# For classifier free guidance, we need to do two forward passes.
# Here we concatenate the unconditional and text embeddings into a single batch
# to avoid doing two forward passes
negative_prompt_embeds_list.append(negative_prompt_embeds)
# negative_prompt_embeds = torch.concat(negative_prompt_embeds_list, dim=-1)
negative_prompt_embeds = negative_prompt_embeds_list[0]
negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1)
negative_prompt_embeds = negative_prompt_embeds.view(
batch_size * num_images_per_prompt, seq_len, -1
)
bs_embed = text_proj.shape[0]
text_proj = text_proj.repeat(1, num_images_per_prompt).view(
@@ -246,7 +224,9 @@ class KolorsTextEncode:
negative_text_proj = negative_text_proj.repeat(1, num_images_per_prompt).view(
bs_embed * num_images_per_prompt, -1
)
text_encoder.to(offload_device)
mm.soft_empty_cache()
gc.collect()
kolors_embeds = {
'prompt_embeds': prompt_embeds,
'negative_prompt_embeds': negative_prompt_embeds,
@@ -296,6 +276,9 @@ class KolorsSampler:
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.soft_empty_cache()
gc.collect()
pipeline = kolors_model['pipeline']
scheduler_config = {
@@ -334,6 +317,8 @@ class KolorsSampler:
generator= torch.Generator(device).manual_seed(seed)
pipeline.unet.to(device)
latent = pipeline(
prompt=None,
prompt_embeds = kolors_embeds['prompt_embeds'],
@@ -348,6 +333,7 @@ class KolorsSampler:
generator= generator,
).images
pipeline.unet.to(offload_device)
vae_scaling_factor = 0.13025 #SDXL scaling factor
latent = latent / vae_scaling_factor