Add noise augmentation option

This commit is contained in:
Kijai
2024-06-04 14:29:55 +03:00
parent 6d2b666b7f
commit 161fc8064e
+21 -3
View File
@@ -320,6 +320,7 @@ class DynamiCrafterI2V:
"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}),
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001})
}
}
@@ -328,7 +329,8 @@ class DynamiCrafterI2V:
FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper"
def process(self, model, image, clip_vision, positive, negative, cfg, steps, eta, seed, fs, keep_model_loaded, frames, vae_dtype, frame_window_size=16, frame_window_stride=4, mask=None, image2=None):
def process(self, model, image, clip_vision, positive, negative, cfg, steps, eta, seed, fs, keep_model_loaded,
frames, vae_dtype, frame_window_size=16, frame_window_stride=4, mask=None, image2=None, augmentation_level=0):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.unload_all_models()
@@ -354,6 +356,8 @@ class DynamiCrafterI2V:
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
image = image * 2 - 1
image = image.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0:
image += torch.randn_like(image) * augmentation_level
B, C, H, W = image.shape
orig_H, orig_W = H, W
@@ -374,6 +378,10 @@ class DynamiCrafterI2V:
if image2 is not None:
image2 = image2 * 2 - 1
image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0:
image2 += torch.randn_like(image2) * augmentation_level
if image2.shape != image.shape:
image2 = F.interpolate(image, size=(H, W), mode="bicubic")
z2 = get_latent_z(self.model, image2.unsqueeze(2)) #bc,1,hw
@@ -516,6 +524,7 @@ class ToonCrafterInterpolation:
},
"optional": {
"image_embed_ratio": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001})
}
}
@@ -524,7 +533,7 @@ class ToonCrafterInterpolation:
FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper"
def process(self, model, clip_vision, images, positive, negative, cfg, steps, eta, seed, fs, frames, vae_dtype, image_embed_ratio=1.0):
def process(self, model, clip_vision, images, positive, negative, cfg, steps, eta, seed, fs, frames, vae_dtype, image_embed_ratio=1.0, augmentation_level=0):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.unload_all_models()
@@ -574,13 +583,22 @@ class ToonCrafterInterpolation:
self.model.first_stage_model.to(device)
if augmentation_level > 0:
image += torch.randn_like(image) * augmentation_level
image2 += torch.randn_like(image) * augmentation_level
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)
encode_pixels = videos
# if augmentation_level > 0:
# encode_pixels += torch.randn_like(videos) * augmentation_level
z, hs = get_latent_z_with_hidden_states(self.model, encode_pixels)
hidden_states.append(hs)
img_tensor_repeat = torch.zeros_like(z)