From 209c65ba41993c7a9425697f217990df9d13d6d2 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 25 Feb 2025 22:29:12 +0200 Subject: [PATCH] fix I2V --- nodes.py | 105 +++++++++++++++++++++++++++++++++++++------------------ 1 file changed, 71 insertions(+), 34 deletions(-) diff --git a/nodes.py b/nodes.py index 6566549..3e733fb 100644 --- a/nodes.py +++ b/nodes.py @@ -107,12 +107,35 @@ class WanVideoModel(comfy.model_base.BaseModel): def __setitem__(self, k, v): self.pipeline[k] = v +from comfy.latent_formats import LatentFormat +class WanVideo(LatentFormat): + latent_channels = 16 + latent_dimensions = 3 + scale_factor = 1 + latent_rgb_factors = [[0.00015850099142733433, -0.00022336468679047376, 0.0012986971243192447], + [0.0005663412275114018, 0.0007861870689866654, 0.0019476977664748178], + [0.0015309811379193067, -0.00033673814517196977, 0.0008630955780593666], + [0.0018870473626590339, 0.002189569168822193, 0.002116887481625339], + [0.0020317996817509638, 0.0007815605945740026, -0.0005115450818442968], + [0.0016342762423795499, 0.0012601010475289658, 0.0016851194126694853], + [0.0013595118767078232, -0.0002916941854044625, 0.00018943102991771045], + [0.001410154384770368, 0.0007686194765380242, 0.001934588961229726], + [-0.00036535934889578086, 0.00021113539535803916, 0.00039667130378409675], + [-9.082161140163455e-05, 0.0013325911783233892, 0.001812325948391347], + [0.00020121251545686565, 0.0018655655639155274, 0.0005459994991828317], + [0.0018891414023764184, 0.0005440105401015541, -0.0002365743607780385], + [0.0017790556111146022, 2.214497568459961e-05, 0.0017639757911266463], + [0.001456174042709703, 0.00043078224591916133, 0.0015744138130009018], + [0.0017913272306190463, 0.0017379684971510461, -0.00012070215501199769], + [-3.400939289556232e-05, -0.0004053172077750597, 0.0007082065661536516]] + + latent_rgb_factors_bias = [-0.0011, 0.0, -0.0002] class WanVideoModelConfig: def __init__(self, dtype): self.unet_config = {} self.unet_extra_config = {} - self.latent_format = comfy.latent_formats.HunyuanVideo #todo: change to WanVideo + self.latent_format = WanVideo #todo: change to WanVideo self.latent_format.latent_channels = 16 self.manual_cast_dtype = dtype self.sampling_settings = {"multiplier": 1.0} @@ -797,14 +820,14 @@ class WanVideoSampler: else: #t2v target_shape = image_embeds["target_shape"] seq_len = image_embeds["max_seq_len"] - noise = torch.randn( - target_shape[0], - target_shape[1], - target_shape[2], - target_shape[3], - dtype=torch.float32, - device=torch.device("cpu"), - generator=seed_g) + noise = torch.randn( + target_shape[0], + target_shape[1], + target_shape[2], + target_shape[3], + dtype=torch.float32, + device=torch.device("cpu"), + generator=seed_g) latent = noise.to(device) @@ -835,8 +858,12 @@ class WanVideoSampler: pbar = ProgressBar(steps) + #from latent_preview import prepare_callback + #callback = prepare_callback(self.comfy_model, steps) + callback=None + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): - for _, t in enumerate(tqdm(timesteps)): + for i, t in enumerate(tqdm(timesteps)): latent_model_input = [latent.to(device)] timestep = [t] @@ -863,7 +890,16 @@ class WanVideoSampler: x0 = [latent.to(device)] del latent_model_input, timestep - pbar.update(1) + if callback is not None: + callback_latent = (latent_model_input[:, :16, :, :, :] - noise_pred * t / 1000).detach()[0].permute(1,0,2,3) + callback( + i, + callback_latent, + None, + steps + ) + else: + pbar.update(1) if force_offload: transformer.to(offload_device) @@ -975,11 +1011,11 @@ class WanVideoLatentPreview: "required": { "samples": ("LATENT",), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "min_val": ("FLOAT", {"default": -0.15, "min": -1.0, "max": 0.0, "step": 0.001}), - "max_val": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001}), - "r_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.001}), - "g_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.001}), - "b_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.001}), + "min_val": ("FLOAT", {"default": -0.15, "min": -1.0, "max": 0.0, "step": 0.0001}), + "max_val": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.0001}), + "r_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.0001}), + "g_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.0001}), + "b_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.0001}), }, } @@ -994,31 +1030,32 @@ class WanVideoLatentPreview: latents = samples["samples"].clone() print("in sample", latents.shape) #latent_rgb_factors =[[-0.02531045419704009, -0.00504800612542497, 0.13293717293982546], [-0.03421835830845858, 0.13996708548892614, -0.07081038680118075], [0.011091819063647063, -0.03372949685846012, -0.0698232210116172], [-0.06276524604742019, -0.09322986677909442, 0.01826383612148913], [0.021290659938126788, -0.07719530444034409, -0.08247812477766273], [0.04401102991215147, -0.0026401932105894754, -0.01410913586718443], [0.08979717602613707, 0.05361221258740831, 0.11501425309699129], [0.04695121980405198, -0.13053491609675175, 0.05025986885867986], [-0.09704684176098193, 0.03397687417738002, -0.1105886644677771], [0.14694697234804935, -0.12316902186157716, 0.04210404546699645], [0.14432470831243552, -0.002580008133591355, -0.08490676947390643], [0.051502750076553944, -0.10071695490292451, -0.01786223610178095], [-0.12503276881774464, 0.08877830923879379, 0.1076584501927316], [-0.020191205513213406, -0.1493425056303128, -0.14289740371758308], [-0.06470138952271293, -0.07410426095060325, 0.00980804676890873], [0.11747671720735695, 0.10916082743849789, -0.12235599365235904]] - latent_rgb_factors = [[-0.41, -0.25, -0.26], - [-0.26, -0.49, -0.24], - [-0.37, -0.54, -0.3], - [-0.04, -0.29, -0.29], - [-0.52, -0.59, -0.39], - [-0.56, -0.6, -0.02], - [-0.53, -0.06, -0.48], - [-0.51, -0.28, -0.18], - [-0.59, -0.1, -0.33], - [-0.56, -0.54, -0.41], - [-0.61, -0.19, -0.5], - [-0.05, -0.25, -0.17], - [-0.23, -0.04, -0.22], - [-0.51, -0.56, -0.43], - [-0.13, -0.4, -0.05], - [-0.01, -0.01, -0.48]] + latent_rgb_factors = [[0.00015850099142733433, -0.00022336468679047376, 0.0012986971243192447], + [0.0005663412275114018, 0.0007861870689866654, 0.0019476977664748178], + [0.0015309811379193067, -0.00033673814517196977, 0.0008630955780593666], + [0.0018870473626590339, 0.002189569168822193, 0.002116887481625339], + [0.0020317996817509638, 0.0007815605945740026, -0.0005115450818442968], + [0.0016342762423795499, 0.0012601010475289658, 0.0016851194126694853], + [0.0013595118767078232, -0.0002916941854044625, 0.00018943102991771045], + [0.001410154384770368, 0.0007686194765380242, 0.001934588961229726], + [-0.00036535934889578086, 0.00021113539535803916, 0.00039667130378409675], + [-9.082161140163455e-05, 0.0013325911783233892, 0.001812325948391347], + [0.00020121251545686565, 0.0018655655639155274, 0.0005459994991828317], + [0.0018891414023764184, 0.0005440105401015541, -0.0002365743607780385], + [0.0017790556111146022, 2.214497568459961e-05, 0.0017639757911266463], + [0.001456174042709703, 0.00043078224591916133, 0.0015744138130009018], + [0.0017913272306190463, 0.0017379684971510461, -0.00012070215501199769], + [-3.400939289556232e-05, -0.0004053172077750597, 0.0007082065661536516]] import random random.seed(seed) #latent_rgb_factors = [[random.uniform(min_val, max_val) for _ in range(3)] for _ in range(16)] + #latent_rgb_factors = [[0.1 for _ in range(3)] for _ in range(16)] out_factors = latent_rgb_factors print(latent_rgb_factors) - #latent_rgb_factors_bias = [0.138, 0.025, -0.299] - latent_rgb_factors_bias = [r_bias, g_bias, b_bias] + latent_rgb_factors_bias = [-0.0011, 0.0, -0.0002] + #latent_rgb_factors_bias = [r_bias, g_bias, b_bias] latent_rgb_factors = torch.tensor(latent_rgb_factors, device=latents.device, dtype=latents.dtype).transpose(0, 1) latent_rgb_factors_bias = torch.tensor(latent_rgb_factors_bias, device=latents.device, dtype=latents.dtype)