improve ToonCrafter batch processing memory use

This commit is contained in:
kijai
2024-06-08 17:29:23 +03:00
parent ad9ce82f2d
commit 5ce1764825
2 changed files with 47 additions and 38 deletions
+10 -10
View File
@@ -530,17 +530,17 @@ class LatentDiffusion(DDPM):
n_samples = default(self.en_and_decode_n_samples_a_time, self.temporal_length)
n_rounds = math.ceil(z.shape[0] / n_samples)
with torch.autocast(mm.get_autocast_device(device), enabled=True):
for n in range(n_rounds):
if isinstance(self.first_stage_model.decoder, VideoDecoder):
kwargs.update({"timesteps": len(z[n * n_samples : (n + 1) * n_samples])})
else:
kwargs = {}
#with torch.autocast(mm.get_autocast_device(device), enabled=True):
for n in range(n_rounds):
if isinstance(self.first_stage_model.decoder, VideoDecoder):
kwargs.update({"timesteps": len(z[n * n_samples : (n + 1) * n_samples])})
else:
kwargs = {}
out = self.first_stage_model.decode(
z[n * n_samples : (n + 1) * n_samples], **kwargs
)
results.append(out)
out = self.first_stage_model.decode(
z[n * n_samples : (n + 1) * n_samples], **kwargs
)
results.append(out)
results = torch.cat(results, dim=0)
if reshape_back:
+34 -25
View File
@@ -570,11 +570,12 @@ class ToonCrafterInterpolation:
out = []
hidden_states = []
pbar = comfy.utils.ProgressBar(len(images) - 1)
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
for i in range(len(images) - 1):
videos, videos2 = None, None
mm.soft_empty_cache()
image = images[i].unsqueeze(0)
image2 = images[i+1].unsqueeze(0)
@@ -596,6 +597,7 @@ class ToonCrafterInterpolation:
videos = torch.cat([videos, videos2], dim=2)
z, hs = get_latent_z_with_hidden_states(self.model, videos)
hs = [t.to("cpu") for t in hs]
hidden_states.append(hs)
img_tensor_repeat = torch.zeros_like(z)
@@ -603,7 +605,6 @@ class ToonCrafterInterpolation:
img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:]
self.model.first_stage_model.to(offload_device)
print("first stage model device: ", self.model.first_stage_model.device)
text_emb = positive[0][0].to(device)
@@ -678,8 +679,9 @@ class ToonCrafterInterpolation:
)
print(f"Sampled {i+1} out of {(len(images) - 1)}")
assert not torch.isnan(samples).any().item(), "Resulting tensor containts NaNs. I'm unsure why this happens, changing step count and/or image dimensions might help."
samples = samples.squeeze(0).permute(1, 0, 2, 3)
samples = samples.squeeze(0).permute(1, 0, 2, 3).to("cpu").to(model.first_stage_model.dtype)
out.append(samples)
pbar.update(1)
self.model.to(offload_device)
mm.soft_empty_cache()
@@ -721,10 +723,15 @@ class ToonCrafterDecode:
def process(self, model, latent, vae_dtype, prune_last_frame=False):
device = mm.get_torch_device()
mm.unload_all_models()
mm.soft_empty_cache()
samples = latent["samples"]
num_samples = samples.shape[0]
samples = samples * 0.18215
model.first_stage_model.to(device)
#samples = samples.to(model.first_stage_model.device)
hs = latent["hidden_states"]
model.en_and_decode_n_samples_a_time = 16
@@ -742,31 +749,33 @@ class ToonCrafterDecode:
print(f"VAE using dtype: {model.first_stage_model.dtype}")
out = []
iteration_counter = 0
for i in range(0, samples.shape[0], 16):
pbar = comfy.utils.ProgressBar(num_samples // 16)
autocast_condition = (model.first_stage_model.dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
for i in range(0, num_samples, 16):
batch_start = i
batch_end = min(i + 16, samples.shape[0]) # Ensure we don't go beyond the tensor's size
batch_samples = samples[batch_start:batch_end]
model.first_stage_model.to(device)
if mm.XFORMERS_IS_AVAILABLE:
print("Using xformers")
additional_decode_kwargs = {'ref_context': hs[iteration_counter]}
decoded_images = model.decode_first_stage(batch_samples, **additional_decode_kwargs) #b c t h w
else:
raise Exception("XFormers not available, it is required for ToonCrafter decoder. Alternatively you can use a standard VAE Decode -node instead, but this has a negative effect on the image quality though.")
#print("xformers not available, ToonCrafter does not work well without it.")
#decoded_images = model.decode_first_stage(batch_samples) #b c t h w
batch_end = min(i + 16, num_samples) # Ensure we don't go beyond the tensor's size
batch_samples = samples[batch_start:batch_end].to(model.first_stage_model.device)
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=model.first_stage_model.dtype) if autocast_condition else nullcontext():
if mm.XFORMERS_IS_AVAILABLE:
print(f"Decoding frames {iteration_counter * 16} - {16 + iteration_counter * 16} out of {num_samples} using xformers")
hs_ = hs[iteration_counter]
hs_ = [t.to(model.first_stage_model.device) for t in hs_]
additional_decode_kwargs = {'ref_context': hs_}
decoded_images = model.decode_first_stage(batch_samples, **additional_decode_kwargs) #b c t h w
else:
raise Exception("XFormers not available, it is required for ToonCrafter decoder. Alternatively you can use a standard VAE Decode -node instead, but this has a negative effect on the image quality though.")
video = decoded_images.detach().cpu()
video = torch.clamp(video.float(), -1., 1.)
video = (video + 1.0) / 2.0
video = video.squeeze(0).permute(0, 2, 3, 1)
iteration_counter += 1
out.append(video)
del decoded_images
mm.soft_empty_cache()
video_out = torch.cat(out, dim=0)
video = decoded_images.detach().cpu()
video = torch.clamp(video.float(), -1., 1.)
video = (video + 1.0) / 2.0
video = video.squeeze(0).permute(0, 2, 3, 1)
iteration_counter += 1
pbar.update(1)
out.append(video)
del decoded_images
mm.soft_empty_cache()
model.first_stage_model.to('cpu')
video_out = torch.cat(out, dim=0)
if prune_last_frame:
video_out = video_out[torch.arange(video_out.shape[0]) % 16!= 0]