improve ToonCrafter batch processing memory use
This commit is contained in:
+11
-11
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user