From 7866e6b239821d0079cf1e312cc91662b9efefb1 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 1 Jun 2024 18:15:08 +0300 Subject: [PATCH] Tooncrafter refactor --- examples/tooncrafter_3_frames_example_01.json | 794 ++++++++++++++++++ nodes.py | 316 +++---- 2 files changed, 964 insertions(+), 146 deletions(-) create mode 100644 examples/tooncrafter_3_frames_example_01.json diff --git a/examples/tooncrafter_3_frames_example_01.json b/examples/tooncrafter_3_frames_example_01.json new file mode 100644 index 0000000..66aeb66 --- /dev/null +++ b/examples/tooncrafter_3_frames_example_01.json @@ -0,0 +1,794 @@ +{ + "last_node_id": 45, + "last_link_id": 100, + "nodes": [ + { + "id": 5, + "type": "ImageResizeKJ", + "pos": [ + 849, + 197 + ], + "size": { + "0": 315, + "1": 242 + }, + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 2 + }, + { + "name": "get_image_size", + "type": "IMAGE", + "link": null + }, + { + "name": "width_input", + "type": "INT", + "link": null, + "widget": { + "name": "width_input" + } + }, + { + "name": "height_input", + "type": "INT", + "link": null, + "widget": { + "name": "height_input" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 71, + 73, + 98 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "width", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "height", + "type": "INT", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "ImageResizeKJ" + }, + "widgets_values": [ + 512, + 512, + "lanczos", + true, + 64, + 0, + 0 + ] + }, + { + "id": 42, + "type": "ToonCrafterInterpolation", + "pos": [ + 1488, + 192 + ], + "size": [ + 400.61801488194305, + 308.4813612777184 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "DCMODEL", + "link": 96, + "slot_index": 0 + }, + { + "name": "images", + "type": "IMAGE", + "link": 91 + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "links": [ + 92 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "ToonCrafterInterpolation" + }, + "widgets_values": [ + 20, + 7, + 1, + 16, + "anime scene", + 2, + "fixed", + 10, + "auto" + ] + }, + { + "id": 4, + "type": "DownloadAndLoadDynamiCrafterModel", + "pos": [ + 1055, + 6 + ], + "size": { + "0": 393, + "1": 106 + }, + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [ + { + "name": "opt_openclippath", + "type": "OPENCLIPVISIONPATH", + "link": null + } + ], + "outputs": [ + { + "name": "DynCraft_model", + "type": "DCMODEL", + "links": [ + 95, + 96 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadDynamiCrafterModel" + }, + "widgets_values": [ + "tooncrafter_512_interp-fp16.safetensors", + "auto", + false + ] + }, + { + "id": 2, + "type": "LoadImage", + "pos": [ + 486, + 567 + ], + "size": { + "0": 315, + "1": 314 + }, + "flags": {}, + "order": 1, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 6 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "72109_125.mp4_00-00 (2).png", + "image" + ] + }, + { + "id": 1, + "type": "LoadImage", + "pos": [ + 490, + 196 + ], + "size": { + "0": 315, + "1": 314 + }, + "flags": {}, + "order": 2, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 2 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "clipspace/clipspace-mask-8168.9000000003725.png [input]", + "image" + ] + }, + { + "id": 7, + "type": "ImageResizeKJ", + "pos": [ + 845, + 504 + ], + "size": { + "0": 315, + "1": 242 + }, + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 6 + }, + { + "name": "get_image_size", + "type": "IMAGE", + "link": 73 + }, + { + "name": "width_input", + "type": "INT", + "link": null, + "widget": { + "name": "width_input" + } + }, + { + "name": "height_input", + "type": "INT", + "link": null, + "widget": { + "name": "height_input" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 68 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "width", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "height", + "type": "INT", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "ImageResizeKJ" + }, + "widgets_values": [ + 512, + 512, + "lanczos", + true, + 64, + 0, + 0 + ] + }, + { + "id": 44, + "type": "LoadImage", + "pos": [ + 484, + 938 + ], + "size": { + "0": 315, + "1": 314 + }, + "flags": {}, + "order": 3, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 99 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "clipspace/clipspace-mask-8168.9000000003725.png [input]", + "image" + ] + }, + { + "id": 28, + "type": "ImageBatchMulti", + "pos": [ + 1237, + 195 + ], + "size": { + "0": 210, + "1": 122 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "image_1", + "type": "IMAGE", + "link": 71 + }, + { + "name": "image_2", + "type": "IMAGE", + "link": 68 + }, + { + "name": "image_3", + "type": "IMAGE", + "link": 100 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 93 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "ImageBatchMulti" + }, + "widgets_values": [ + 3, + null + ] + }, + { + "id": 6, + "type": "GetImageSizeAndCount", + "pos": [ + 1234, + 379 + ], + "size": { + "0": 210, + "1": 86 + }, + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 93 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 91 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "512 width", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "320 height", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "3 count", + "type": "INT", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "GetImageSizeAndCount" + } + }, + { + "id": 45, + "type": "ImageResizeKJ", + "pos": [ + 850, + 918 + ], + "size": [ + 315, + 242 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 99 + }, + { + "name": "get_image_size", + "type": "IMAGE", + "link": 98 + }, + { + "name": "width_input", + "type": "INT", + "link": null, + "widget": { + "name": "width_input" + } + }, + { + "name": "height_input", + "type": "INT", + "link": null, + "widget": { + "name": "height_input" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 100 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "width", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "height", + "type": "INT", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "ImageResizeKJ" + }, + "widgets_values": [ + 512, + 512, + "lanczos", + true, + 64, + 0, + 0 + ] + }, + { + "id": 25, + "type": "ToonCrafterDecode", + "pos": [ + 1929, + 6 + ], + "size": { + "0": 315, + "1": 102 + }, + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "DCMODEL", + "link": 95 + }, + { + "name": "latent", + "type": "LATENT", + "link": 92 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 85 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "ToonCrafterDecode" + }, + "widgets_values": [ + "auto" + ] + }, + { + "id": 29, + "type": "VHS_VideoCombine", + "pos": [ + 1917, + 194 + ], + "size": [ + 1271.3231201171875, + 1086.0769500732422 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 85 + }, + { + "name": "audio", + "type": "VHS_AUDIO", + "link": null + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 8, + "loop_count": 0, + "filename_prefix": "AnimateDiff", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": false, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "AnimateDiff_00006.mp4", + "subfolder": "", + "type": "temp", + "format": "video/h264-mp4" + } + } + } + } + ], + "links": [ + [ + 2, + 1, + 0, + 5, + 0, + "IMAGE" + ], + [ + 6, + 2, + 0, + 7, + 0, + "IMAGE" + ], + [ + 68, + 7, + 0, + 28, + 1, + "IMAGE" + ], + [ + 71, + 5, + 0, + 28, + 0, + "IMAGE" + ], + [ + 73, + 5, + 0, + 7, + 1, + "IMAGE" + ], + [ + 85, + 25, + 0, + 29, + 0, + "IMAGE" + ], + [ + 91, + 6, + 0, + 42, + 1, + "IMAGE" + ], + [ + 92, + 42, + 0, + 25, + 1, + "LATENT" + ], + [ + 93, + 28, + 0, + 6, + 0, + "IMAGE" + ], + [ + 95, + 4, + 0, + 25, + 0, + "DCMODEL" + ], + [ + 96, + 4, + 0, + 42, + 0, + "DCMODEL" + ], + [ + 98, + 5, + 0, + 45, + 1, + "IMAGE" + ], + [ + 99, + 44, + 0, + 45, + 0, + "IMAGE" + ], + [ + 100, + 45, + 0, + 28, + 2, + "IMAGE" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 0.5644739300537778, + "offset": [ + -260.37146856633785, + 316.4086463973054 + ] + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/nodes.py b/nodes.py index 0b989da..cc40652 100644 --- a/nodes.py +++ b/nodes.py @@ -400,7 +400,7 @@ class DynamiCrafterI2V: ) 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." - + ## reconstruct from latent to pixel space self.model.first_stage_model.to(device) decoded_images = self.model.decode_first_stage(samples) #b c t h w @@ -424,12 +424,12 @@ class DynamiCrafterI2V: last_image = video[-1].unsqueeze(0) return (video, last_image) -class ToonCrafterI2V: +class ToonCrafterInterpolation: @classmethod def INPUT_TYPES(s): return {"required": { "model": ("DCMODEL",), - "image": ("IMAGE",), + "images": ("IMAGE",), "steps": ("INT", {"default": 20, "min": 1, "max": 200, "step": 1}), "cfg": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 200.0, "step": 0.01}), "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), @@ -437,7 +437,6 @@ class ToonCrafterI2V: "prompt": ("STRING", {"multiline": True, "default": "",}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "fs": ("INT", {"default": 10, "min": 2, "max": 100, "step": 1}), - "keep_model_loaded": ("BOOLEAN", {"default": True}), "vae_dtype": ( [ 'fp32', @@ -447,24 +446,15 @@ class ToonCrafterI2V: ], { "default": 'auto' }), - }, - "optional": { - "image2": ("IMAGE",), - "mask": ("MASK",), - "frame_window_size": ("INT", {"default": 16, "min": 1, "max": 200, "step": 1}), - "frame_window_stride": ("INT", {"default": 4, "min": 1, "max": 200, "step": 1}), - "num_videos": ("INT", {"default": 1, "min": 1, "max": 1000, "step": 1}), - "prune_first_last": ("BOOLEAN", {"default": True}), - } } - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("images",) + RETURN_TYPES = ("LATENT",) + RETURN_NAMES = ("samples",) FUNCTION = "process" CATEGORY = "DynamiCrafterWrapper" - def process(self, model, image, image2, prompt, cfg, steps, eta, seed, fs, keep_model_loaded, frames, vae_dtype, frame_window_size=16, frame_window_stride=4, mask=None, prune_first_last=True, num_videos=1, **kwargs): + def process(self, model, images, prompt, cfg, steps, eta, seed, fs, frames, vae_dtype): device = mm.get_torch_device() mm.unload_all_models() mm.soft_empty_cache() @@ -483,111 +473,95 @@ class ToonCrafterI2V: model.first_stage_model.to(convert_dtype(vae_dtype)) print(f"VAE using dtype: {model.first_stage_model.dtype}") + images = images * 2 - 1 + images = images.permute(0, 3, 1, 2).to(dtype).to(device) + + B, C, H, W = images.shape + orig_H, orig_W = H, W + if W % 64 != 0: + W = W - (W % 64) + if H % 64 != 0: + H = H - (H % 64) + if orig_H % 64 != 0 or orig_W % 64 != 0: + images = F.interpolate(images, size=(H, W), mode="bicubic") + self.model = model self.model.to(device) + + out = [] + hidden_states = [] + 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(): - videos, videos2 = None, None - image = image * 2 - 1 - image = image.permute(0, 3, 1, 2).to(dtype).to(device) + for i in range(len(images) - 1): + videos, videos2 = None, None + image = images[i].unsqueeze(0) + image2 = images[i+1].unsqueeze(0) + + B, C, H, W = image.shape + noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8] - B, C, H, W = image.shape - orig_H, orig_W = H, W - if W % 64 != 0: - W = W - (W % 64) - if H % 64 != 0: - H = H - (H % 64) - if orig_H % 64 != 0 or orig_W % 64 != 0: - image = F.interpolate(image, size=(H, W), mode="bicubic") - - B, C, H, W = image.shape - noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8] + self.model.first_stage_model.to(device) - self.model.first_stage_model.to(device) + videos = image.unsqueeze(2) # bc1hw + videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames//2) + videos2 = image2.unsqueeze(2) # bc1hw + videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames//2) + videos = torch.cat([videos, videos2], dim=2) + z, hs = get_latent_z_with_hidden_states(self.model, videos) + hidden_states.append(hs) - image2 = image2 * 2 - 1 - image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device) - if image2.shape != image.shape: - image2 = F.interpolate(image2, size=(H, W), mode="bicubic") + img_tensor_repeat = torch.zeros_like(z) + img_tensor_repeat[:,:,:1,:,:] = z[:,:,:1,:,:] + img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:] - videos = image.unsqueeze(2) # bc1hw - videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames//2) + self.model.first_stage_model.to('cpu') - videos2 = image2.unsqueeze(2) # bc1hw - videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames//2) + self.model.cond_stage_model.to(device) + self.model.embedder.to(device) + self.model.image_proj_model.to(device) - videos = torch.cat([videos, videos2], dim=2) + text_emb = self.model.get_learned_conditioning([prompt]) + cond_images = self.model.embedder(image) + img_emb = self.model.image_proj_model(cond_images) + imtext_cond = torch.cat([text_emb, img_emb], dim=1) + del cond_images, img_emb, text_emb - z, hs = get_latent_z_with_hidden_states(self.model, videos) + fs = torch.tensor([fs], dtype=torch.long, device=self.model.device) + cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]} - img_tensor_repeat = torch.zeros_like(z) - img_tensor_repeat[:,:,:1,:,:] = z[:,:,:1,:,:] - img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:] - - - self.model.first_stage_model.to('cpu') - - self.model.cond_stage_model.to(device) - self.model.embedder.to(device) - self.model.image_proj_model.to(device) - - text_emb = self.model.get_learned_conditioning([prompt]) - cond_images = self.model.embedder(image) - img_emb = self.model.image_proj_model(cond_images) - imtext_cond = torch.cat([text_emb, img_emb], dim=1) - del cond_images, img_emb, text_emb - - fs = torch.tensor([fs], dtype=torch.long, device=self.model.device) - cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]} - - if noise_shape[-1] == 32: - timestep_spacing = "uniform" - guidance_rescale = 0.0 - else: - timestep_spacing = "uniform_trailing" - guidance_rescale = 0.7 - - ## construct unconditional guidance - if cfg != 1.0: - uc_emb = self.model.get_learned_conditioning([""]) - ## process image embedding token - if hasattr(self.model, 'embedder'): - uc_img = torch.zeros(noise_shape[0],3,224,224).to(self.model.device) - ## img: b c h w >> b l c - uc_img = self.model.embedder(uc_img) - uc_img = self.model.image_proj_model(uc_img) - uc_emb = torch.cat([uc_emb, uc_img], dim=1) - if isinstance(cond, dict): - uc = {key:cond[key] for key in cond.keys()} - uc.update({'c_crossattn': [uc_emb]}) + if noise_shape[-1] == 32: + timestep_spacing = "uniform" + guidance_rescale = 0.0 else: - uc = uc_emb - else: - uc = None + timestep_spacing = "uniform_trailing" + guidance_rescale = 0.7 - self.model.cond_stage_model.to('cpu') - self.model.embedder.to('cpu') - self.model.image_proj_model.to('cpu') - - if mask is not None: - mask = mask.to(dtype).to(device) - mask = F.interpolate(mask.unsqueeze(0), size=(H // 8, W // 8), mode="nearest").squeeze(0) - mask = (1 - mask) - mask = mask.unsqueeze(1) - B, C, H, W = mask.shape - if B < frames: - mask = mask.unsqueeze(2) - mask = mask.expand(-1, -1, frames, -1, -1) + ## construct unconditional guidance + if cfg != 1.0: + uc_emb = self.model.get_learned_conditioning([""]) + ## process image embedding token + if hasattr(self.model, 'embedder'): + uc_img = torch.zeros(noise_shape[0],3,224,224).to(self.model.device) + ## img: b c h w >> b l c + uc_img = self.model.embedder(uc_img) + uc_img = self.model.image_proj_model(uc_img) + uc_emb = torch.cat([uc_emb, uc_img], dim=1) + if isinstance(cond, dict): + uc = {key:cond[key] for key in cond.keys()} + uc.update({'c_crossattn': [uc_emb]}) + else: + uc = uc_emb else: - mask = mask.unsqueeze(0) - mask = mask.permute(0, 2, 1, 3, 4) - mask = torch.where(mask < 1.0, torch.tensor(0.0, device=device, dtype=dtype), torch.tensor(1.0, device=device, dtype=dtype)) + uc = None + + self.model.cond_stage_model.to('cpu') + self.model.embedder.to('cpu') + self.model.image_proj_model.to('cpu') + + #inference - #inference - - video_list = [] - for i in range(num_videos): self.model.model.diffusion_model.to(device) ddim_sampler = DDIMSampler(self.model) samples, _ = ddim_sampler.sample(S=steps, @@ -605,51 +579,99 @@ class ToonCrafterI2V: timestep_spacing=timestep_spacing, guidance_rescale=guidance_rescale, clean_cond=True, - mask=mask, - x0=img_tensor_repeat.clone() if mask is not None else None, - frame_window_size = frame_window_size, - frame_window_stride = frame_window_stride, + mask=None, + x0=None, + frame_window_size = 16, + frame_window_stride = 4, ) - + 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) + out.append(samples) - ## reconstruct from latent to pixel space - self.model.model.diffusion_model.to('cpu') - mm.soft_empty_cache() - self.model.first_stage_model.to(device) - if mm.XFORMERS_IS_AVAILABLE: - print("Using xformers") - additional_decode_kwargs = {'ref_context': hs} - decoded_images = self.model.decode_first_stage(samples, **additional_decode_kwargs) #b c t h w + self.model.to('cpu') + mm.soft_empty_cache() + + samples = torch.cat(out, dim=0) + samples = samples / 0.18215 + + latent = { + "samples": samples, + "hidden_states": hidden_states, + } + return (latent,) + +class ToonCrafterDecode: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model": ("DCMODEL",), + "latent": ("LATENT",), + "vae_dtype": ( + [ + 'fp32', + 'fp16', + 'bf16', + 'auto' + ], { + "default": 'auto' + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "process" + CATEGORY = "DynamiCrafterWrapper" + + def process(self, model, latent, vae_dtype): + device = mm.get_torch_device() + mm.soft_empty_cache() + + samples = latent["samples"] + samples = samples * 0.18215 + hs = latent["hidden_states"] + + if vae_dtype == "auto": + try: + if mm.should_use_bf16(): + model.first_stage_model.to(convert_dtype('bf16')) else: - print("xformers not available, ToonCrafter does not work well without it.") - decoded_images = self.model.decode_first_stage(samples) #b c t h w - self.model.first_stage_model.to('cpu') - - video = decoded_images.detach().cpu() - video = torch.clamp(video.float(), -1., 1.) - video = (video + 1.0) / 2.0 - video = video.squeeze(0).permute(1, 2, 3, 0) - if prune_first_last: - video = video[1:-1] - video_list.append(video) - del decoded_images, samples, video - - if not keep_model_loaded: - self.model.to('cpu') - mm.soft_empty_cache() - # Ensure the final dimensions are divisible by 2 - final_H = (orig_H // 2) * 2 - final_W = (orig_W // 2) * 2 - - video_out = torch.cat(video_list, dim=0) - - if video_out.shape[1] != final_H or video_out.shape[2] != final_W: - video_out = F.interpolate(video_out.permute(0, 3, 1, 2), size=(final_H, final_W), mode="bicubic").permute(0, 2, 3, 1) + model.first_stage_model.to(convert_dtype('fp32')) + except: + raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtype manually.") + else: + model.first_stage_model.to(convert_dtype(vae_dtype)) + print(f"VAE using dtype: {model.first_stage_model.dtype}") + out = [] + iteration_counter = 0 + for i in range(0, samples.shape[0], 16): - return (video_out, ) - + 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: + 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) + model.first_stage_model.to('cpu') + + return (video_out,) + class DynamiCrafterBatchInterpolation: @classmethod def INPUT_TYPES(s): @@ -858,7 +880,8 @@ NODE_CLASS_MAPPINGS = { "DynamiCrafterI2V": DynamiCrafterI2V, "DynamiCrafterModelLoader": DynamiCrafterModelLoader, "DynamiCrafterBatchInterpolation": DynamiCrafterBatchInterpolation, - "ToonCrafterI2V": ToonCrafterI2V, + "ToonCrafterInterpolation": ToonCrafterInterpolation, + "ToonCrafterDecode": ToonCrafterDecode, "OpenCLIPVisionSelect": OpenCLIPVisionSelect, "DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel @@ -868,6 +891,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "DynamiCrafterModelLoader": "DynamiCrafterModelLoader", "DynamiCrafterBatchInterpolation": "DynamiCrafterBatchInterpolation", "OpenCLIPVisionSelect": "OpenCLIPVisionSelect", - "ToonCrafterI2V": "ToonCrafterI2V", + "ToonCrafterInterpolation": "ToonCrafterInterpolation", + "ToonCrafterDecode": "ToonCrafterDecode", "DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel" }