comdoleger commited on
Commit
f9d7131
·
verified ·
1 Parent(s): d1b24ee

Upload extensions_built_in/diffusion_models/wan22/wan22_14b_i2v_model.py with huggingface_hub

Browse files
extensions_built_in/diffusion_models/wan22/wan22_14b_i2v_model.py ADDED
@@ -0,0 +1,144 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from toolkit.models.wan21.wan_utils import add_first_frame_conditioning
3
+ from toolkit.prompt_utils import PromptEmbeds
4
+ from PIL import Image
5
+ import torch
6
+ from toolkit.config_modules import GenerateImageConfig
7
+ from .wan22_pipeline import Wan22Pipeline
8
+
9
+ from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
10
+ from diffusers import WanImageToVideoPipeline
11
+ from torchvision.transforms import functional as TF
12
+
13
+ from .wan22_14b_model import Wan2214bModel
14
+
15
+ class Wan2214bI2VModel(Wan2214bModel):
16
+ arch = "wan22_14b_i2v"
17
+
18
+
19
+ def generate_single_image(
20
+ self,
21
+ pipeline: Wan22Pipeline,
22
+ gen_config: GenerateImageConfig,
23
+ conditional_embeds: PromptEmbeds,
24
+ unconditional_embeds: PromptEmbeds,
25
+ generator: torch.Generator,
26
+ extra: dict,
27
+ ):
28
+
29
+ # todo
30
+ # reactivate progress bar since this is slooooow
31
+ pipeline.set_progress_bar_config(disable=False)
32
+
33
+ num_frames = (
34
+ (gen_config.num_frames - 1) // 4
35
+ ) * 4 + 1 # make sure it is divisible by 4 + 1
36
+ gen_config.num_frames = num_frames
37
+
38
+ height = gen_config.height
39
+ width = gen_config.width
40
+ first_frame_n1p1 = None
41
+ if gen_config.ctrl_img is not None:
42
+ control_img = Image.open(gen_config.ctrl_img).convert("RGB")
43
+
44
+ d = self.get_bucket_divisibility()
45
+
46
+ # make sure they are divisible by d
47
+ height = height // d * d
48
+ width = width // d * d
49
+
50
+ # resize the control image
51
+ control_img = control_img.resize((width, height), Image.LANCZOS)
52
+
53
+ # 5. Prepare latent variables
54
+ # num_channels_latents = self.transformer.config.in_channels
55
+ num_channels_latents = 16
56
+ latents = pipeline.prepare_latents(
57
+ 1,
58
+ num_channels_latents,
59
+ height,
60
+ width,
61
+ gen_config.num_frames,
62
+ torch.float32,
63
+ self.device_torch,
64
+ generator,
65
+ None,
66
+ ).to(self.torch_dtype)
67
+
68
+ first_frame_n1p1 = (
69
+ TF.to_tensor(control_img)
70
+ .unsqueeze(0)
71
+ .to(self.device_torch, dtype=self.torch_dtype)
72
+ * 2.0
73
+ - 1.0
74
+ ) # normalize to [-1, 1]
75
+
76
+ # Add conditioning using the standalone function
77
+ gen_config.latents = add_first_frame_conditioning(
78
+ latent_model_input=latents,
79
+ first_frame=first_frame_n1p1,
80
+ vae=self.vae
81
+ )
82
+
83
+ output = pipeline(
84
+ prompt_embeds=conditional_embeds.text_embeds.to(
85
+ self.device_torch, dtype=self.torch_dtype
86
+ ),
87
+ negative_prompt_embeds=unconditional_embeds.text_embeds.to(
88
+ self.device_torch, dtype=self.torch_dtype
89
+ ),
90
+ height=height,
91
+ width=width,
92
+ num_inference_steps=gen_config.num_inference_steps,
93
+ guidance_scale=gen_config.guidance_scale,
94
+ latents=gen_config.latents,
95
+ num_frames=gen_config.num_frames,
96
+ generator=generator,
97
+ return_dict=False,
98
+ output_type="pil",
99
+ **extra,
100
+ )[0]
101
+
102
+ # shape = [1, frames, channels, height, width]
103
+ batch_item = output[0] # list of pil images
104
+ if gen_config.num_frames > 1:
105
+ return batch_item # return the frames.
106
+ else:
107
+ # get just the first image
108
+ img = batch_item[0]
109
+ return img
110
+
111
+ def get_noise_prediction(
112
+ self,
113
+ latent_model_input: torch.Tensor,
114
+ timestep: torch.Tensor, # 0 to 1000 scale
115
+ text_embeddings: PromptEmbeds,
116
+ batch: DataLoaderBatchDTO,
117
+ **kwargs
118
+ ):
119
+ # videos come in (bs, num_frames, channels, height, width)
120
+ # images come in (bs, channels, height, width)
121
+ with torch.no_grad():
122
+ frames = batch.tensor
123
+ if len(frames.shape) == 4:
124
+ first_frames = frames
125
+ elif len(frames.shape) == 5:
126
+ first_frames = frames[:, 0]
127
+ else:
128
+ raise ValueError(f"Unknown frame shape {frames.shape}")
129
+
130
+ # Add conditioning using the standalone function
131
+ conditioned_latent = add_first_frame_conditioning(
132
+ latent_model_input=latent_model_input,
133
+ first_frame=first_frames,
134
+ vae=self.vae
135
+ )
136
+
137
+ noise_pred = self.model(
138
+ hidden_states=conditioned_latent,
139
+ timestep=timestep,
140
+ encoder_hidden_states=text_embeddings.text_embeds,
141
+ return_dict=False,
142
+ **kwargs
143
+ )[0]
144
+ return noise_pred