noise augment for i2v

This commit is contained in:
kijai
2025-02-01 15:13:40 +02:00
parent 43cb25a038
commit e34492ad47
+20 -2
View File
@@ -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)