diff --git a/nodes.py b/nodes.py index 016ca37..a62fabc 100644 --- a/nodes.py +++ b/nodes.py @@ -55,6 +55,18 @@ script_directory = os.path.dirname(os.path.abspath(__file__)) VAE_SCALING_FACTOR = 0.476986 +def add_noise_to_reference_video(image, ratio=None): + if ratio is None: + sigma = torch.normal(mean=-3.0, std=0.5, size=(image.shape[0],)).to(image.device) + sigma = torch.exp(sigma).to(image.dtype) + else: + sigma = torch.ones((image.shape[0],)).to(image.device, image.dtype) * ratio + + image_noise = torch.randn_like(image) * sigma[:, None, None, None, None] + image_noise = torch.where(image==-1, torch.zeros_like(image), image_noise) + image = image + image_noise + return image + def filter_state_dict_by_blocks(state_dict, blocks_mapping): filtered_dict = {} @@ -1388,6 +1400,9 @@ class HyVideoEncode: "spatial_tile_sample_min_size": ("INT", {"default": 256, "min": 32, "max": 2048, "step": 32, "tooltip": "Spatial tile minimum size in pixels, smaller values use less VRAM, may introduce more seams"}), "auto_tile_size": ("BOOLEAN", {"default": True, "tooltip": "Automatically set tile size based on defaults, above settings are ignored"}), }, + "optional": { + "noise_aug_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Strength of noise augmentation"}), + } } RETURN_TYPES = ("LATENT",) @@ -1395,7 +1410,7 @@ class HyVideoEncode: FUNCTION = "encode" CATEGORY = "HunyuanVideoWrapper" - def encode(self, vae, image, enable_vae_tiling, temporal_tiling_sample_size, auto_tile_size, spatial_tile_sample_min_size): + def encode(self, vae, image, enable_vae_tiling, temporal_tiling_sample_size, auto_tile_size, spatial_tile_sample_min_size, noise_aug_strength=0.0): device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -1415,7 +1430,10 @@ class HyVideoEncode: vae.tile_sample_min_size = 256 vae.tile_latent_min_size = 32 - image = (image * 2.0 - 1.0).to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W + image = (image.clone() * 2.0 - 1.0).to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W + if noise_aug_strength > 0.0: + image = add_noise_to_reference_video(image, ratio=noise_aug_strength) + if enable_vae_tiling: vae.enable_tiling() latents = vae.encode(image).latent_dist.sample(generator)