From 5ce17648251e9f3f7c4fdf8f7a2c34cf4ae2e597 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 8 Jun 2024 17:29:23 +0300 Subject: [PATCH] improve ToonCrafter batch processing memory use --- lvdm/models/ddpm3d.py | 22 +++++++-------- nodes.py | 63 ++++++++++++++++++++++++------------------- 2 files changed, 47 insertions(+), 38 deletions(-) diff --git a/lvdm/models/ddpm3d.py b/lvdm/models/ddpm3d.py index de554da..0083855 100644 --- a/lvdm/models/ddpm3d.py +++ b/lvdm/models/ddpm3d.py @@ -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 = {} - - out = self.first_stage_model.decode( - z[n * n_samples : (n + 1) * n_samples], **kwargs - ) - results.append(out) + #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) results = torch.cat(results, dim=0) if reshape_back: diff --git a/nodes.py b/nodes.py index 0b4be44..e3b76dc 100644 --- a/nodes.py +++ b/nodes.py @@ -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 - - 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) + 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 + 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]