diff --git a/README.md b/README.md index da29d55..725d071 100644 --- a/README.md +++ b/README.md @@ -23,6 +23,7 @@ CogVideoX-Fun is a modified pipeline based on the CogVideoX structure, designed We will support quick pull-ups from different platforms, refer to [Quick Start](#quick-start). What's New: +- Retrain the i2v model and add noise to increase the motion amplitude of the video. Upload the control model training code and control model. [ 2024.09.29 ] - Create code! Now supporting Windows and Linux. Supports 2b and 5b models. Supports video generation at any resolution from 256x256x49 to 1024x1024x49. [ 2024.09.18 ] Function: @@ -68,10 +69,10 @@ cd CogVideoX-Fun mkdir models/Diffusion_Transformer mkdir models/Personalized_Model -wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/Diffusion_Transformer/CogVideoX-Fun-2b-InP.tar.gz -O models/Diffusion_Transformer/CogVideoX-Fun-2b-InP.tar.gz +wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP.tar.gz -O models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP.tar.gz cd models/Diffusion_Transformer/ -tar -xvf CogVideoX-Fun-2b-InP.tar.gz +tar -xvf CogVideoX-Fun-V1.1-2b-InP.tar.gz cd ../../ ``` @@ -103,8 +104,8 @@ We'd better place the [weights](#model-zoo) along the specified path: ``` 📦 models/ ├── 📂 Diffusion_Transformer/ -│ ├── 📂 CogVideoX-Fun-2b-InP/ -│ └── 📂 CogVideoX-Fun-5b-InP/ +│ ├── 📂 CogVideoX-Fun-V1.1-2b-InP/ +│ └── 📂 CogVideoX-Fun-V1.1-5b-InP/ ├── 📂 Personalized_Model/ │ └── your trained trainformer model / your trained lora model (for UI load) ``` @@ -112,42 +113,43 @@ We'd better place the [weights](#model-zoo) along the specified path: # Video Result The results displayed are all based on image. -### CogVideoX-Fun-5B +### CogVideoX-Fun-V1.1-5B Resolution-1024
- + - + - + - +
+ Resolution-768
- + - + - + - +
@@ -157,35 +159,89 @@ Resolution-512
- + - + - + - +
-### CogVideoX-Fun-2B +### CogVideoX-Fun-V1.1-5B-Pose + + + + + +
- + Resolution-512 - + Resolution-768 - + Resolution-1024 +
+ + + + + +
+ +### CogVideoX-Fun-V1.1-2B + +Resolution-768 + + + + + + + +
+ + + + + - + +
+ +### CogVideoX-Fun-V1.1-2B-Pose + + + + + + + + + +
+ Resolution-512 + + Resolution-768 + + Resolution-1024 +
+ + + + +
@@ -283,11 +339,22 @@ Then, we run scripts/train.sh. sh scripts/train.sh ``` -For details on setting some parameters, please refer to [Readme Train](scripts/README_TRAIN.md) and [Readme Lora](scripts/README_TRAIN_LORA.md). +For details on setting some parameters, please refer to [Readme Train](scripts/README_TRAIN.md), [Readme Lora](scripts/README_TRAIN_LORA.md) and [Readme Control](scripts/README_TRAIN_CONTROL.md). # Model zoo +V1.1: + +| 名称 | 存储空间 | Hugging Face | Model Scope | 描述 | +|--|--|--|--|--| +| CogVideoX-Fun-V1.1-2b-InP.tar.gz | Before extraction:9.7 GB \/ After extraction: 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. Noise has been added to the reference image, and the amplitude of motion is greater compared to V1.0. | +| CogVideoX-Fun-V1.1-5b-InP.tar.gz | Before extraction:16.0 GB \/ After extraction: 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. Noise has been added to the reference image, and the amplitude of motion is greater compared to V1.0. | +| CogVideoX-Fun-V1.1-2b-Pose.tar.gz | Before extraction:9.7 GB \/ After extraction: 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-Pose) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-Pose) | Our official pose-control video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second.| +| CogVideoX-Fun-V1.1-5b-Pose.tar.gz | Before extraction:16.0 GB \/ After extraction: 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-Pose) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-Pose) | Our official pose-control video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second.| + +V1.0: + | Name | Storage Space | Hugging Face | Model Scope | Description | |--|--|--|--|--| | CogVideoX-Fun-2b-InP.tar.gz | Before extraction:9.7 GB \/ After extraction: 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-2b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-2b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. | diff --git a/README_zh-CN.md b/README_zh-CN.md index 975dd84..0b9da3a 100644 --- a/README_zh-CN.md +++ b/README_zh-CN.md @@ -23,6 +23,7 @@ CogVideoX-Fun是一个基于CogVideoX结构修改后的的pipeline,是一个 我们会逐渐支持从不同平台快速启动,请参阅 [快速启动](#快速启动)。 新特性: +- 重新训练i2v模型,添加Noise,使得视频的运动幅度更大。上传控制模型训练代码与Control模型。[ 2024.09.29 ] - 创建代码!现在支持 Windows 和 Linux。支持2b与5b最大256x256x49到1024x1024x49的任意分辨率的视频生成。[ 2024.09.18 ] 功能概览: @@ -66,10 +67,10 @@ cd CogVideoX-Fun mkdir models/Diffusion_Transformer mkdir models/Personalized_Model -wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/Diffusion_Transformer/CogVideoX-Fun-2b-InP.tar.gz -O models/Diffusion_Transformer/CogVideoX-Fun-2b-InP.tar.gz +wget https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP.tar.gz -O models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP.tar.gz cd models/Diffusion_Transformer/ -tar -xvf CogVideoX-Fun-2b-InP.tar.gz +tar -xvf CogVideoX-Fun-V1.1-2b-InP.tar.gz cd ../../ ``` @@ -101,8 +102,8 @@ Linux 的详细信息: ``` 📦 models/ ├── 📂 Diffusion_Transformer/ -│ ├── 📂 CogVideoX-Fun-2b-InP/ -│ └── 📂 CogVideoX-Fun-5b-InP/ +│ ├── 📂 CogVideoX-Fun-V1.1-2b-InP/ +│ └── 📂 CogVideoX-Fun-V1.1-5b-InP/ ├── 📂 Personalized_Model/ │ └── your trained trainformer model / your trained lora model (for UI load) ``` @@ -110,42 +111,43 @@ Linux 的详细信息: # 视频作品 所展示的结果都是图生视频获得。 -### CogVideoX-Fun-5B +### CogVideoX-Fun-V1.1-5B Resolution-1024
- + - + - + - +
+ Resolution-768
- + - + - + - +
@@ -155,41 +157,92 @@ Resolution-512
- + - + - + - +
-### CogVideoX-Fun-2B +### CogVideoX-Fun-V1.1-5B-Pose + + + + + + + + + + + +
+ Resolution-512 + + Resolution-768 + + Resolution-1024 +
+ + + + + +
+ +### CogVideoX-Fun-V1.1-2B Resolution-768
- + - + - + - +
+### CogVideoX-Fun-V1.1-2B-Pose + + + + + + + + + + + +
+ Resolution-512 + + Resolution-768 + + Resolution-1024 +
+ + + + + +
# 如何使用 @@ -289,6 +342,18 @@ sh scripts/train.sh 关于一些参数的设置细节,可以查看[Readme Train](scripts/README_TRAIN.md)与[Readme Lora](scripts/README_TRAIN_LORA.md) # 模型地址 + +V1.1: + +| 名称 | 存储空间 | Hugging Face | Model Scope | 描述 | +|--|--|--|--|--| +| CogVideoX-Fun-V1.1-2b-InP.tar.gz | 解压前 9.7 GB / 解压后 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-InP) | 官方的图生视频权重。添加了Noise,运动幅度相比于V1.0更大。支持多分辨率(512,768,1024,1280)的视频预测,以49帧、每秒8帧进行训练 | +| CogVideoX-Fun-V1.1-5b-InP.tar.gz | 解压前 16.0GB / 解压后 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-InP) | 官方的图生视频权重。添加了Noise,运动幅度相比于V1.0更大。支持多分辨率(512,768,1024,1280)的视频预测,以49帧、每秒8帧进行训练 | +| CogVideoX-Fun-V1.1-2b-Pose.tar.gz | 解压前 9.7 GB / 解压后 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-Pose) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-Pose) | 官方的姿态控制生视频权重。支持多分辨率(512,768,1024,1280)的视频预测,以49帧、每秒8帧进行训练 | +| CogVideoX-Fun-V1.1-5b-Pose.tar.gz | 解压前 16.0GB / 解压后 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-Pose) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-Pose) | 官方的姿态控制生视频权重。支持多分辨率(512,768,1024,1280)的视频预测,以49帧、每秒8帧进行训练 | + +V1.0: + | 名称 | 存储空间 | Hugging Face | Model Scope | 描述 | |--|--|--|--|--| | CogVideoX-Fun-2b-InP.tar.gz | 解压前 9.7 GB / 解压后 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-2b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-2b-InP) | 官方的图生视频权重。支持多分辨率(512,768,1024,1280)的视频预测,以49帧、每秒8帧进行训练 | @@ -306,4 +371,4 @@ sh scripts/train.sh CogVideoX-2B 模型 (包括其对应的Transformers模块,VAE模块) 根据 [Apache 2.0 协议](LICENSE) 许可证发布。 -CogVideoX-5B 模型(Transformer 模块)在[CogVideoX许可证](https://huggingface.co/THUDM/CogVideoX-5b/blob/main/LICENSE)下发布. \ No newline at end of file +CogVideoX-5B 模型(Transformer 模块)在[CogVideoX许可证](https://huggingface.co/THUDM/CogVideoX-5b/blob/main/LICENSE)下发布. diff --git a/app.py b/app.py index a453e7d..d7206fd 100644 --- a/app.py +++ b/app.py @@ -19,11 +19,14 @@ if __name__ == "__main__": server_port = 7860 # Params below is used when ui_mode = "modelscope" - model_name = "models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" + model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP" + # "Inpaint" or "Control" + model_type = "Inpaint" + # Save dir of this model savedir_sample = "samples" if ui_mode == "modelscope": - demo, controller = ui_modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype) + demo, controller = ui_modelscope(model_name, model_type, savedir_sample, low_gpu_memory_mode, weight_dtype) elif ui_mode == "eas": demo, controller = ui_eas(model_name, savedir_sample) else: diff --git a/cogvideox/api/api.py b/cogvideox/api/api.py index 5ec9395..8940559 100644 --- a/cogvideox/api/api.py +++ b/cogvideox/api/api.py @@ -68,6 +68,20 @@ def save_base64_video(base64_string): return file_path +def save_base64_image(base64_string): + video_data = base64.b64decode(base64_string) + + md5_hash = hashlib.md5(video_data).hexdigest() + filename = f"{md5_hash}.jpg" + + temp_dir = tempfile.gettempdir() + file_path = os.path.join(temp_dir, filename) + + with open(file_path, 'wb') as video_file: + video_file.write(video_data) + + return file_path + def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): @app.post("/cogvideox_fun/infer_forward") def _infer_forward_api( @@ -77,7 +91,7 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): lora_model_path = datas.get('lora_model_path', 'none') lora_alpha_slider = datas.get('lora_alpha_slider', 0.55) prompt_textbox = datas.get('prompt_textbox', None) - negative_prompt_textbox = datas.get('negative_prompt_textbox', 'The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. ') + negative_prompt_textbox = datas.get('negative_prompt_textbox', 'The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. ') sampler_dropdown = datas.get('sampler_dropdown', 'Euler') sample_step_slider = datas.get('sample_step_slider', 30) resize_method = datas.get('resize_method', "Generate by") @@ -93,6 +107,8 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): start_image = datas.get('start_image', None) end_image = datas.get('end_image', None) validation_video = datas.get('validation_video', None) + validation_video_mask = datas.get('validation_video_mask', None) + control_video = datas.get('control_video', None) denoise_strength = datas.get('denoise_strength', 0.70) seed_textbox = datas.get("seed_textbox", 43) @@ -109,6 +125,12 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): if validation_video is not None: validation_video = save_base64_video(validation_video) + if validation_video_mask is not None: + validation_video_mask = save_base64_image(validation_video_mask) + + if control_video is not None: + control_video = save_base64_video(control_video) + try: save_sample_path, comment = controller.generate( "", @@ -131,6 +153,8 @@ def infer_forward_api(_: gr.Blocks, app: FastAPI, controller): start_image, end_image, validation_video, + validation_video_mask, + control_video, denoise_strength, seed_textbox, is_api = True, diff --git a/cogvideox/api/post_infer.py b/cogvideox/api/post_infer.py index bf7e3e5..57f6ffe 100644 --- a/cogvideox/api/post_infer.py +++ b/cogvideox/api/post_infer.py @@ -33,7 +33,7 @@ def post_infer(generation_method, length_slider, url='http://127.0.0.1:7860'): "lora_model_path": "none", "lora_alpha_slider": 0.55, "prompt_textbox": "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.", - "negative_prompt_textbox": "The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. ", + "negative_prompt_textbox": "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. ", "sampler_dropdown": "Euler", "sample_step_slider": 50, "width_slider": 672, diff --git a/cogvideox/data/dataset_image_video.py b/cogvideox/data/dataset_image_video.py index 8d0bb58..032ae9e 100644 --- a/cogvideox/data/dataset_image_video.py +++ b/cogvideox/data/dataset_image_video.py @@ -322,3 +322,225 @@ class ImageVideoDataset(Dataset): return sample + +class ImageVideoControlDataset(Dataset): + def __init__( + self, + ann_path, data_root=None, + video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16, + image_sample_size=512, + video_repeat=0, + text_drop_ratio=-1, + enable_bucket=False, + video_length_drop_start=0.1, + video_length_drop_end=0.9, + enable_inpaint=False, + ): + # Loading annotations from files + print(f"loading annotations from {ann_path} ...") + if ann_path.endswith('.csv'): + with open(ann_path, 'r') as csvfile: + dataset = list(csv.DictReader(csvfile)) + elif ann_path.endswith('.json'): + dataset = json.load(open(ann_path)) + + self.data_root = data_root + + # It's used to balance num of images and videos. + self.dataset = [] + for data in dataset: + if data.get('type', 'image') != 'video': + self.dataset.append(data) + if video_repeat > 0: + for _ in range(video_repeat): + for data in dataset: + if data.get('type', 'image') == 'video': + self.dataset.append(data) + del dataset + + self.length = len(self.dataset) + print(f"data scale: {self.length}") + # TODO: enable bucket training + self.enable_bucket = enable_bucket + self.text_drop_ratio = text_drop_ratio + self.enable_inpaint = enable_inpaint + + self.video_length_drop_start = video_length_drop_start + self.video_length_drop_end = video_length_drop_end + + # Video params + self.video_sample_stride = video_sample_stride + self.video_sample_n_frames = video_sample_n_frames + self.video_sample_size = tuple(video_sample_size) if not isinstance(video_sample_size, int) else (video_sample_size, video_sample_size) + self.video_transforms = transforms.Compose( + [ + transforms.Resize(min(self.video_sample_size)), + transforms.CenterCrop(self.video_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ] + ) + + # Image params + self.image_sample_size = tuple(image_sample_size) if not isinstance(image_sample_size, int) else (image_sample_size, image_sample_size) + self.image_transforms = transforms.Compose([ + transforms.Resize(min(self.image_sample_size)), + transforms.CenterCrop(self.image_sample_size), + transforms.ToTensor(), + transforms.Normalize([0.5, 0.5, 0.5],[0.5, 0.5, 0.5]) + ]) + + self.larger_side_of_image_and_video = max(min(self.image_sample_size), min(self.video_sample_size)) + + def get_batch(self, idx): + data_info = self.dataset[idx % len(self.dataset)] + video_id, control_video_id, text = data_info['file_path'], data_info['control_file_path'], data_info['text'] + + if data_info.get('type', 'image')=='video': + if self.data_root is None: + video_dir = video_id + else: + video_dir = os.path.join(self.data_root, video_id) + + with VideoReader_contextmanager(video_dir, num_threads=2) as video_reader: + min_sample_n_frames = min( + self.video_sample_n_frames, + int(len(video_reader) * (self.video_length_drop_end - self.video_length_drop_start) // self.video_sample_stride) + ) + if min_sample_n_frames == 0: + raise ValueError(f"No Frames in video.") + + video_length = int(self.video_length_drop_end * len(video_reader)) + clip_length = min(video_length, (min_sample_n_frames - 1) * self.video_sample_stride + 1) + start_idx = random.randint(int(self.video_length_drop_start * video_length), video_length - clip_length) if video_length != clip_length else 0 + batch_index = np.linspace(start_idx, start_idx + clip_length - 1, min_sample_n_frames, dtype=int) + + try: + sample_args = (video_reader, batch_index) + pixel_values = func_timeout( + VIDEO_READER_TIMEOUT, get_video_reader_batch, args=sample_args + ) + resized_frames = [] + for i in range(len(pixel_values)): + frame = pixel_values[i] + resized_frame = resize_frame(frame, self.larger_side_of_image_and_video) + resized_frames.append(resized_frame) + pixel_values = np.array(resized_frames) + except FunctionTimedOut: + raise ValueError(f"Read {idx} timeout.") + except Exception as e: + raise ValueError(f"Failed to extract frames from video. Error is {e}.") + + if not self.enable_bucket: + pixel_values = torch.from_numpy(pixel_values).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + del video_reader + else: + pixel_values = pixel_values + + if not self.enable_bucket: + pixel_values = self.video_transforms(pixel_values) + + # Random use no text generation + if random.random() < self.text_drop_ratio: + text = '' + + if self.data_root is None: + control_video_id = control_video_id + else: + control_video_id = os.path.join(self.data_root, control_video_id) + + with VideoReader_contextmanager(control_video_id, num_threads=2) as control_video_reader: + try: + sample_args = (control_video_reader, batch_index) + control_pixel_values = func_timeout( + VIDEO_READER_TIMEOUT, get_video_reader_batch, args=sample_args + ) + resized_frames = [] + for i in range(len(control_pixel_values)): + frame = control_pixel_values[i] + resized_frame = resize_frame(frame, self.larger_side_of_image_and_video) + resized_frames.append(resized_frame) + control_pixel_values = np.array(resized_frames) + except FunctionTimedOut: + raise ValueError(f"Read {idx} timeout.") + except Exception as e: + raise ValueError(f"Failed to extract frames from video. Error is {e}.") + + if not self.enable_bucket: + control_pixel_values = torch.from_numpy(control_pixel_values).permute(0, 3, 1, 2).contiguous() + control_pixel_values = control_pixel_values / 255. + del control_video_reader + else: + control_pixel_values = control_pixel_values + + if not self.enable_bucket: + control_pixel_values = self.video_transforms(control_pixel_values) + return pixel_values, control_pixel_values, text, "video" + else: + image_path, text = data_info['file_path'], data_info['text'] + if self.data_root is not None: + image_path = os.path.join(self.data_root, image_path) + image = Image.open(image_path).convert('RGB') + if not self.enable_bucket: + image = self.image_transforms(image).unsqueeze(0) + else: + image = np.expand_dims(np.array(image), 0) + + if random.random() < self.text_drop_ratio: + text = '' + + if self.data_root is None: + control_image_id = control_image_id + else: + control_image_id = os.path.join(self.data_root, control_image_id) + + control_image = Image.open(control_image_id).convert('RGB') + if not self.enable_bucket: + control_image = self.image_transforms(control_image).unsqueeze(0) + else: + control_image = np.expand_dims(np.array(control_image), 0) + return image, control_image, text, 'image' + + def __len__(self): + return self.length + + def __getitem__(self, idx): + data_info = self.dataset[idx % len(self.dataset)] + data_type = data_info.get('type', 'image') + while True: + sample = {} + try: + data_info_local = self.dataset[idx % len(self.dataset)] + data_type_local = data_info_local.get('type', 'image') + if data_type_local != data_type: + raise ValueError("data_type_local != data_type") + + pixel_values, control_pixel_values, name, data_type = self.get_batch(idx) + sample["pixel_values"] = pixel_values + sample["control_pixel_values"] = control_pixel_values + sample["text"] = name + sample["data_type"] = data_type + sample["idx"] = idx + + if len(sample) > 0: + break + except Exception as e: + print(e, self.dataset[idx % len(self.dataset)]) + idx = random.randint(0, self.length-1) + + if self.enable_inpaint and not self.enable_bucket: + mask = get_random_mask(pixel_values.size()) + mask_pixel_values = pixel_values * (1 - mask) + torch.ones_like(pixel_values) * -1 * mask + sample["mask_pixel_values"] = mask_pixel_values + sample["mask"] = mask + + clip_pixel_values = sample["pixel_values"][0].permute(1, 2, 0).contiguous() + clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255 + sample["clip_pixel_values"] = clip_pixel_values + + ref_pixel_values = sample["pixel_values"][0].unsqueeze(0) + if (mask == 1).all(): + ref_pixel_values = torch.ones_like(ref_pixel_values) * -1 + sample["ref_pixel_values"] = ref_pixel_values + + return sample diff --git a/cogvideox/models/transformer3d.py b/cogvideox/models/transformer3d.py index b80af91..88c8013 100644 --- a/cogvideox/models/transformer3d.py +++ b/cogvideox/models/transformer3d.py @@ -277,6 +277,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): spatial_interpolation_scale: float = 1.875, temporal_interpolation_scale: float = 1.0, use_rotary_positional_embeddings: bool = False, + add_noise_in_inpaint_model: bool = False, ): super().__init__() inner_dim = num_attention_heads * attention_head_dim @@ -452,6 +453,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): timestep: Union[int, float, torch.LongTensor], timestep_cond: Optional[torch.Tensor] = None, inpaint_latents: Optional[torch.Tensor] = None, + control_latents: Optional[torch.Tensor] = None, image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, return_dict: bool = True, ): @@ -470,6 +472,8 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): # 2. Patch embedding if inpaint_latents is not None: hidden_states = torch.concat([hidden_states, inpaint_latents], 2) + if control_latents is not None: + hidden_states = torch.concat([hidden_states, control_latents], 2) hidden_states = self.patch_embed(encoder_hidden_states, hidden_states) # 3. Position embedding diff --git a/cogvideox/pipeline/pipeline_cogvideox_control.py b/cogvideox/pipeline/pipeline_cogvideox_control.py new file mode 100644 index 0000000..4e82c84 --- /dev/null +++ b/cogvideox/pipeline/pipeline_cogvideox_control.py @@ -0,0 +1,843 @@ +# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team. +# All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import inspect +import math +from dataclasses import dataclass +from typing import Callable, Dict, List, Optional, Tuple, Union + +import torch +import torch.nn.functional as F +from einops import rearrange +from transformers import T5EncoderModel, T5Tokenizer + +from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback +from diffusers.models import AutoencoderKLCogVideoX, CogVideoXTransformer3DModel +from diffusers.models.embeddings import get_3d_rotary_pos_embed +from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from diffusers.schedulers import CogVideoXDDIMScheduler, CogVideoXDPMScheduler +from diffusers.utils import BaseOutput, logging, replace_example_docstring +from diffusers.utils.torch_utils import randn_tensor +from diffusers.video_processor import VideoProcessor +from diffusers.image_processor import VaeImageProcessor +from einops import rearrange + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +EXAMPLE_DOC_STRING = """ + Examples: + ```python + >>> import torch + >>> from diffusers import CogVideoX_Fun_Pipeline + >>> from diffusers.utils import export_to_video + + >>> # Models: "THUDM/CogVideoX-2b" or "THUDM/CogVideoX-5b" + >>> pipe = CogVideoX_Fun_Pipeline.from_pretrained("THUDM/CogVideoX-2b", torch_dtype=torch.float16).to("cuda") + >>> prompt = ( + ... "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. " + ... "The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other " + ... "pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, " + ... "casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. " + ... "The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical " + ... "atmosphere of this unique musical performance." + ... ) + >>> video = pipe(prompt=prompt, guidance_scale=6, num_inference_steps=50).frames[0] + >>> export_to_video(video, "output.mp4", fps=8) + ``` +""" + + +# Similar to diffusers.pipelines.hunyuandit.pipeline_hunyuandit.get_resize_crop_region_for_grid +def get_resize_crop_region_for_grid(src, tgt_width, tgt_height): + tw = tgt_width + th = tgt_height + h, w = src + r = h / w + if r > (th / tw): + resize_height = th + resize_width = int(round(th / h * w)) + else: + resize_width = tw + resize_height = int(round(tw / w * h)) + + crop_top = int(round((th - resize_height) / 2.0)) + crop_left = int(round((tw - resize_width) / 2.0)) + + return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width) + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps +def retrieve_timesteps( + scheduler, + num_inference_steps: Optional[int] = None, + device: Optional[Union[str, torch.device]] = None, + timesteps: Optional[List[int]] = None, + sigmas: Optional[List[float]] = None, + **kwargs, +): + """ + Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles + custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. + + Args: + scheduler (`SchedulerMixin`): + The scheduler to get timesteps from. + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` + must be `None`. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + timesteps (`List[int]`, *optional*): + Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, + `num_inference_steps` and `sigmas` must be `None`. + sigmas (`List[float]`, *optional*): + Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, + `num_inference_steps` and `timesteps` must be `None`. + + Returns: + `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the + second element is the number of inference steps. + """ + if timesteps is not None and sigmas is not None: + raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") + if timesteps is not None: + accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accepts_timesteps: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" timestep schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + elif sigmas is not None: + accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accept_sigmas: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" sigmas schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + else: + scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) + timesteps = scheduler.timesteps + return timesteps, num_inference_steps + + +@dataclass +class CogVideoX_Fun_PipelineOutput(BaseOutput): + r""" + Output class for CogVideo pipelines. + + Args: + video (`torch.Tensor`, `np.ndarray`, or List[List[PIL.Image.Image]]): + List of video outputs - It can be a nested list of length `batch_size,` with each sub-list containing + denoised PIL image sequences of length `num_frames.` It can also be a NumPy array or Torch tensor of shape + `(batch_size, num_frames, channels, height, width)`. + """ + + videos: torch.Tensor + + +class CogVideoX_Fun_Pipeline_Control(DiffusionPipeline): + r""" + Pipeline for text-to-video generation using CogVideoX. + + This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods the + library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.) + + Args: + vae ([`AutoencoderKL`]): + Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations. + text_encoder ([`T5EncoderModel`]): + Frozen text-encoder. CogVideoX_Fun uses + [T5](https://huggingface.co/docs/transformers/model_doc/t5#transformers.T5EncoderModel); specifically the + [t5-v1_1-xxl](https://huggingface.co/PixArt-alpha/PixArt-alpha/tree/main/t5-v1_1-xxl) variant. + tokenizer (`T5Tokenizer`): + Tokenizer of class + [T5Tokenizer](https://huggingface.co/docs/transformers/model_doc/t5#transformers.T5Tokenizer). + transformer ([`CogVideoXTransformer3DModel`]): + A text conditioned `CogVideoXTransformer3DModel` to denoise the encoded video latents. + scheduler ([`SchedulerMixin`]): + A scheduler to be used in combination with `transformer` to denoise the encoded video latents. + """ + + _optional_components = [] + model_cpu_offload_seq = "text_encoder->vae->transformer->vae" + + _callback_tensor_inputs = [ + "latents", + "prompt_embeds", + "negative_prompt_embeds", + ] + + def __init__( + self, + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + vae: AutoencoderKLCogVideoX, + transformer: CogVideoXTransformer3DModel, + scheduler: Union[CogVideoXDDIMScheduler, CogVideoXDPMScheduler], + ): + super().__init__() + + self.register_modules( + tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer, scheduler=scheduler + ) + self.vae_scale_factor_spatial = ( + 2 ** (len(self.vae.config.block_out_channels) - 1) if hasattr(self, "vae") and self.vae is not None else 8 + ) + self.vae_scale_factor_temporal = ( + self.vae.config.temporal_compression_ratio if hasattr(self, "vae") and self.vae is not None else 4 + ) + + self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial) + + self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor) + self.mask_processor = VaeImageProcessor( + vae_scale_factor=self.vae_scale_factor, do_normalize=False, do_binarize=True, do_convert_grayscale=True + ) + + def _get_t5_prompt_embeds( + self, + prompt: Union[str, List[str]] = None, + num_videos_per_prompt: int = 1, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + ): + device = device or self._execution_device + dtype = dtype or self.text_encoder.dtype + + prompt = [prompt] if isinstance(prompt, str) else prompt + batch_size = len(prompt) + + text_inputs = self.tokenizer( + prompt, + padding="max_length", + max_length=max_sequence_length, + truncation=True, + add_special_tokens=True, + return_tensors="pt", + ) + text_input_ids = text_inputs.input_ids + untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids + + if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids): + removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because `max_sequence_length` is set to " + f" {max_sequence_length} tokens: {removed_text}" + ) + + prompt_embeds = self.text_encoder(text_input_ids.to(device))[0] + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + # duplicate text embeddings for each generation per prompt, using mps friendly method + _, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1) + + return prompt_embeds + + def encode_prompt( + self, + prompt: Union[str, List[str]], + negative_prompt: Optional[Union[str, List[str]]] = None, + do_classifier_free_guidance: bool = True, + num_videos_per_prompt: int = 1, + prompt_embeds: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[torch.Tensor] = None, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + ): + r""" + Encodes the prompt into text encoder hidden states. + + Args: + prompt (`str` or `List[str]`, *optional*): + prompt to be encoded + negative_prompt (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is + less than `1`). + do_classifier_free_guidance (`bool`, *optional*, defaults to `True`): + Whether to use classifier free guidance or not. + num_videos_per_prompt (`int`, *optional*, defaults to 1): + Number of videos that should be generated per prompt. torch device to place the resulting embeddings on + prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. + device: (`torch.device`, *optional*): + torch device + dtype: (`torch.dtype`, *optional*): + torch dtype + """ + device = device or self._execution_device + + prompt = [prompt] if isinstance(prompt, str) else prompt + if prompt is not None: + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if prompt_embeds is None: + prompt_embeds = self._get_t5_prompt_embeds( + prompt=prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + if do_classifier_free_guidance and negative_prompt_embeds is None: + negative_prompt = negative_prompt or "" + negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt + + if prompt is not None and type(prompt) is not type(negative_prompt): + raise TypeError( + f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" + f" {type(prompt)}." + ) + elif batch_size != len(negative_prompt): + raise ValueError( + f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" + f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" + " the batch size of `prompt`." + ) + + negative_prompt_embeds = self._get_t5_prompt_embeds( + prompt=negative_prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + return prompt_embeds, negative_prompt_embeds + + def prepare_latents( + self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, latents=None + ): + shape = ( + batch_size, + (num_frames - 1) // self.vae_scale_factor_temporal + 1, + num_channels_latents, + height // self.vae_scale_factor_spatial, + width // self.vae_scale_factor_spatial, + ) + if isinstance(generator, list) and len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}. Make sure the batch size matches the length of the generators." + ) + + if latents is None: + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + else: + latents = latents.to(device) + + # scale the initial noise by the standard deviation required by the scheduler + latents = latents * self.scheduler.init_noise_sigma + return latents + + def prepare_control_latents( + self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance + ): + # resize the mask to latents shape as we concatenate the mask to the latents + # we do that before converting to dtype to avoid breaking in case we're using cpu_offload + # and half precision + + if mask is not None: + mask = mask.to(device=device, dtype=self.vae.dtype) + bs = 1 + new_mask = [] + for i in range(0, mask.shape[0], bs): + mask_bs = mask[i : i + bs] + mask_bs = self.vae.encode(mask_bs)[0] + mask_bs = mask_bs.mode() + new_mask.append(mask_bs) + mask = torch.cat(new_mask, dim = 0) + mask = mask * self.vae.config.scaling_factor + + if masked_image is not None: + masked_image = masked_image.to(device=device, dtype=self.vae.dtype) + bs = 1 + new_mask_pixel_values = [] + for i in range(0, masked_image.shape[0], bs): + mask_pixel_values_bs = masked_image[i : i + bs] + mask_pixel_values_bs = self.vae.encode(mask_pixel_values_bs)[0] + mask_pixel_values_bs = mask_pixel_values_bs.mode() + new_mask_pixel_values.append(mask_pixel_values_bs) + masked_image_latents = torch.cat(new_mask_pixel_values, dim = 0) + masked_image_latents = masked_image_latents * self.vae.config.scaling_factor + else: + masked_image_latents = None + + return mask, masked_image_latents + + def decode_latents(self, latents: torch.Tensor) -> torch.Tensor: + latents = latents.permute(0, 2, 1, 3, 4) # [batch_size, num_channels, num_frames, height, width] + latents = 1 / self.vae.config.scaling_factor * latents + + frames = self.vae.decode(latents).sample + frames = (frames / 2 + 0.5).clamp(0, 1) + # we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16 + frames = frames.cpu().float().numpy() + return frames + + # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_extra_step_kwargs + def prepare_extra_step_kwargs(self, generator, eta): + # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature + # eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers. + # eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502 + # and should be between [0, 1] + + accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys()) + extra_step_kwargs = {} + if accepts_eta: + extra_step_kwargs["eta"] = eta + + # check if the scheduler accepts generator + accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys()) + if accepts_generator: + extra_step_kwargs["generator"] = generator + return extra_step_kwargs + + # Copied from diffusers.pipelines.latte.pipeline_latte.LattePipeline.check_inputs + def check_inputs( + self, + prompt, + height, + width, + negative_prompt, + callback_on_step_end_tensor_inputs, + prompt_embeds=None, + negative_prompt_embeds=None, + ): + if height % 8 != 0 or width % 8 != 0: + raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.") + + if callback_on_step_end_tensor_inputs is not None and not all( + k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs + ): + raise ValueError( + f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}" + ) + if prompt is not None and prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to" + " only forward one of the two." + ) + elif prompt is None and prompt_embeds is None: + raise ValueError( + "Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined." + ) + elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): + raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") + + if prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `negative_prompt_embeds`:" + f" {negative_prompt_embeds}. Please make sure to only forward one of the two." + ) + + if negative_prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:" + f" {negative_prompt_embeds}. Please make sure to only forward one of the two." + ) + + if prompt_embeds is not None and negative_prompt_embeds is not None: + if prompt_embeds.shape != negative_prompt_embeds.shape: + raise ValueError( + "`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but" + f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`" + f" {negative_prompt_embeds.shape}." + ) + + def fuse_qkv_projections(self) -> None: + r"""Enables fused QKV projections.""" + self.fusing_transformer = True + self.transformer.fuse_qkv_projections() + + def unfuse_qkv_projections(self) -> None: + r"""Disable QKV projection fusion if enabled.""" + if not self.fusing_transformer: + logger.warning("The Transformer was not initially fused for QKV projections. Doing nothing.") + else: + self.transformer.unfuse_qkv_projections() + self.fusing_transformer = False + + def _prepare_rotary_positional_embeddings( + self, + height: int, + width: int, + num_frames: int, + device: torch.device, + ) -> Tuple[torch.Tensor, torch.Tensor]: + grid_height = height // (self.vae_scale_factor_spatial * self.transformer.config.patch_size) + grid_width = width // (self.vae_scale_factor_spatial * self.transformer.config.patch_size) + base_size_width = 720 // (self.vae_scale_factor_spatial * self.transformer.config.patch_size) + base_size_height = 480 // (self.vae_scale_factor_spatial * self.transformer.config.patch_size) + + grid_crops_coords = get_resize_crop_region_for_grid( + (grid_height, grid_width), base_size_width, base_size_height + ) + freqs_cos, freqs_sin = get_3d_rotary_pos_embed( + embed_dim=self.transformer.config.attention_head_dim, + crops_coords=grid_crops_coords, + grid_size=(grid_height, grid_width), + temporal_size=num_frames, + use_real=True, + ) + + freqs_cos = freqs_cos.to(device=device) + freqs_sin = freqs_sin.to(device=device) + return freqs_cos, freqs_sin + + @property + def guidance_scale(self): + return self._guidance_scale + + @property + def num_timesteps(self): + return self._num_timesteps + + @property + def interrupt(self): + return self._interrupt + + # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.StableDiffusionImg2ImgPipeline.get_timesteps + def get_timesteps(self, num_inference_steps, strength, device): + # get the original timestep using init_timestep + init_timestep = min(int(num_inference_steps * strength), num_inference_steps) + + t_start = max(num_inference_steps - init_timestep, 0) + timesteps = self.scheduler.timesteps[t_start * self.scheduler.order :] + + return timesteps, num_inference_steps - t_start + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: Optional[Union[str, List[str]]] = None, + negative_prompt: Optional[Union[str, List[str]]] = None, + height: int = 480, + width: int = 720, + video: Union[torch.FloatTensor] = None, + control_video: Union[torch.FloatTensor] = None, + num_frames: int = 49, + num_inference_steps: int = 50, + timesteps: Optional[List[int]] = None, + guidance_scale: float = 6, + use_dynamic_cfg: bool = False, + num_videos_per_prompt: int = 1, + eta: float = 0.0, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.FloatTensor] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + output_type: str = "numpy", + return_dict: bool = False, + callback_on_step_end: Optional[ + Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks] + ] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + max_sequence_length: int = 226, + comfyui_progressbar: bool = False, + ) -> Union[CogVideoX_Fun_PipelineOutput, Tuple]: + """ + Function invoked when calling the pipeline for generation. + + Args: + prompt (`str` or `List[str]`, *optional*): + The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`. + instead. + negative_prompt (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is + less than `1`). + height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The height in pixels of the generated image. This is set to 1024 by default for the best results. + width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The width in pixels of the generated image. This is set to 1024 by default for the best results. + num_frames (`int`, defaults to `48`): + Number of frames to generate. Must be divisible by self.vae_scale_factor_temporal. Generated video will + contain 1 extra frame because CogVideoX_Fun is conditioned with (num_seconds * fps + 1) frames where + num_seconds is 6 and fps is 4. However, since videos can be saved at any fps, the only condition that + needs to be satisfied is that of divisibility mentioned above. + num_inference_steps (`int`, *optional*, defaults to 50): + The number of denoising steps. More denoising steps usually lead to a higher quality image at the + expense of slower inference. + timesteps (`List[int]`, *optional*): + Custom timesteps to use for the denoising process with schedulers which support a `timesteps` argument + in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is + passed will be used. Must be in descending order. + guidance_scale (`float`, *optional*, defaults to 7.0): + Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598). + `guidance_scale` is defined as `w` of equation 2. of [Imagen + Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale > + 1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`, + usually at the expense of lower image quality. + num_videos_per_prompt (`int`, *optional*, defaults to 1): + The number of videos to generate per prompt. + generator (`torch.Generator` or `List[torch.Generator]`, *optional*): + One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html) + to make generation deterministic. + latents (`torch.FloatTensor`, *optional*): + Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image + generation. Can be used to tweak the same generation with different prompts. If not provided, a latents + tensor will ge generated by sampling using the supplied random `generator`. + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. + output_type (`str`, *optional*, defaults to `"pil"`): + The output format of the generate image. Choose between + [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] instead + of a plain tuple. + callback_on_step_end (`Callable`, *optional*): + A function that calls at the end of each denoising steps during the inference. The function is called + with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int, + callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by + `callback_on_step_end_tensor_inputs`. + callback_on_step_end_tensor_inputs (`List`, *optional*): + The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list + will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the + `._callback_tensor_inputs` attribute of your pipeline class. + max_sequence_length (`int`, defaults to `226`): + Maximum sequence length in encoded prompt. Must be consistent with + `self.transformer.config.max_text_seq_length` otherwise may lead to poor results. + + Examples: + + Returns: + [`~pipelines.cogvideo.pipeline_cogvideox.CogVideoX_Fun_PipelineOutput`] or `tuple`: + [`~pipelines.cogvideo.pipeline_cogvideox.CogVideoX_Fun_PipelineOutput`] if `return_dict` is True, otherwise a + `tuple`. When returning a tuple, the first element is a list with the generated images. + """ + + if num_frames > 49: + raise ValueError( + "The number of frames must be less than 49 for now due to static positional embeddings. This will be updated in the future to remove this limitation." + ) + + if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)): + callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs + + height = height or self.transformer.config.sample_size * self.vae_scale_factor_spatial + width = width or self.transformer.config.sample_size * self.vae_scale_factor_spatial + num_videos_per_prompt = 1 + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + height, + width, + negative_prompt, + callback_on_step_end_tensor_inputs, + prompt_embeds, + negative_prompt_embeds, + ) + self._guidance_scale = guidance_scale + self._interrupt = False + + # 2. Default call parameters + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + device = self._execution_device + + # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2) + # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1` + # corresponds to doing no classifier free guidance. + do_classifier_free_guidance = guidance_scale > 1.0 + + # 3. Encode input prompt + prompt_embeds, negative_prompt_embeds = self.encode_prompt( + prompt, + negative_prompt, + do_classifier_free_guidance, + num_videos_per_prompt=num_videos_per_prompt, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + max_sequence_length=max_sequence_length, + device=device, + ) + if do_classifier_free_guidance: + prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) + + # 4. Prepare timesteps + timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps) + self._num_timesteps = len(timesteps) + if comfyui_progressbar: + from comfy.utils import ProgressBar + pbar = ProgressBar(num_inference_steps + 2) + + # 5. Prepare latents. + latent_channels = self.vae.config.latent_channels + latents = self.prepare_latents( + batch_size * num_videos_per_prompt, + latent_channels, + num_frames, + height, + width, + prompt_embeds.dtype, + device, + generator, + latents, + ) + if comfyui_progressbar: + pbar.update(1) + + if control_video is not None: + video_length = control_video.shape[2] + control_video = self.image_processor.preprocess(rearrange(control_video, "b c f h w -> (b f) c h w"), height=height, width=width) + control_video = control_video.to(dtype=torch.float32) + control_video = rearrange(control_video, "(b f) c h w -> b c f h w", f=video_length) + else: + control_video = None + control_video_latents = self.prepare_control_latents( + None, + control_video, + batch_size, + height, + width, + prompt_embeds.dtype, + device, + generator, + do_classifier_free_guidance + )[1] + control_video_latents_input = ( + torch.cat([control_video_latents] * 2) if do_classifier_free_guidance else control_video_latents + ) + control_latents = rearrange(control_video_latents_input, "b c f h w -> b f c h w") + + if comfyui_progressbar: + pbar.update(1) + + # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) + + # 7. Create rotary embeds if required + image_rotary_emb = ( + self._prepare_rotary_positional_embeddings(height, width, latents.size(1), device) + if self.transformer.config.use_rotary_positional_embeddings + else None + ) + + # 8. Denoising loop + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + + with self.progress_bar(total=num_inference_steps) as progress_bar: + # for DPM-solver++ + old_pred_original_sample = None + for i, t in enumerate(timesteps): + if self.interrupt: + continue + + latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents + latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) + + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timestep = t.expand(latent_model_input.shape[0]) + + # predict noise model_output + noise_pred = self.transformer( + hidden_states=latent_model_input, + encoder_hidden_states=prompt_embeds, + timestep=timestep, + image_rotary_emb=image_rotary_emb, + return_dict=False, + control_latents=control_latents, + )[0] + noise_pred = noise_pred.float() + + # perform guidance + if use_dynamic_cfg: + self._guidance_scale = 1 + guidance_scale * ( + (1 - math.cos(math.pi * ((num_inference_steps - t.item()) / num_inference_steps) ** 5.0)) / 2 + ) + if do_classifier_free_guidance: + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond) + + # compute the previous noisy sample x_t -> x_t-1 + if not isinstance(self.scheduler, CogVideoXDPMScheduler): + latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0] + else: + latents, old_pred_original_sample = self.scheduler.step( + noise_pred, + old_pred_original_sample, + t, + timesteps[i - 1] if i > 0 else None, + latents, + **extra_step_kwargs, + return_dict=False, + ) + latents = latents.to(prompt_embeds.dtype) + + # call the callback, if provided + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds) + + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update() + if comfyui_progressbar: + pbar.update(1) + + if output_type == "numpy": + video = self.decode_latents(latents) + elif not output_type == "latent": + video = self.decode_latents(latents) + video = self.video_processor.postprocess_video(video=video, output_type=output_type) + else: + video = latents + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + video = torch.from_numpy(video) + + return CogVideoX_Fun_PipelineOutput(videos=video) diff --git a/cogvideox/pipeline/pipeline_cogvideox_inpaint.py b/cogvideox/pipeline/pipeline_cogvideox_inpaint.py index d86160d..a2dfaa7 100644 --- a/cogvideox/pipeline/pipeline_cogvideox_inpaint.py +++ b/cogvideox/pipeline/pipeline_cogvideox_inpaint.py @@ -177,6 +177,19 @@ def resize_mask(mask, latent, process_first_frame_only=True): return resized_mask +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 + + @dataclass class CogVideoX_Fun_PipelineOutput(BaseOutput): r""" @@ -444,7 +457,7 @@ class CogVideoX_Fun_Pipeline_Inpaint(DiffusionPipeline): return outputs def prepare_mask_latents( - self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance + self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance, noise_aug_strength ): # resize the mask to latents shape as we concatenate the mask to the latents # we do that before converting to dtype to avoid breaking in case we're using cpu_offload @@ -463,6 +476,8 @@ class CogVideoX_Fun_Pipeline_Inpaint(DiffusionPipeline): mask = mask * self.vae.config.scaling_factor if masked_image is not None: + if self.transformer.config.add_noise_in_inpaint_model: + masked_image = add_noise_to_reference_video(masked_image, ratio=noise_aug_strength) masked_image = masked_image.to(device=device, dtype=self.vae.dtype) bs = 1 new_mask_pixel_values = [] @@ -650,6 +665,7 @@ class CogVideoX_Fun_Pipeline_Inpaint(DiffusionPipeline): callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 226, strength: float = 1, + noise_aug_strength: float = 0.0563, comfyui_progressbar: bool = False, ) -> Union[CogVideoX_Fun_PipelineOutput, Tuple]: """ @@ -866,6 +882,7 @@ class CogVideoX_Fun_Pipeline_Inpaint(DiffusionPipeline): device, generator, do_classifier_free_guidance, + noise_aug_strength=noise_aug_strength, ) mask_latents = resize_mask(1 - mask_condition, masked_video_latents) mask_latents = mask_latents.to(masked_video_latents.device) * self.vae.config.scaling_factor diff --git a/cogvideox/ui/ui.py b/cogvideox/ui/ui.py index f556ec9..04aef80 100644 --- a/cogvideox/ui/ui.py +++ b/cogvideox/ui/ui.py @@ -30,6 +30,8 @@ from cogvideox.data.bucket_sampler import ASPECT_RATIO_512, get_closest_ratio from cogvideox.models.autoencoder_magvit import AutoencoderKLCogVideoX from cogvideox.models.transformer3d import CogVideoXTransformer3DModel from cogvideox.pipeline.pipeline_cogvideox import CogVideoX_Fun_Pipeline +from cogvideox.pipeline.pipeline_cogvideox_control import \ + CogVideoX_Fun_Pipeline_Control from cogvideox.pipeline.pipeline_cogvideox_inpaint import \ CogVideoX_Fun_Pipeline_Inpaint from cogvideox.utils.lora_utils import merge_lora, unmerge_lora @@ -58,7 +60,7 @@ css = """ } """ -class CogVideoX_I2VController: +class CogVideoX_Fun_Controller: def __init__(self, low_gpu_memory_mode, weight_dtype): # config dirs self.basedir = os.getcwd() @@ -68,6 +70,7 @@ class CogVideoX_I2VController: self.personalized_model_dir = os.path.join(self.basedir, "models", "Personalized_Model") self.savedir = os.path.join(self.basedir, "samples", datetime.now().strftime("Gradio-%Y-%m-%dT%H-%M-%S")) self.savedir_sample = os.path.join(self.savedir, "sample") + self.model_type = "Inpaint" os.makedirs(self.savedir, exist_ok=True) self.diffusion_transformer_list = [] @@ -102,6 +105,9 @@ class CogVideoX_I2VController: personalized_model_list = sorted(glob(os.path.join(self.personalized_model_dir, "*.safetensors"))) self.personalized_model_list = [os.path.basename(p) for p in personalized_model_list] + def update_model_type(self, model_type): + self.model_type = model_type + def update_diffusion_transformer(self, diffusion_transformer_dropdown): print("Update diffusion transformer") if diffusion_transformer_dropdown == "none": @@ -118,16 +124,25 @@ class CogVideoX_I2VController: ).to(self.weight_dtype) # Get pipeline - if self.transformer.config.in_channels != self.vae.config.latent_channels: - self.pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( - diffusion_transformer_dropdown, - vae=self.vae, - transformer=self.transformer, - scheduler=scheduler_dict["Euler"].from_pretrained(diffusion_transformer_dropdown, subfolder="scheduler"), - torch_dtype=self.weight_dtype - ) + if self.model_type == "Inpaint": + if self.transformer.config.in_channels != self.vae.config.latent_channels: + self.pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( + diffusion_transformer_dropdown, + vae=self.vae, + transformer=self.transformer, + scheduler=scheduler_dict["Euler"].from_pretrained(diffusion_transformer_dropdown, subfolder="scheduler"), + torch_dtype=self.weight_dtype + ) + else: + self.pipeline = CogVideoX_Fun_Pipeline.from_pretrained( + diffusion_transformer_dropdown, + vae=self.vae, + transformer=self.transformer, + scheduler=scheduler_dict["Euler"].from_pretrained(diffusion_transformer_dropdown, subfolder="scheduler"), + torch_dtype=self.weight_dtype + ) else: - self.pipeline = CogVideoX_Fun_Pipeline.from_pretrained( + self.pipeline = CogVideoX_Fun_Pipeline_Control.from_pretrained( diffusion_transformer_dropdown, vae=self.vae, transformer=self.transformer, @@ -191,6 +206,8 @@ class CogVideoX_I2VController: start_image, end_image, validation_video, + validation_video_mask, + control_video, denoise_strength, seed_textbox, is_api = False, @@ -208,20 +225,34 @@ class CogVideoX_I2VController: if self.lora_model_path != lora_model_dropdown: print("Update lora model") self.update_lora_model(lora_model_dropdown) - + + if control_video is not None and self.model_type == "Inpaint": + if is_api: + return "", f"If specifying the control video, please set the model_type == \"Control\". " + else: + raise gr.Error(f"If specifying the control video, please set the model_type == \"Control\". ") + + if control_video is None and self.model_type == "Control": + if is_api: + return "", f"If set the model_type == \"Control\", please specifying the control video. " + else: + raise gr.Error(f"If set the model_type == \"Control\", please specifying the control video. ") + if resize_method == "Resize according to Reference": - if start_image is None and validation_video is None: + if start_image is None and validation_video is None and control_video is None: if is_api: return "", f"Please upload an image when using \"Resize according to Reference\"." else: raise gr.Error(f"Please upload an image when using \"Resize according to Reference\".") aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} - - if validation_video is not None: - original_width, original_height = Image.fromarray(cv2.VideoCapture(validation_video).read()[1]).size + if self.model_type == "Inpaint": + if validation_video is not None: + original_width, original_height = Image.fromarray(cv2.VideoCapture(validation_video).read()[1]).size + else: + original_width, original_height = start_image[0].size if type(start_image) is list else Image.open(start_image).size else: - original_width, original_height = start_image[0].size if type(start_image) is list else Image.open(start_image).size + original_width, original_height = Image.fromarray(cv2.VideoCapture(control_video).read()[1]).size closest_size, closest_ratio = get_closest_ratio(original_height, original_width, ratios=aspect_ratio_sample_size) height_slider, width_slider = [int(x / 16) * 16 for x in closest_size] @@ -255,75 +286,91 @@ class CogVideoX_I2VController: generator = torch.Generator(device="cuda").manual_seed(int(seed_textbox)) try: - if self.transformer.config.in_channels != self.vae.config.latent_channels: - if generation_method == "Long Video Generation": - if validation_video is not None: - raise gr.Error(f"Video to Video is not Support Long Video Generation now.") - init_frames = 0 - last_frames = init_frames + partial_video_length - while init_frames < length_slider: - if last_frames >= length_slider: - _partial_video_length = length_slider - init_frames - _partial_video_length = int((_partial_video_length - 1) // self.vae.config.temporal_compression_ratio * self.vae.config.temporal_compression_ratio) + 1 + if self.model_type == "Inpaint": + if self.transformer.config.in_channels != self.vae.config.latent_channels: + if generation_method == "Long Video Generation": + if validation_video is not None: + raise gr.Error(f"Video to Video is not Support Long Video Generation now.") + init_frames = 0 + last_frames = init_frames + partial_video_length + while init_frames < length_slider: + if last_frames >= length_slider: + _partial_video_length = length_slider - init_frames + _partial_video_length = int((_partial_video_length - 1) // self.vae.config.temporal_compression_ratio * self.vae.config.temporal_compression_ratio) + 1 + + if _partial_video_length <= 0: + break + else: + _partial_video_length = partial_video_length + + if last_frames >= length_slider: + input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=_partial_video_length, sample_size=(height_slider, width_slider)) + else: + input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, None, video_length=_partial_video_length, sample_size=(height_slider, width_slider)) + + with torch.no_grad(): + sample = self.pipeline( + prompt_textbox, + negative_prompt = negative_prompt_textbox, + num_inference_steps = sample_step_slider, + guidance_scale = cfg_scale_slider, + width = width_slider, + height = height_slider, + num_frames = _partial_video_length, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + strength = 1, + ).videos - if _partial_video_length <= 0: + if init_frames != 0: + mix_ratio = torch.from_numpy( + np.array([float(_index) / float(overlap_video_length) for _index in range(overlap_video_length)], np.float32) + ).unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1) + + new_sample[:, :, -overlap_video_length:] = new_sample[:, :, -overlap_video_length:] * (1 - mix_ratio) + \ + sample[:, :, :overlap_video_length] * mix_ratio + new_sample = torch.cat([new_sample, sample[:, :, overlap_video_length:]], dim = 2) + + sample = new_sample + else: + new_sample = sample + + if last_frames >= length_slider: break - else: - _partial_video_length = partial_video_length - if last_frames >= length_slider: - input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=_partial_video_length, sample_size=(height_slider, width_slider)) - else: - input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, None, video_length=_partial_video_length, sample_size=(height_slider, width_slider)) + start_image = [ + Image.fromarray( + (sample[0, :, _index].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8) + ) for _index in range(-overlap_video_length, 0) + ] - with torch.no_grad(): - sample = self.pipeline( - prompt_textbox, - negative_prompt = negative_prompt_textbox, - num_inference_steps = sample_step_slider, - guidance_scale = cfg_scale_slider, - width = width_slider, - height = height_slider, - num_frames = _partial_video_length, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - strength = 1, - ).videos - - if init_frames != 0: - mix_ratio = torch.from_numpy( - np.array([float(_index) / float(overlap_video_length) for _index in range(overlap_video_length)], np.float32) - ).unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1) - - new_sample[:, :, -overlap_video_length:] = new_sample[:, :, -overlap_video_length:] * (1 - mix_ratio) + \ - sample[:, :, :overlap_video_length] * mix_ratio - new_sample = torch.cat([new_sample, sample[:, :, overlap_video_length:]], dim = 2) - - sample = new_sample - else: - new_sample = sample - - if last_frames >= length_slider: - break - - start_image = [ - Image.fromarray( - (sample[0, :, _index].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8) - ) for _index in range(-overlap_video_length, 0) - ] - - init_frames = init_frames + _partial_video_length - overlap_video_length - last_frames = init_frames + _partial_video_length - else: - if validation_video is not None: - input_video, input_video_mask, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider)) - strength = denoise_strength + init_frames = init_frames + _partial_video_length - overlap_video_length + last_frames = init_frames + _partial_video_length else: - input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider)) - strength = 1 + if validation_video is not None: + input_video, input_video_mask, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=8) + strength = denoise_strength + else: + input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider)) + strength = 1 + sample = self.pipeline( + prompt_textbox, + negative_prompt = negative_prompt_textbox, + num_inference_steps = sample_step_slider, + guidance_scale = cfg_scale_slider, + width = width_slider, + height = height_slider, + num_frames = length_slider if not is_image else 1, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + strength = strength, + ).videos + else: sample = self.pipeline( prompt_textbox, negative_prompt = negative_prompt_textbox, @@ -332,13 +379,11 @@ class CogVideoX_I2VController: width = width_slider, height = height_slider, num_frames = length_slider if not is_image else 1, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - strength = strength, + generator = generator ).videos else: + input_video, input_video_mask, clip_image = get_video_to_video_latent(control_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), fps=8) + sample = self.pipeline( prompt_textbox, negative_prompt = negative_prompt_textbox, @@ -347,7 +392,9 @@ class CogVideoX_I2VController: width = width_slider, height = height_slider, num_frames = length_slider if not is_image else 1, - generator = generator + generator = generator, + + control_video = input_video, ).videos except Exception as e: gc.collect() @@ -422,7 +469,7 @@ class CogVideoX_I2VController: def ui(low_gpu_memory_mode, weight_dtype): - controller = CogVideoX_I2VController(low_gpu_memory_mode, weight_dtype) + controller = CogVideoX_Fun_Controller(low_gpu_memory_mode, weight_dtype) with gr.Blocks(css=css) as demo: gr.Markdown( @@ -437,7 +484,20 @@ def ui(low_gpu_memory_mode, weight_dtype): with gr.Column(variant="panel"): gr.Markdown( """ - ### 1. Model checkpoints (模型路径). + ### 1. CogVideoX-Fun Model Type (CogVideoX-Fun模型的种类,正常模型还是控制模型). + """ + ) + with gr.Row(): + model_type = gr.Dropdown( + label="The model type of CogVideoX-Fun (CogVideoX-Fun模型的种类,正常模型还是控制模型)", + choices=["Inpaint", "Control"], + value="Inpaint", + interactive=True, + ) + + gr.Markdown( + """ + ### 2. Model checkpoints (模型路径). """ ) with gr.Row(): @@ -488,12 +548,12 @@ def ui(low_gpu_memory_mode, weight_dtype): with gr.Column(variant="panel"): gr.Markdown( """ - ### 2. Configs for Generation (生成参数配置). + ### 3. Configs for Generation (生成参数配置). """ ) prompt_textbox = gr.Textbox(label="Prompt (正向提示词)", lines=2, value="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.") - negative_prompt_textbox = gr.Textbox(label="Negative prompt (负向提示词)", lines=2, value="The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " ) + negative_prompt_textbox = gr.Textbox(label="Negative prompt (负向提示词)", lines=2, value="The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. " ) with gr.Row(): with gr.Column(): @@ -522,7 +582,7 @@ def ui(low_gpu_memory_mode, weight_dtype): partial_video_length = gr.Slider(label="Partial video generation length (每个部分的视频生成帧数)", value=25, minimum=5, maximum=49, step=4, visible=False) source_method = gr.Radio( - ["Text to Video (文本到视频)", "Image to Video (图片到视频)", "Video to Video (视频到视频)"], + ["Text to Video (文本到视频)", "Image to Video (图片到视频)", "Video to Video (视频到视频)", "Video Control (视频控制)"], value="Text to Video (文本到视频)", show_label=False, ) @@ -557,13 +617,36 @@ def ui(low_gpu_memory_mode, weight_dtype): end_image = gr.Image(label="The image at the ending of the video (图片到视频的结束图片[非必需, Optional])", show_label=False, elem_id="i2v_end", sources="upload", type="filepath") with gr.Column(visible = False) as video_to_video_col: - validation_video = gr.Video( - label="The video to convert (视频转视频的参考视频)", show_label=True, - elem_id="v2v", sources="upload", - ) - denoise_strength = gr.Slider(label="Denoise strength (重绘系数)", value=0.70, minimum=0.10, maximum=0.95, step=0.01) + with gr.Row(): + validation_video = gr.Video( + label="The video to convert (视频转视频的参考视频)", show_label=True, + elem_id="v2v", sources="upload", + ) + with gr.Accordion("The mask of the video to inpaint (视频重新绘制的mask[非必需, Optional])", open=False): + gr.Markdown( + """ + - Please set a larger denoise_strength when using validation_video_mask, such as 1.00 instead of 0.70 + - (请设置更大的denoise_strength,当使用validation_video_mask的时候,比如1而不是0.70) + """ + ) + validation_video_mask = gr.Image( + label="The mask of the video to inpaint (视频重新绘制的mask[非必需, Optional])", + show_label=False, elem_id="v2v_mask", sources="upload", type="filepath" + ) + denoise_strength = gr.Slider(label="Denoise strength (重绘系数)", value=0.70, minimum=0.10, maximum=1.00, step=0.01) - cfg_scale_slider = gr.Slider(label="CFG Scale (引导系数)", value=7.0, minimum=0, maximum=20) + with gr.Column(visible = False) as control_video_col: + gr.Markdown( + """ + Demo pose control video can be downloaded here [URL](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/pose.mp4). + """ + ) + control_video = gr.Video( + label="The control video (用于提供控制信号的video)", show_label=True, + elem_id="v2v_control", sources="upload", + ) + + cfg_scale_slider = gr.Slider(label="CFG Scale (引导系数)", value=6.0, minimum=0, maximum=20) with gr.Row(): seed_textbox = gr.Textbox(label="Seed (随机种子)", value=43) @@ -585,6 +668,12 @@ def ui(low_gpu_memory_mode, weight_dtype): interactive=False ) + model_type.change( + fn=controller.update_model_type, + inputs=[model_type], + outputs=[] + ) + def upload_generation_method(generation_method): if generation_method == "Video Generation": return [gr.update(visible=True, maximum=49, value=49), gr.update(visible=False), gr.update(visible=False)] @@ -598,13 +687,18 @@ def ui(low_gpu_memory_mode, weight_dtype): def upload_source_method(source_method): if source_method == "Text to Video (文本到视频)": - return [gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)] elif source_method == "Image to Video (图片到视频)": - return [gr.update(visible=True), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None)] + return [gr.update(visible=True), gr.update(visible=False), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + elif source_method == "Video to Video (视频到视频)": + return [gr.update(visible=False), gr.update(visible=True), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(), gr.update(), gr.update(value=None)] else: - return [gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update()] + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update()] source_method.change( - upload_source_method, source_method, [image_to_video_col, video_to_video_col, start_image, end_image, validation_video] + upload_source_method, source_method, [ + image_to_video_col, video_to_video_col, control_video_col, start_image, end_image, + validation_video, validation_video_mask, control_video + ] ) def upload_resize_method(resize_method): @@ -639,6 +733,8 @@ def ui(low_gpu_memory_mode, weight_dtype): start_image, end_image, validation_video, + validation_video_mask, + control_video, denoise_strength, seed_textbox, ], @@ -647,8 +743,8 @@ def ui(low_gpu_memory_mode, weight_dtype): return demo, controller -class CogVideoX_I2VController_Modelscope: - def __init__(self, model_name, savedir_sample, low_gpu_memory_mode, weight_dtype): +class CogVideoX_Fun_Controller_Modelscope: + def __init__(self, model_name, model_type, savedir_sample, low_gpu_memory_mode, weight_dtype): # Basic dir self.basedir = os.getcwd() self.personalized_model_dir = os.path.join(self.basedir, "models", "Personalized_Model") @@ -658,6 +754,7 @@ class CogVideoX_I2VController_Modelscope: os.makedirs(self.savedir_sample, exist_ok=True) # model path + self.model_type = model_type self.weight_dtype = weight_dtype self.vae = AutoencoderKLCogVideoX.from_pretrained( @@ -672,16 +769,25 @@ class CogVideoX_I2VController_Modelscope: ).to(self.weight_dtype) # Get pipeline - if self.transformer.config.in_channels != self.vae.config.latent_channels: - self.pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( - model_name, - vae=self.vae, - transformer=self.transformer, - scheduler=scheduler_dict["Euler"].from_pretrained(model_name, subfolder="scheduler"), - torch_dtype=self.weight_dtype - ) + if model_type == "Inpaint": + if self.transformer.config.in_channels != self.vae.config.latent_channels: + self.pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( + model_name, + vae=self.vae, + transformer=self.transformer, + scheduler=scheduler_dict["Euler"].from_pretrained(model_name, subfolder="scheduler"), + torch_dtype=self.weight_dtype + ) + else: + self.pipeline = CogVideoX_Fun_Pipeline.from_pretrained( + model_name, + vae=self.vae, + transformer=self.transformer, + scheduler=scheduler_dict["Euler"].from_pretrained(model_name, subfolder="scheduler"), + torch_dtype=self.weight_dtype + ) else: - self.pipeline = CogVideoX_Fun_Pipeline.from_pretrained( + self.pipeline = CogVideoX_Fun_Pipeline_Control.from_pretrained( model_name, vae=self.vae, transformer=self.transformer, @@ -733,6 +839,8 @@ class CogVideoX_I2VController_Modelscope: start_image, end_image, validation_video, + validation_video_mask, + control_video, denoise_strength, seed_textbox, is_api = False, @@ -747,25 +855,48 @@ class CogVideoX_I2VController_Modelscope: if self.lora_model_path != lora_model_dropdown: print("Update lora model") self.update_lora_model(lora_model_dropdown) + + if control_video is not None and self.model_type == "Inpaint": + if is_api: + return "", f"If specifying the control video, please set the model_type == \"Control\". " + else: + raise gr.Error(f"If specifying the control video, please set the model_type == \"Control\". ") + + if control_video is None and self.model_type == "Control": + if is_api: + return "", f"If set the model_type == \"Control\", please specifying the control video. " + else: + raise gr.Error(f"If set the model_type == \"Control\", please specifying the control video. ") if resize_method == "Resize according to Reference": - if start_image is None and validation_video is None: - raise gr.Error(f"Please upload an image when using \"Resize according to Reference\".") + if start_image is None and validation_video is None and control_video is None: + if is_api: + return "", f"Please upload an image when using \"Resize according to Reference\"." + else: + raise gr.Error(f"Please upload an image when using \"Resize according to Reference\".") - aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} - - if validation_video is not None: - original_width, original_height = Image.fromarray(cv2.VideoCapture(validation_video).read()[1]).size + aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + if self.model_type == "Inpaint": + if validation_video is not None: + original_width, original_height = Image.fromarray(cv2.VideoCapture(validation_video).read()[1]).size + else: + original_width, original_height = start_image[0].size if type(start_image) is list else Image.open(start_image).size else: - original_width, original_height = start_image[0].size if type(start_image) is list else Image.open(start_image).size + original_width, original_height = Image.fromarray(cv2.VideoCapture(control_video).read()[1]).size closest_size, closest_ratio = get_closest_ratio(original_height, original_width, ratios=aspect_ratio_sample_size) height_slider, width_slider = [int(x / 16) * 16 for x in closest_size] if self.transformer.config.in_channels == self.vae.config.latent_channels and start_image is not None: - raise gr.Error(f"Please select an image to video pretrained model while using image to video.") - + if is_api: + return "", f"Please select an image to video pretrained model while using image to video." + else: + raise gr.Error(f"Please select an image to video pretrained model while using image to video.") + if start_image is None and end_image is not None: - raise gr.Error(f"If specifying the ending image of the video, please specify a starting image of the video.") + if is_api: + return "", f"If specifying the ending image of the video, please specify a starting image of the video." + else: + raise gr.Error(f"If specifying the ending image of the video, please specify a starting image of the video.") is_image = True if generation_method == "Image Generation" else False @@ -779,13 +910,42 @@ class CogVideoX_I2VController_Modelscope: generator = torch.Generator(device="cuda").manual_seed(int(seed_textbox)) try: - if self.transformer.config.in_channels != self.vae.config.latent_channels: - if validation_video is not None: - input_video, input_video_mask, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider)) - strength = denoise_strength + if self.model_type == "Inpaint": + if self.transformer.config.in_channels != self.vae.config.latent_channels: + if validation_video is not None: + input_video, input_video_mask, clip_image = get_video_to_video_latent(validation_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), validation_video_mask=validation_video_mask, fps=8) + strength = denoise_strength + else: + input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider)) + strength = 1 + + sample = self.pipeline( + prompt_textbox, + negative_prompt = negative_prompt_textbox, + num_inference_steps = sample_step_slider, + guidance_scale = cfg_scale_slider, + width = width_slider, + height = height_slider, + num_frames = length_slider if not is_image else 1, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + strength = strength, + ).videos else: - input_video, input_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, length_slider if not is_image else 1, sample_size=(height_slider, width_slider)) - strength = 1 + sample = self.pipeline( + prompt_textbox, + negative_prompt = negative_prompt_textbox, + num_inference_steps = sample_step_slider, + guidance_scale = cfg_scale_slider, + width = width_slider, + height = height_slider, + num_frames = length_slider if not is_image else 1, + generator = generator + ).videos + else: + input_video, input_video_mask, clip_image = get_video_to_video_latent(control_video, length_slider if not is_image else 1, sample_size=(height_slider, width_slider), fps=8) sample = self.pipeline( prompt_textbox, @@ -797,20 +957,7 @@ class CogVideoX_I2VController_Modelscope: num_frames = length_slider if not is_image else 1, generator = generator, - video = input_video, - mask_video = input_video_mask, - strength = strength, - ).videos - else: - sample = self.pipeline( - prompt_textbox, - negative_prompt = negative_prompt_textbox, - num_inference_steps = sample_step_slider, - guidance_scale = cfg_scale_slider, - width = width_slider, - height = height_slider, - num_frames = length_slider if not is_image else 1, - generator = generator + control_video = input_video, ).videos except Exception as e: gc.collect() @@ -866,8 +1013,8 @@ class CogVideoX_I2VController_Modelscope: return gr.Image.update(visible=False, value=None), gr.Video.update(value=save_sample_path, visible=True), "Success" -def ui_modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype): - controller = CogVideoX_I2VController_Modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype) +def ui_modelscope(model_name, model_type, savedir_sample, low_gpu_memory_mode, weight_dtype): + controller = CogVideoX_Fun_Controller_Modelscope(model_name, model_type, savedir_sample, low_gpu_memory_mode, weight_dtype) with gr.Blocks(css=css) as demo: gr.Markdown( @@ -882,7 +1029,20 @@ def ui_modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype) with gr.Column(variant="panel"): gr.Markdown( """ - ### 1. Model checkpoints (模型路径). + ### 1. CogVideoX-Fun Model Type (CogVideoX-Fun模型的种类,正常模型还是控制模型). + """ + ) + with gr.Row(): + model_type = gr.Dropdown( + label="The model type of CogVideoX-Fun (CogVideoX-Fun模型的种类,正常模型还是控制模型)", + choices=[model_type], + value=model_type, + interactive=False, + ) + + gr.Markdown( + """ + ### 2. Model checkpoints (模型路径). """ ) with gr.Row(): @@ -919,12 +1079,12 @@ def ui_modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype) with gr.Column(variant="panel"): gr.Markdown( """ - ### 2. Configs for Generation (生成参数配置). + ### 3. Configs for Generation (生成参数配置). """ ) prompt_textbox = gr.Textbox(label="Prompt (正向提示词)", lines=2, value="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.") - negative_prompt_textbox = gr.Textbox(label="Negative prompt (负向提示词)", lines=2, value="The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " ) + negative_prompt_textbox = gr.Textbox(label="Negative prompt (负向提示词)", lines=2, value="The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. " ) with gr.Row(): with gr.Column(): @@ -953,7 +1113,7 @@ def ui_modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype) partial_video_length = gr.Slider(label="Partial video generation length (每个部分的视频生成帧数)", value=25, minimum=5, maximum=49, step=4, visible=False) source_method = gr.Radio( - ["Text to Video (文本到视频)", "Image to Video (图片到视频)", "Video to Video (视频到视频)"], + ["Text to Video (文本到视频)", "Image to Video (图片到视频)", "Video to Video (视频到视频)", "Video Control (视频控制)"], value="Text to Video (文本到视频)", show_label=False, ) @@ -986,13 +1146,36 @@ def ui_modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype) end_image = gr.Image(label="The image at the ending of the video (图片到视频的结束图片[非必需, Optional])", show_label=False, elem_id="i2v_end", sources="upload", type="filepath") with gr.Column(visible = False) as video_to_video_col: - validation_video = gr.Video( - label="The video to convert (视频转视频的参考视频)", show_label=True, - elem_id="v2v", sources="upload", - ) - denoise_strength = gr.Slider(label="Denoise strength (重绘系数)", value=0.70, minimum=0.10, maximum=0.95, step=0.01) + with gr.Row(): + validation_video = gr.Video( + label="The video to convert (视频转视频的参考视频)", show_label=True, + elem_id="v2v", sources="upload", + ) + with gr.Accordion("The mask of the video to inpaint (视频重新绘制的mask[非必需, Optional])", open=False): + gr.Markdown( + """ + - Please set a larger denoise_strength when using validation_video_mask, such as 1.00 instead of 0.70 + - (请设置更大的denoise_strength,当使用validation_video_mask的时候,比如1而不是0.70) + """ + ) + validation_video_mask = gr.Image( + label="The mask of the video to inpaint (视频重新绘制的mask[非必需, Optional])", + show_label=False, elem_id="v2v_mask", sources="upload", type="filepath" + ) + denoise_strength = gr.Slider(label="Denoise strength (重绘系数)", value=0.70, minimum=0.10, maximum=1.00, step=0.01) - cfg_scale_slider = gr.Slider(label="CFG Scale (引导系数)", value=7.0, minimum=0, maximum=20) + with gr.Column(visible = False) as control_video_col: + gr.Markdown( + """ + Demo pose control video can be downloaded here [URL](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/pose.mp4). + """ + ) + control_video = gr.Video( + label="The control video (用于提供控制信号的video)", show_label=True, + elem_id="v2v_control", sources="upload", + ) + + cfg_scale_slider = gr.Slider(label="CFG Scale (引导系数)", value=6.0, minimum=0, maximum=20) with gr.Row(): seed_textbox = gr.Textbox(label="Seed (随机种子)", value=43) @@ -1025,13 +1208,18 @@ def ui_modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype) def upload_source_method(source_method): if source_method == "Text to Video (文本到视频)": - return [gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)] elif source_method == "Image to Video (图片到视频)": - return [gr.update(visible=True), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None)] + return [gr.update(visible=True), gr.update(visible=False), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + elif source_method == "Video to Video (视频到视频)": + return [gr.update(visible=False), gr.update(visible=True), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(), gr.update(), gr.update(value=None)] else: - return [gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update()] + return [gr.update(visible=False), gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update()] source_method.change( - upload_source_method, source_method, [image_to_video_col, video_to_video_col, start_image, end_image, validation_video] + upload_source_method, source_method, [ + image_to_video_col, video_to_video_col, control_video_col, start_image, end_image, + validation_video, validation_video_mask, control_video + ] ) def upload_resize_method(resize_method): @@ -1066,6 +1254,8 @@ def ui_modelscope(model_name, savedir_sample, low_gpu_memory_mode, weight_dtype) start_image, end_image, validation_video, + validation_video_mask, + control_video, denoise_strength, seed_textbox, ], @@ -1080,7 +1270,7 @@ def post_eas( prompt_textbox, negative_prompt_textbox, sampler_dropdown, sample_step_slider, resize_method, width_slider, height_slider, base_resolution, generation_method, length_slider, cfg_scale_slider, - start_image, end_image, validation_video, denoise_strength, seed_textbox, + start_image, end_image, validation_video, validation_video_mask, denoise_strength, seed_textbox, ): if start_image is not None: with open(start_image, 'rb') as file: @@ -1100,6 +1290,12 @@ def post_eas( validation_video_encoded_content = base64.b64encode(file_content) validation_video = validation_video_encoded_content.decode('utf-8') + if validation_video_mask is not None: + with open(validation_video_mask, 'rb') as file: + file_content = file.read() + validation_video_mask_encoded_content = base64.b64encode(file_content) + validation_video_mask = validation_video_mask_encoded_content.decode('utf-8') + datas = { "base_model_path": base_model_dropdown, "lora_model_path": lora_model_dropdown, @@ -1118,6 +1314,7 @@ def post_eas( "start_image": start_image, "end_image": end_image, "validation_video": validation_video, + "validation_video_mask": validation_video_mask, "denoise_strength": denoise_strength, "seed_textbox": seed_textbox, } @@ -1131,7 +1328,7 @@ def post_eas( return outputs -class CogVideoX_I2VController_EAS: +class CogVideoX_Fun_Controller_EAS: def __init__(self, model_name, savedir_sample): self.savedir_sample = savedir_sample os.makedirs(self.savedir_sample, exist_ok=True) @@ -1156,6 +1353,7 @@ class CogVideoX_I2VController_EAS: start_image, end_image, validation_video, + validation_video_mask, denoise_strength, seed_textbox ): @@ -1167,7 +1365,7 @@ class CogVideoX_I2VController_EAS: prompt_textbox, negative_prompt_textbox, sampler_dropdown, sample_step_slider, resize_method, width_slider, height_slider, base_resolution, generation_method, length_slider, cfg_scale_slider, - start_image, end_image, validation_video, denoise_strength, + start_image, end_image, validation_video, validation_video_mask, denoise_strength, seed_textbox ) try: @@ -1201,7 +1399,7 @@ class CogVideoX_I2VController_EAS: def ui_eas(model_name, savedir_sample): - controller = CogVideoX_I2VController_EAS(model_name, savedir_sample) + controller = CogVideoX_Fun_Controller_EAS(model_name, savedir_sample) with gr.Blocks(css=css) as demo: gr.Markdown( @@ -1216,7 +1414,7 @@ def ui_eas(model_name, savedir_sample): with gr.Column(variant="panel"): gr.Markdown( """ - ### 1. Model checkpoints. + ### 1. Model checkpoints (模型路径). """ ) with gr.Row(): @@ -1258,7 +1456,7 @@ def ui_eas(model_name, savedir_sample): ) prompt_textbox = gr.Textbox(label="Prompt", lines=2, value="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.") - negative_prompt_textbox = gr.Textbox(label="Negative prompt", lines=2, value="The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " ) + negative_prompt_textbox = gr.Textbox(label="Negative prompt", lines=2, value="The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. " ) with gr.Row(): with gr.Column(): @@ -1317,13 +1515,25 @@ def ui_eas(model_name, savedir_sample): end_image = gr.Image(label="The image at the ending of the video (Optional)", show_label=True, elem_id="i2v_end", sources="upload", type="filepath") with gr.Column(visible = False) as video_to_video_col: - validation_video = gr.Video( - label="The video to convert (视频转视频的参考视频)", show_label=True, - elem_id="v2v", sources="upload", - ) - denoise_strength = gr.Slider(label="Denoise strength (重绘系数)", value=0.70, minimum=0.10, maximum=0.95, step=0.01) + with gr.Row(): + validation_video = gr.Video( + label="The video to convert (视频转视频的参考视频)", show_label=True, + elem_id="v2v", sources="upload", + ) + with gr.Accordion("The mask of the video to inpaint (视频重新绘制的mask[非必需, Optional])", open=False): + gr.Markdown( + """ + - Please set a larger denoise_strength when using validation_video_mask, such as 1.00 instead of 0.70 + - (请设置更大的denoise_strength,当使用validation_video_mask的时候,比如1而不是0.70) + """ + ) + validation_video_mask = gr.Image( + label="The mask of the video to inpaint (视频重新绘制的mask[非必需, Optional])", + show_label=False, elem_id="v2v_mask", sources="upload", type="filepath" + ) + denoise_strength = gr.Slider(label="Denoise strength (重绘系数)", value=0.70, minimum=0.10, maximum=1.00, step=0.01) - cfg_scale_slider = gr.Slider(label="CFG Scale (引导系数)", value=7.0, minimum=0, maximum=20) + cfg_scale_slider = gr.Slider(label="CFG Scale (引导系数)", value=6.0, minimum=0, maximum=20) with gr.Row(): seed_textbox = gr.Textbox(label="Seed", value=43) @@ -1356,13 +1566,13 @@ def ui_eas(model_name, savedir_sample): def upload_source_method(source_method): if source_method == "Text to Video (文本到视频)": - return [gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None)] + return [gr.update(visible=False), gr.update(visible=False), gr.update(value=None), gr.update(value=None), gr.update(value=None), gr.update(value=None)] elif source_method == "Image to Video (图片到视频)": - return [gr.update(visible=True), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None)] + return [gr.update(visible=True), gr.update(visible=False), gr.update(), gr.update(), gr.update(value=None), gr.update(value=None)] else: - return [gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update()] + return [gr.update(visible=False), gr.update(visible=True), gr.update(value=None), gr.update(value=None), gr.update(), gr.update()] source_method.change( - upload_source_method, source_method, [image_to_video_col, video_to_video_col, start_image, end_image, validation_video] + upload_source_method, source_method, [image_to_video_col, video_to_video_col, start_image, end_image, validation_video, validation_video_mask] ) def upload_resize_method(resize_method): @@ -1395,6 +1605,7 @@ def ui_eas(model_name, savedir_sample): start_image, end_image, validation_video, + validation_video_mask, denoise_strength, seed_textbox, ], diff --git a/cogvideox/utils/utils.py b/cogvideox/utils/utils.py index 7fc1e83..a9298a3 100644 --- a/cogvideox/utils/utils.py +++ b/cogvideox/utils/utils.py @@ -166,16 +166,27 @@ def get_image_to_video_latent(validation_image_start, validation_image_end, vide return input_video, input_video_mask, clip_image -def get_video_to_video_latent(input_video_path, video_length, sample_size): - if type(input_video_path) is str: +def get_video_to_video_latent(input_video_path, video_length, sample_size, fps=None, validation_video_mask=None): + if isinstance(input_video_path, str): cap = cv2.VideoCapture(input_video_path) input_video = [] + + original_fps = cap.get(cv2.CAP_PROP_FPS) + frame_skip = 1 if fps is None else int(original_fps // fps) + + frame_count = 0 + while True: ret, frame = cap.read() if not ret: break - frame = cv2.resize(frame, (sample_size[1], sample_size[0])) - input_video.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)) + + if frame_count % frame_skip == 0: + frame = cv2.resize(frame, (sample_size[1], sample_size[0])) + input_video.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)) + + frame_count += 1 + cap.release() else: input_video = input_video_path @@ -183,7 +194,15 @@ def get_video_to_video_latent(input_video_path, video_length, sample_size): input_video = torch.from_numpy(np.array(input_video))[:video_length] input_video = input_video.permute([3, 0, 1, 2]).unsqueeze(0) / 255 - input_video_mask = torch.zeros_like(input_video[:, :1]) - input_video_mask[:, :, :] = 255 + if validation_video_mask is not None: + validation_video_mask = Image.open(validation_video_mask).convert('L').resize((sample_size[1], sample_size[0])) + input_video_mask = np.where(np.array(validation_video_mask) < 240, 0, 255) + + input_video_mask = torch.from_numpy(np.array(input_video_mask)).unsqueeze(0).unsqueeze(-1).permute([3, 0, 1, 2]).unsqueeze(0) + input_video_mask = torch.tile(input_video_mask, [1, 1, input_video.size()[2], 1, 1]) + input_video_mask = input_video_mask.to(input_video.device, input_video.dtype) + else: + input_video_mask = torch.zeros_like(input_video[:, :1]) + input_video_mask[:, :, :] = 255 return input_video, input_video_mask, None \ No newline at end of file diff --git a/comfyui/README.md b/comfyui/README.md index e029304..305664c 100644 --- a/comfyui/README.md +++ b/comfyui/README.md @@ -28,6 +28,17 @@ python install.py ### 2. Download models into `ComfyUI/models/CogVideoX_Fun/` +V1.1: + +| 名称 | 存储空间 | Hugging Face | Model Scope | 描述 | +|--|--|--|--|--| +| CogVideoX-Fun-V1.1-2b-InP.tar.gz | Before extraction:9.7 GB \/ After extraction: 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. Noise has been added to the reference image, and the amplitude of motion is greater compared to V1.0. | +| CogVideoX-Fun-V1.1-5b-InP.tar.gz | Before extraction:16.0 GB \/ After extraction: 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. Noise has been added to the reference image, and the amplitude of motion is greater compared to V1.0. | +| CogVideoX-Fun-V1.1-2b-Pose.tar.gz | Before extraction:9.7 GB \/ After extraction: 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-Pose) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-Pose) | Our official pose-control video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second.| +| CogVideoX-Fun-V1.1-5b-Pose.tar.gz | Before extraction:16.0 GB \/ After extraction: 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-Pose) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-Pose) | Our official pose-control video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second.| + +V1.0: + | Name | Storage Space | Hugging Face | Model Scope | Description | |--|--|--|--|--| | CogVideoX-Fun-2b-InP.tar.gz | Before extraction:9.7 GB \/ After extraction: 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-2b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-2b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. | @@ -48,19 +59,26 @@ python install.py ## Example workflows ### Video to video generation -Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/cogvideoxfunv1_workflow_v2v.json) of the json: -![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/cogvideoxfunv1_workflow_v2v.jpg) +Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/cogvideoxfunv1.1_workflow_v2v.json) of the json: +![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/cogvideoxfunv1.1_workflow_v2v.jpg) You can run the demo using following video: [demo video](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/play_guitar.mp4) +### Control video generation +Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/cogvideoxfunv1.1_workflow_v2v_control.json) of the json: +![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/cogvideoxfunv1.1_workflow_v2v_control.jpg) + +You can run the demo using following video: +[demo video](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/pose.mp4) + ### Image to video generation -Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/cogvideoxfunv1_workflow_i2v.json) of the json: -![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/cogvideoxfunv1_workflow_i2v.jpg) +Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/cogvideoxfunv1.1_workflow_i2v.json) of the json: +![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/cogvideoxfunv1.1_workflow_i2v.jpg) You can run the demo using following photo: ![demo image](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/firework.png) ### Text to video generation -Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/cogvideoxfunv1_workflow_t2v.json) of the json: -![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/cogvideoxfunv1_workflow_t2v.jpg) \ No newline at end of file +Our ui is shown as follow, this is the [download link](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/cogvideoxfunv1.1_workflow_t2v.json) of the json: +![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/cogvideoxfunv1.1_workflow_t2v.jpg) \ No newline at end of file diff --git a/comfyui/comfyui_nodes.py b/comfyui/comfyui_nodes.py index cc21102..9e7eeb1 100644 --- a/comfyui/comfyui_nodes.py +++ b/comfyui/comfyui_nodes.py @@ -23,6 +23,8 @@ from ..cogvideox.data.bucket_sampler import ASPECT_RATIO_512, get_closest_ratio from ..cogvideox.models.autoencoder_magvit import AutoencoderKLCogVideoX from ..cogvideox.models.transformer3d import CogVideoXTransformer3DModel from ..cogvideox.pipeline.pipeline_cogvideox import CogVideoX_Fun_Pipeline +from ..cogvideox.pipeline.pipeline_cogvideox_control import \ + CogVideoX_Fun_Pipeline_Control from ..cogvideox.pipeline.pipeline_cogvideox_inpaint import ( CogVideoX_Fun_Pipeline_Inpaint) from ..cogvideox.utils.lora_utils import merge_lora, unmerge_lora @@ -59,9 +61,19 @@ class LoadCogVideoX_Fun_Model: [ 'CogVideoX-Fun-2b-InP', 'CogVideoX-Fun-5b-InP', + 'CogVideoX-Fun-V1.1-2b-InP', + 'CogVideoX-Fun-V1.1-5b-InP', + 'CogVideoX-Fun-V1.1-2b-Pose', + 'CogVideoX-Fun-V1.1-5b-Pose', ], { - "default": 'CogVideoX-Fun-2b-InP', + "default": 'CogVideoX-Fun-V1.1-2b-InP', + } + ), + "model_type": ( + ["Inpaint", "Control"], + { + "default": "Inpaint", } ), "low_gpu_memory_mode":( @@ -85,7 +97,7 @@ class LoadCogVideoX_Fun_Model: FUNCTION = "loadmodel" CATEGORY = "CogVideoXFUNWrapper" - def loadmodel(self, low_gpu_memory_mode, model, precision): + def loadmodel(self, low_gpu_memory_mode, model, model_type, precision): # Init weight_dtype and device device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -131,22 +143,31 @@ class LoadCogVideoX_Fun_Model: pbar.update(1) # Get pipeline - if transformer.config.in_channels != vae.config.latent_channels: - pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( - model_path, - vae=vae, - transformer=transformer, - scheduler=scheduler, - torch_dtype=weight_dtype - ) + if model_type == "Inpaint": + if transformer.config.in_channels != vae.config.latent_channels: + pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( + model_path, + vae=vae, + transformer=transformer, + scheduler=scheduler, + torch_dtype=weight_dtype + ) + else: + pipeline = CogVideoX_Fun_Pipeline.from_pretrained( + model_path, + vae=vae, + transformer=transformer, + scheduler=scheduler, + torch_dtype=weight_dtype + ) else: - pipeline = CogVideoX_Fun_Pipeline.from_pretrained( - model_path, - vae=vae, - transformer=transformer, - scheduler=scheduler, - torch_dtype=weight_dtype - ) + pipeline = CogVideoX_Fun_Pipeline_Control.from_pretrained( + model_path, + vae=vae, + transformer=transformer, + scheduler=scheduler, + torch_dtype=weight_dtype + ) if low_gpu_memory_mode: pipeline.enable_sequential_cpu_offload() else: @@ -156,6 +177,7 @@ class LoadCogVideoX_Fun_Model: 'pipeline': pipeline, 'dtype': weight_dtype, 'model_path': model_path, + 'model_type': model_type, 'loras': [], 'strength_model': [], } @@ -491,8 +513,11 @@ class CogVideoX_Fun_V2VSampler: "default": 'DDIM' } ), + }, + "optional":{ "validation_video": ("IMAGE",), - } + "control_video": ("IMAGE",), + }, } RETURN_TYPES = ("IMAGE",) @@ -500,26 +525,34 @@ class CogVideoX_Fun_V2VSampler: FUNCTION = "process" CATEGORY = "CogVideoXFUNWrapper" - def process(self, cogvideoxfun_model, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, denoise_strength, scheduler, validation_video): + def process(self, cogvideoxfun_model, prompt, negative_prompt, video_length, base_resolution, seed, steps, cfg, denoise_strength, scheduler, validation_video=None, control_video=None): device = mm.get_torch_device() offload_device = mm.unet_offload_device() mm.soft_empty_cache() gc.collect() - - # Count most suitable height and width - aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} - if type(validation_video) is str: - original_width, original_height = Image.fromarray(cv2.VideoCapture(validation_video).read()[1]).size - else: - validation_video = np.array(validation_video.cpu().numpy() * 255, np.uint8) - original_width, original_height = Image.fromarray(validation_video[0]).size - closest_size, closest_ratio = get_closest_ratio(original_height, original_width, ratios=aspect_ratio_sample_size) - height, width = [int(x / 16) * 16 for x in closest_size] # Get Pipeline pipeline = cogvideoxfun_model['pipeline'] model_path = cogvideoxfun_model['model_path'] + model_type = cogvideoxfun_model['model_type'] + + # Count most suitable height and width + aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + if model_type == "Inpaint": + if type(validation_video) is str: + original_width, original_height = Image.fromarray(cv2.VideoCapture(validation_video).read()[1]).size + else: + validation_video = np.array(validation_video.cpu().numpy() * 255, np.uint8) + original_width, original_height = Image.fromarray(validation_video[0]).size + else: + if type(control_video) is str: + original_width, original_height = Image.fromarray(cv2.VideoCapture(control_video).read()[1]).size + else: + control_video = np.array(control_video.cpu().numpy() * 255, np.uint8) + original_width, original_height = Image.fromarray(control_video[0]).size + closest_size, closest_ratio = get_closest_ratio(original_height, original_width, ratios=aspect_ratio_sample_size) + height, width = [int(x / 16) * 16 for x in closest_size] # Load Sampler if scheduler == "DPM++": @@ -535,29 +568,47 @@ class CogVideoX_Fun_V2VSampler: pipeline.scheduler = noise_scheduler generator= torch.Generator(device).manual_seed(seed) - + with torch.no_grad(): video_length = int((video_length - 1) // pipeline.vae.config.temporal_compression_ratio * pipeline.vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 - input_video, input_video_mask, clip_image = get_video_to_video_latent(validation_video, video_length=video_length, sample_size=(height, width)) + if model_type == "Inpaint": + input_video, input_video_mask, clip_image = get_video_to_video_latent(validation_video, video_length=video_length, sample_size=(height, width), fps=8) + else: + input_video, input_video_mask, clip_image = get_video_to_video_latent(control_video, video_length=video_length, sample_size=(height, width), fps=8) for _lora_path, _lora_weight in zip(cogvideoxfun_model.get("loras", []), cogvideoxfun_model.get("strength_model", [])): pipeline = merge_lora(pipeline, _lora_path, _lora_weight) + + if model_type == "Inpaint": + sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = negative_prompt, + height = height, + width = width, + generator = generator, + guidance_scale = cfg, + num_inference_steps = steps, - sample = pipeline( - prompt, - num_frames = video_length, - negative_prompt = negative_prompt, - height = height, - width = width, - generator = generator, - guidance_scale = cfg, - num_inference_steps = steps, + video = input_video, + mask_video = input_video_mask, + strength = float(denoise_strength), + comfyui_progressbar = True, + ).videos + else: + sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = negative_prompt, + height = height, + width = width, + generator = generator, + guidance_scale = cfg, + num_inference_steps = steps, - video = input_video, - mask_video = input_video_mask, - strength = float(denoise_strength), - comfyui_progressbar = True, - ).videos + control_video = input_video, + comfyui_progressbar = True, + ).videos videos = rearrange(sample, "b c t h w -> (b t) h w c") for _lora_path, _lora_weight in zip(cogvideoxfun_model.get("loras", []), cogvideoxfun_model.get("strength_model", [])): diff --git a/comfyui/v1.1/cogvideoxfunv1.1_workflow_i2v.json b/comfyui/v1.1/cogvideoxfunv1.1_workflow_i2v.json new file mode 100644 index 0000000..b046768 --- /dev/null +++ b/comfyui/v1.1/cogvideoxfunv1.1_workflow_i2v.json @@ -0,0 +1,451 @@ +{ + "last_node_id": 83, + "last_link_id": 46, + "nodes": [ + { + "id": 7, + "type": "LoadImage", + "pos": [ + 258.76883544921907, + 468.15773315429715 + ], + "size": [ + 378.07147216796875, + 314.0000114440918 + ], + "flags": {}, + "order": 0, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 45 + ], + "shape": 3, + "label": "图像", + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": null, + "shape": 3, + "label": "遮罩" + } + ], + "title": "Start Image(图片到视频的开始图片)", + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "firework.png", + "image" + ] + }, + { + "id": 79, + "type": "Note", + "pos": [ + 16, + 460 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": {}, + "order": 1, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "You can upload image here\n(在此上传开始图像)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 80, + "type": "Note", + "pos": [ + 20, + -300 + ], + "size": { + "0": 210, + "1": 66.98204040527344 + }, + "flags": {}, + "order": 2, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "Load model here\n(在此选择要使用的模型)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": {}, + "order": 3, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 75, + "type": "CogVideoX_FUN_TextBox", + "pos": [ + 250, + -50 + ], + "size": { + "0": 383.54010009765625, + "1": 156.71620178222656 + }, + "flags": {}, + "order": 4, + "mode": 0, + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 43 + ], + "shape": 3, + "slot_index": 0 + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "CogVideoX_FUN_TextBox" + }, + "widgets_values": [ + "fireworks display over night city. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + ] + }, + { + "id": 82, + "type": "CogVideoX_Fun_I2VSampler", + "pos": [ + 758, + 93 + ], + "size": { + "0": 336, + "1": 282 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "cogvideoxfun_model", + "type": "CogVideoXFUNSMODEL", + "link": 42 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 43 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 44 + }, + { + "name": "start_img", + "type": "IMAGE", + "link": 45, + "slot_index": 3 + }, + { + "name": "end_img", + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 46 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "CogVideoX_Fun_I2VSampler" + }, + "widgets_values": [ + 49, + 512, + 43, + "fixed", + 50, + 6, + "DDIM" + ] + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1134, + 93 + ], + "size": [ + 390.9534912109375, + 535.9734235491071 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 46, + "label": "图像", + "slot_index": 0 + }, + { + "name": "audio", + "type": "VHS_AUDIO", + "link": null, + "label": "音频" + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null, + "label": "批次管理" + }, + { + "name": "vae", + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null, + "shape": 3, + "label": "文件名", + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 8, + "loop_count": 0, + "filename_prefix": "CogVideoX-Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "CogVideoX-Fun_00003.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 8 + } + } + } + }, + { + "id": 83, + "type": "LoadCogVideoX_Fun_Model", + "pos": [ + 300, + -294 + ], + "size": { + "0": 315, + "1": 130 + }, + "flags": {}, + "order": 5, + "mode": 0, + "outputs": [ + { + "name": "cogvideoxfun_model", + "type": "CogVideoXFUNSMODEL", + "links": [ + 42 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "LoadCogVideoX_Fun_Model" + }, + "widgets_values": [ + "CogVideoX-Fun-V1.1-2b-InP", + "Inpaint", + false, + "bf16" + ] + }, + { + "id": 73, + "type": "CogVideoX_FUN_TextBox", + "pos": [ + 250, + 160 + ], + "size": { + "0": 383.7149963378906, + "1": 183.83506774902344 + }, + "flags": {}, + "order": 6, + "mode": 0, + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 44 + ], + "shape": 3, + "slot_index": 0 + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "CogVideoX_FUN_TextBox" + }, + "widgets_values": [ + "The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " + ] + } + ], + "links": [ + [ + 42, + 83, + 0, + 82, + 0, + "CogVideoXFUNSMODEL" + ], + [ + 43, + 75, + 0, + 82, + 1, + "STRING_PROMPT" + ], + [ + 44, + 73, + 0, + 82, + 2, + "STRING_PROMPT" + ], + [ + 45, + 7, + 0, + 82, + 3, + "IMAGE" + ], + [ + 46, + 82, + 0, + 17, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24 + }, + { + "title": "Load CogVideoX-Fun", + "bounding": [ + 220, + -380, + 472, + 232 + ], + "color": "#b06634", + "font_size": 24 + }, + { + "title": "Upload Your Start Image", + "bounding": [ + 218, + 382, + 452, + 418 + ], + "color": "#a1309b", + "font_size": 24 + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7513148009015778, + "offset": [ + 268.77277812624413, + 436.3236112390962 + ] + }, + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/v1.1/cogvideoxfunv1.1_workflow_t2v.json b/comfyui/v1.1/cogvideoxfunv1.1_workflow_t2v.json new file mode 100644 index 0000000..1ac9c23 --- /dev/null +++ b/comfyui/v1.1/cogvideoxfunv1.1_workflow_t2v.json @@ -0,0 +1,359 @@ +{ + "last_node_id": 88, + "last_link_id": 52, + "nodes": [ + { + "id": 80, + "type": "Note", + "pos": [ + 20, + -300 + ], + "size": { + "0": 210, + "1": 66.98204040527344 + }, + "flags": {}, + "order": 0, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "Load model here\n(在此选择要使用的模型)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": {}, + "order": 1, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 75, + "type": "CogVideoX_FUN_TextBox", + "pos": [ + 250, + -50 + ], + "size": { + "0": 383.54010009765625, + "1": 156.71620178222656 + }, + "flags": {}, + "order": 2, + "mode": 0, + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 50 + ], + "shape": 3, + "slot_index": 0 + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "CogVideoX_FUN_TextBox" + }, + "widgets_values": [ + "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." + ] + }, + { + "id": 88, + "type": "CogVideoX_Fun_T2VSampler", + "pos": [ + 728, + -68 + ], + "size": { + "0": 327.6000061035156, + "1": 290 + }, + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "cogvideoxfun_model", + "type": "CogVideoXFUNSMODEL", + "link": 49 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 50 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 51, + "slot_index": 2 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 52 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "CogVideoX_Fun_T2VSampler" + }, + "widgets_values": [ + 49, + 672, + 384, + false, + 43, + "fixed", + 50, + 6, + "DDIM" + ] + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1110, + -67 + ], + "size": [ + 390.9534912109375, + 535.9734235491071 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 52, + "label": "图像", + "slot_index": 0 + }, + { + "name": "audio", + "type": "VHS_AUDIO", + "link": null, + "label": "音频" + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null, + "label": "批次管理" + }, + { + "name": "vae", + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null, + "shape": 3, + "label": "文件名", + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 8, + "loop_count": 0, + "filename_prefix": "CogVideoX-Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "CogVideoX-Fun_00004.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 8 + } + } + } + }, + { + "id": 73, + "type": "CogVideoX_FUN_TextBox", + "pos": [ + 250, + 160 + ], + "size": { + "0": 383.7149963378906, + "1": 183.83506774902344 + }, + "flags": {}, + "order": 3, + "mode": 0, + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 51 + ], + "shape": 3, + "slot_index": 0 + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "CogVideoX_FUN_TextBox" + }, + "widgets_values": [ + "The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " + ] + }, + { + "id": 87, + "type": "LoadCogVideoX_Fun_Model", + "pos": [ + 302, + -285 + ], + "size": { + "0": 315, + "1": 130 + }, + "flags": {}, + "order": 4, + "mode": 0, + "outputs": [ + { + "name": "cogvideoxfun_model", + "type": "CogVideoXFUNSMODEL", + "links": [ + 49 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "LoadCogVideoX_Fun_Model" + }, + "widgets_values": [ + "CogVideoX-Fun-V1.1-2b-InP", + "Inpaint", + false, + "bf16" + ] + } + ], + "links": [ + [ + 49, + 87, + 0, + 88, + 0, + "CogVideoXFUNSMODEL" + ], + [ + 50, + 75, + 0, + 88, + 1, + "STRING_PROMPT" + ], + [ + 51, + 73, + 0, + 88, + 2, + "STRING_PROMPT" + ], + [ + 52, + 88, + 0, + 17, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24 + }, + { + "title": "Load CogVideoX-Fun", + "bounding": [ + 220, + -380, + 472, + 232 + ], + "color": "#b06634", + "font_size": 24 + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.8264462809917354, + "offset": [ + 181.0702206286297, + 544.9672051634072 + ] + }, + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/v1.1/cogvideoxfunv1.1_workflow_v2v.json b/comfyui/v1.1/cogvideoxfunv1.1_workflow_v2v.json new file mode 100644 index 0000000..66e25a9 --- /dev/null +++ b/comfyui/v1.1/cogvideoxfunv1.1_workflow_v2v.json @@ -0,0 +1,492 @@ +{ + "last_node_id": 90, + "last_link_id": 57, + "nodes": [ + { + "id": 80, + "type": "Note", + "pos": [ + 20, + -300 + ], + "size": { + "0": 210, + "1": 66.98204040527344 + }, + "flags": {}, + "order": 0, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "Load model here\n(在此选择要使用的模型)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": {}, + "order": 1, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 79, + "type": "Note", + "pos": [ + 15.739953613281248, + 462.38664912015946 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": {}, + "order": 2, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "You can upload video here\n(在此上传视频)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 85, + "type": "VHS_LoadVideo", + "pos": [ + 336, + 470 + ], + "size": [ + 235.1999969482422, + 398.971426827567 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [ + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 56 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "frame_count", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "audio", + "type": "AUDIO", + "links": null, + "shape": 3 + }, + { + "name": "video_info", + "type": "VHS_VIDEOINFO", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "VHS_LoadVideo" + }, + "widgets_values": { + "video": "00000125.mp4", + "force_rate": 0, + "force_size": "Disabled", + "custom_width": 512, + "custom_height": 512, + "frame_load_cap": 0, + "skip_first_frames": 0, + "select_every_nth": 1, + "choose video to upload": "image", + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "frame_load_cap": 0, + "skip_first_frames": 0, + "force_rate": 0, + "filename": "00000125.mp4", + "type": "input", + "format": "video/mp4", + "select_every_nth": 1 + } + } + } + }, + { + "id": 88, + "type": "LoadCogVideoX_Fun_Model", + "pos": [ + 309, + -286 + ], + "size": { + "0": 315, + "1": 130 + }, + "flags": {}, + "order": 4, + "mode": 0, + "outputs": [ + { + "name": "cogvideoxfun_model", + "type": "CogVideoXFUNSMODEL", + "links": [ + 53 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "LoadCogVideoX_Fun_Model" + }, + "widgets_values": [ + "CogVideoX-Fun-V1.1-2b-InP", + "Inpaint", + false, + "bf16" + ] + }, + { + "id": 75, + "type": "CogVideoX_FUN_TextBox", + "pos": [ + 250, + -50 + ], + "size": { + "0": 383.54010009765625, + "1": 156.71620178222656 + }, + "flags": {}, + "order": 5, + "mode": 0, + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 54 + ], + "shape": 3, + "slot_index": 0 + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "CogVideoX_FUN_TextBox" + }, + "widgets_values": [ + "A cute cat is playing the guitar." + ] + }, + { + "id": 73, + "type": "CogVideoX_FUN_TextBox", + "pos": [ + 250, + 160 + ], + "size": { + "0": 383.7149963378906, + "1": 183.83506774902344 + }, + "flags": {}, + "order": 6, + "mode": 0, + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 55 + ], + "shape": 3, + "slot_index": 0 + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "CogVideoX_FUN_TextBox" + }, + "widgets_values": [ + "The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " + ] + }, + { + "id": 90, + "type": "CogVideoX_Fun_V2VSampler", + "pos": [ + 754, + 14 + ], + "size": { + "0": 317.4000244140625, + "1": 306 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "cogvideoxfun_model", + "type": "CogVideoXFUNSMODEL", + "link": 53 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 54 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 55 + }, + { + "name": "validation_video", + "type": "IMAGE", + "link": 56, + "slot_index": 3 + }, + { + "name": "control_video", + "type": "IMAGE", + "link": null + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 57 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "CogVideoX_Fun_V2VSampler" + }, + "widgets_values": [ + 49, + 768, + 43, + "randomize", + 50, + 6, + 0.7, + "DDIM" + ] + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1125, + 15 + ], + "size": [ + 390.9534912109375, + 535.9734235491071 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 57, + "label": "图像", + "slot_index": 0 + }, + { + "name": "audio", + "type": "VHS_AUDIO", + "link": null, + "label": "音频" + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null, + "label": "批次管理" + }, + { + "name": "vae", + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null, + "shape": 3, + "label": "文件名", + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 8, + "loop_count": 0, + "filename_prefix": "EasyAnimate", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "EasyAnimate_00045.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 8 + } + } + } + } + ], + "links": [ + [ + 53, + 88, + 0, + 90, + 0, + "CogVideoXFUNSMODEL" + ], + [ + 54, + 75, + 0, + 90, + 1, + "STRING_PROMPT" + ], + [ + 55, + 73, + 0, + 90, + 2, + "STRING_PROMPT" + ], + [ + 56, + 85, + 0, + 90, + 3, + "IMAGE" + ], + [ + 57, + 90, + 0, + 17, + 0, + "IMAGE" + ] + ], + "groups": [ + { + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24 + }, + { + "title": "Load CogVideoX-Fun", + "bounding": [ + 220, + -380, + 472, + 232 + ], + "color": "#b06634", + "font_size": 24 + }, + { + "title": "Upload Your Video", + "bounding": [ + 218, + 385, + 456, + 498 + ], + "color": "#a1309b", + "font_size": 24 + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.683013455365071, + "offset": [ + 314.4077746994681, + 444.69453403364594 + ] + }, + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/v1.1/cogvideoxfunv1.1_workflow_v2v_control.json b/comfyui/v1.1/cogvideoxfunv1.1_workflow_v2v_control.json new file mode 100644 index 0000000..dc7de7a --- /dev/null +++ b/comfyui/v1.1/cogvideoxfunv1.1_workflow_v2v_control.json @@ -0,0 +1,492 @@ +{ + "last_node_id": 90, + "last_link_id": 59, + "nodes": [ + { + "id": 80, + "type": "Note", + "pos": [ + 20, + -300 + ], + "size": { + "0": 210, + "1": 66.98204040527344 + }, + "flags": {}, + "order": 0, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "Load model here\n(在此选择要使用的模型)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 78, + "type": "Note", + "pos": [ + 18, + -46 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": {}, + "order": 1, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "You can write prompt here\n(你可以在此填写提示词)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 79, + "type": "Note", + "pos": [ + 15.739953613281248, + 462.38664912015946 + ], + "size": { + "0": 210, + "1": 58 + }, + "flags": {}, + "order": 2, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "You can upload video here\n(在此上传视频)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 73, + "type": "CogVideoX_FUN_TextBox", + "pos": [ + 250, + 160 + ], + "size": { + "0": 383.7149963378906, + "1": 183.83506774902344 + }, + "flags": {}, + "order": 3, + "mode": 0, + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 55 + ], + "shape": 3, + "slot_index": 0 + } + ], + "title": "Negtive Prompt(反向提示词)", + "properties": { + "Node name for S&R": "CogVideoX_FUN_TextBox" + }, + "widgets_values": [ + "The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " + ] + }, + { + "id": 17, + "type": "VHS_VideoCombine", + "pos": [ + 1125, + 15 + ], + "size": [ + 390.9534912109375, + 973.1686096191406 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 57, + "label": "图像", + "slot_index": 0 + }, + { + "name": "audio", + "type": "VHS_AUDIO", + "link": null, + "label": "音频" + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null, + "label": "批次管理" + }, + { + "name": "vae", + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null, + "shape": 3, + "label": "文件名", + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 8, + "loop_count": 0, + "filename_prefix": "CogVideoX-Fun", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 22, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "CogVideoX-Fun_00007.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 8 + } + } + } + }, + { + "id": 88, + "type": "LoadCogVideoX_Fun_Model", + "pos": [ + 309, + -286 + ], + "size": { + "0": 315, + "1": 130 + }, + "flags": {}, + "order": 4, + "mode": 0, + "outputs": [ + { + "name": "cogvideoxfun_model", + "type": "CogVideoXFUNSMODEL", + "links": [ + 53 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "LoadCogVideoX_Fun_Model" + }, + "widgets_values": [ + "CogVideoX-Fun-V1.1-2b-Pose", + "Control", + false, + "bf16" + ] + }, + { + "id": 85, + "type": "VHS_LoadVideo", + "pos": [ + 336, + 470 + ], + "size": [ + 235.1999969482422, + 658.5777723524305 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 59 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "frame_count", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "audio", + "type": "AUDIO", + "links": null, + "shape": 3 + }, + { + "name": "video_info", + "type": "VHS_VIDEOINFO", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "VHS_LoadVideo" + }, + "widgets_values": { + "video": "pose.mp4", + "force_rate": 8, + "force_size": "Disabled", + "custom_width": 512, + "custom_height": 512, + "frame_load_cap": 0, + "skip_first_frames": 0, + "select_every_nth": 1, + "choose video to upload": "image", + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "frame_load_cap": 0, + "skip_first_frames": 0, + "force_rate": 8, + "filename": "pose.mp4", + "type": "input", + "format": "video/mp4", + "select_every_nth": 1 + } + } + } + }, + { + "id": 75, + "type": "CogVideoX_FUN_TextBox", + "pos": [ + 250, + -50 + ], + "size": { + "0": 383.54010009765625, + "1": 156.71620178222656 + }, + "flags": {}, + "order": 6, + "mode": 0, + "outputs": [ + { + "name": "prompt", + "type": "STRING_PROMPT", + "links": [ + 54 + ], + "shape": 3, + "slot_index": 0 + } + ], + "title": "Positive Prompt(正向提示词)", + "properties": { + "Node name for S&R": "CogVideoX_FUN_TextBox" + }, + "widgets_values": [ + "A person wearing a knee-length white sleeveless dress and white high-heeled sandals performs a dance in a well-lit room with wooden flooring. The room's background features a closed door, a shelf displaying clear glass bottles of alcoholic beverages, and a partially visible dark-colored sofa. " + ] + }, + { + "id": 90, + "type": "CogVideoX_Fun_V2VSampler", + "pos": [ + 754, + 14 + ], + "size": { + "0": 336, + "1": 306 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "cogvideoxfun_model", + "type": "CogVideoXFUNSMODEL", + "link": 53 + }, + { + "name": "prompt", + "type": "STRING_PROMPT", + "link": 54 + }, + { + "name": "negative_prompt", + "type": "STRING_PROMPT", + "link": 55 + }, + { + "name": "validation_video", + "type": "IMAGE", + "link": null, + "slot_index": 3 + }, + { + "name": "control_video", + "type": "IMAGE", + "link": 59 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 57 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "CogVideoX_Fun_V2VSampler" + }, + "widgets_values": [ + 49, + 512, + 43, + "fixed", + 50, + 6, + 1, + "DDIM" + ] + } + ], + "links": [ + [ + 53, + 88, + 0, + 90, + 0, + "CogVideoXFUNSMODEL" + ], + [ + 54, + 75, + 0, + 90, + 1, + "STRING_PROMPT" + ], + [ + 55, + 73, + 0, + 90, + 2, + "STRING_PROMPT" + ], + [ + 57, + 90, + 0, + 17, + 0, + "IMAGE" + ], + [ + 59, + 85, + 0, + 90, + 4, + "IMAGE" + ] + ], + "groups": [ + { + "title": "Prompts", + "bounding": [ + 218, + -127, + 450, + 483 + ], + "color": "#3f789e", + "font_size": 24 + }, + { + "title": "Load CogVideoX-Fun", + "bounding": [ + 220, + -380, + 472, + 232 + ], + "color": "#b06634", + "font_size": 24 + }, + { + "title": "Upload Your Video", + "bounding": [ + 218, + 385, + 457, + 776 + ], + "color": "#a1309b", + "font_size": 24 + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6830134553650712, + "offset": [ + 250.2298948633902, + 399.72391778748613 + ] + }, + "workspace_info": { + "id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea" + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/comfyui/v1/cogvideoxfunv1_workflow_v2v.json b/comfyui/v1/cogvideoxfunv1_workflow_v2v.json index 54f9849..c37af56 100644 --- a/comfyui/v1/cogvideoxfunv1_workflow_v2v.json +++ b/comfyui/v1/cogvideoxfunv1_workflow_v2v.json @@ -217,7 +217,7 @@ "Node name for S&R": "CogVideoX_FUN_TextBox" }, "widgets_values": [ - "A beautiful woman is playing the guitar. The video quality is high and the picture is clear. High quality, masterpiece, the best quality, high resolution, ultra careful." + "A cute cat is playing the guitar." ] }, { diff --git a/predict_i2v.py b/predict_i2v.py index ecd1cb5..fce07a5 100644 --- a/predict_i2v.py +++ b/predict_i2v.py @@ -24,7 +24,7 @@ from cogvideox.utils.utils import get_image_to_video_latent, save_videos_grid low_gpu_memory_mode = False # Config and model path -model_name = "models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" +model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP" # Choose the sampler in "Euler" "Euler A" "DPM++" "PNDM" "DDIM_Cog" and "DDIM_Origin" sampler_name = "DDIM_Origin" @@ -52,7 +52,7 @@ validation_image_end = None # prompts prompt = "The dog is shaking head. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." -negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " +negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. " guidance_scale = 6.0 seed = 43 num_inference_steps = 50 @@ -160,7 +160,7 @@ if partial_video_length is not None: with torch.no_grad(): sample = pipeline( - prompt + ". The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic. ", + prompt, num_frames = _partial_video_length, negative_prompt = negative_prompt, height = sample_size[0], diff --git a/predict_t2v.py b/predict_t2v.py index 8251d1d..f94f6c2 100644 --- a/predict_t2v.py +++ b/predict_t2v.py @@ -24,7 +24,7 @@ from cogvideox.utils.utils import get_image_to_video_latent, save_videos_grid low_gpu_memory_mode = False # model path -model_name = "models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" +model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP" # Choose the sampler in "Euler" "Euler A" "DPM++" "PNDM" and "DDIM" sampler_name = "DDIM_Origin" @@ -43,7 +43,7 @@ fps = 8 # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 weight_dtype = torch.bfloat16 prompt = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." -negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " +negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. " guidance_scale = 6.0 seed = 43 num_inference_steps = 50 diff --git a/predict_v2v.py b/predict_v2v.py index 552db61..62ab21a 100644 --- a/predict_v2v.py +++ b/predict_v2v.py @@ -25,7 +25,7 @@ from cogvideox.utils.utils import get_video_to_video_latent, save_videos_grid low_gpu_memory_mode = False # model path -model_name = "models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" +model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP" # Choose the sampler in "Euler" "Euler A" "DPM++" "PNDM" and "DDIM" sampler_name = "DDIM_Origin" @@ -42,13 +42,17 @@ fps = 8 # Use torch.float16 if GPU does not support torch.bfloat16 # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 weight_dtype = torch.bfloat16 -# If you want to generate from text, please set the validation_image_start = None and validation_image_end = None -validation_video = "asset/03480_03600_good_scenes.mp4" +# If you are preparing to redraw the reference video, set validation_video and validation_video_mask. +# If you do not use validation_video_mask, the entire video will be redrawn; +# if you use validation_video_mask, only a portion of the video will be redrawn. +# Please set a larger denoise_strength when using validation_video_mask, such as 1.00 instead of 0.70 +validation_video = "asset/1.mp4" +validation_video_mask = None denoise_strength = 0.70 # prompts -prompt = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic." -negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. Strange motion trajectory. " +prompt = "A cute cat is playing the guitar. " +negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. " guidance_scale = 6.0 seed = 43 num_inference_steps = 50 @@ -137,7 +141,7 @@ if lora_path is not None: pipeline = merge_lora(pipeline, lora_path, lora_weight, "cuda") video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 -input_video, input_video_mask, clip_image = get_video_to_video_latent(validation_video, video_length=video_length, sample_size=sample_size) +input_video, input_video_mask, clip_image = get_video_to_video_latent(validation_video, video_length=video_length, sample_size=sample_size, validation_video_mask=validation_video_mask, fps=fps) with torch.no_grad(): sample = pipeline( diff --git a/predict_v2v_control.py b/predict_v2v_control.py new file mode 100644 index 0000000..2628e0c --- /dev/null +++ b/predict_v2v_control.py @@ -0,0 +1,163 @@ +import json +import os + +import cv2 +import numpy as np +import torch +from diffusers import (AutoencoderKL, CogVideoXDDIMScheduler, DDIMScheduler, + DPMSolverMultistepScheduler, + EulerAncestralDiscreteScheduler, EulerDiscreteScheduler, + PNDMScheduler) +from omegaconf import OmegaConf +from PIL import Image +from transformers import (CLIPImageProcessor, CLIPVisionModelWithProjection, + T5EncoderModel, T5Tokenizer) + +from cogvideox.models.autoencoder_magvit import AutoencoderKLCogVideoX +from cogvideox.models.transformer3d import CogVideoXTransformer3DModel +from cogvideox.pipeline.pipeline_cogvideox import CogVideoX_Fun_Pipeline +from cogvideox.pipeline.pipeline_cogvideox_control import \ + CogVideoX_Fun_Pipeline_Control +from cogvideox.utils.lora_utils import merge_lora, unmerge_lora +from cogvideox.utils.utils import get_video_to_video_latent, save_videos_grid + +# Low gpu memory mode, this is used when the GPU memory is under 16GB +low_gpu_memory_mode = False + +# model path +model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-Pose" + +# Choose the sampler in "Euler" "Euler A" "DPM++" "PNDM" and "DDIM" +sampler_name = "DDIM_Origin" + +# Load pretrained model if need +transformer_path = None +vae_path = None +lora_path = None +# Other params +sample_size = [672, 384] +video_length = 49 +fps = 8 + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +control_video = "asset/pose.mp4" + +# prompts +prompt = "A person wearing a knee-length white sleeveless dress and white high-heeled sandals performs a dance in a well-lit room with wooden flooring. The room's background features a closed door, a shelf displaying clear glass bottles of alcoholic beverages, and a partially visible dark-colored sofa. " +negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. " +guidance_scale = 6.0 +seed = 43 +num_inference_steps = 50 +lora_weight = 0.55 +save_path = "samples/cogvideox-fun-videos_control" + +transformer = CogVideoXTransformer3DModel.from_pretrained_2d( + model_name, + subfolder="transformer", +).to(weight_dtype) + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Vae +vae = AutoencoderKLCogVideoX.from_pretrained( + model_name, + subfolder="vae" +).to(weight_dtype) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +text_encoder = T5EncoderModel.from_pretrained( + model_name, subfolder="text_encoder", torch_dtype=weight_dtype +) +# Get Scheduler +Choosen_Scheduler = scheduler_dict = { + "Euler": EulerDiscreteScheduler, + "Euler A": EulerAncestralDiscreteScheduler, + "DPM++": DPMSolverMultistepScheduler, + "PNDM": PNDMScheduler, + "DDIM_Cog": CogVideoXDDIMScheduler, + "DDIM_Origin": DDIMScheduler, +}[sampler_name] +scheduler = Choosen_Scheduler.from_pretrained( + model_name, + subfolder="scheduler" +) + +pipeline = CogVideoX_Fun_Pipeline_Control.from_pretrained( + model_name, + vae=vae, + text_encoder=text_encoder, + transformer=transformer, + scheduler=scheduler, + torch_dtype=weight_dtype +) + +if low_gpu_memory_mode: + pipeline.enable_sequential_cpu_offload() +else: + pipeline.enable_model_cpu_offload() + +generator = torch.Generator(device="cuda").manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, "cuda") + +video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 +input_video, input_video_mask, clip_image = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps) + +with torch.no_grad(): + sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = negative_prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = guidance_scale, + num_inference_steps = num_inference_steps, + + control_video = input_video, + ).videos + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, "cuda") + +if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + +index = len([path for path in os.listdir(save_path)]) + 1 +prefix = str(index).zfill(8) + +if video_length == 1: + save_sample_path = os.path.join(save_path, prefix + f".png") + + image = sample[0, :, 0] + image = image.transpose(0, 1).transpose(1, 2) + image = (image * 255).numpy().astype(np.uint8) + image = Image.fromarray(image) + image.save(save_sample_path) +else: + video_path = os.path.join(save_path, prefix + ".mp4") + save_videos_grid(sample, video_path, fps=fps) \ No newline at end of file diff --git a/reports/report_v1_1.md b/reports/report_v1_1.md new file mode 100644 index 0000000..b716d8c --- /dev/null +++ b/reports/report_v1_1.md @@ -0,0 +1,32 @@ +# CogVideoX FUN v1.1 Report + +In CogVideoX-FUN v1.1, we performed additional filtering on the previous dataset, selecting videos with larger motion amplitudes rather than still images in motion, resulting in approximately 0.48 million videos. The model continues to support both image and video prediction, accommodating pixel values from 512x512x49, 768x768x49, 1024x1024x49, and videos with different aspect ratios. We support both image-to-video generation and video-to-video reconstruction. + +Additionally, we have released training and prediction code for adding control signals, along with the initial version of the Control model. + +Compared to version 1.0, CogVideoX-FUN V1.1 highlights the following features: +- In the 5b model, Noise has been added to the reference images, increasing the motion amplitude of the videos. +- Released training and prediction code for adding control signals, along with the initial version of the Control model. + +## Adding Noise to Reference Images + +Building on the original CogVideoX-FUN V1.0, we drew upon [CogVideoX](https://github.com/THUDM/CogVideo/) and [SVD](https://github.com/Stability-AI/generative-models) to add Noise upwards to the non-zero reference images to disrupt the original images, aiming for greater motion amplitude. + +In our 5b model, Noise has been added, while the 2b model only performed fine-tuning with new data. This is because, after attempting to add Noise in the 2b model, the generated videos exhibited excessive motion amplitude, leading to deformation and damaging the output. The 5b model, due to its stronger generative capabilities, maintains relatively stable outputs during motion. + +Furthermore, the prompt words significantly influence the generation results, so please describe the actions in detail to increase dynamism. If unsure how to write positive prompts, you can use phrases like "smooth motion" or "in the wind" to enhance dynamism. Additionally, it is advisable to avoid using dynamic terms like "motion" in negative prompts. + +## Adding Control Signals to CogVideoX-FUN + +On the basis of the original CogVideoX-FUN V1.0, we replaced the original mask signal with Pose control signals. The control signals are encoded using VAE and used as Guidance, along with latent data entering the patch processing flow. + +We filtered the 0.48 million dataset, selecting around 20,000 videos and images containing portraits for pose extraction, which served as condition control signals for training. + +During the training process, the videos are scaled according to different Token lengths. The entire training process is divided into two phases, with each phase comprising 13,312 (corresponding to 512x512x49 videos) and 53,248 (corresponding to 1024x1024x49 videos). + +Taking CogVideoX-Fun-V1.1-5b-Pose as an example: +- In the 13312 phase, the batch size is 128, with 2.4k training steps. +- In the 53248 phase, the batch size is 128, with 1.2k training steps. + +The working principle diagram is shown below: +ui diff --git a/reports/report_v1_1_zh-CN.md b/reports/report_v1_1_zh-CN.md new file mode 100644 index 0000000..81dfb7a --- /dev/null +++ b/reports/report_v1_1_zh-CN.md @@ -0,0 +1,31 @@ +# CogVideoX FUN v1.1 Report + +在CogVideoX-FUN v1.1中,我们在之前的数据集中再次做了筛选,选出其中动作幅度较大,而不是静止画面移动的视频,数量大约为0.48m。模型依然支持图片与视频预测,支持像素值从512x512x49、768x768x49、1024x1024x49与不同纵横比的视频生成。我们支持图像到视频的生成与视频到视频的重建。 + +另外,我们还发布了添加控制信号的训练代码与预测代码,并发布了初版的Control模型。 + +对比V1.0版本,CogVideoX-FUN V1.1突出了以下功能: + +- 在5b模型中,给参考图片添加了Noise,增加了视频的运动幅度。 +- 发布了添加控制信号的训练代码与预测代码,并发布了初版的Control模型。 + +## 参考图片添加Noise +在原本CogVideoX-FUN V1.0的基础上,我们参考[CogVideoX](https://github.com/THUDM/CogVideo/)和[SVD](https://github.com/Stability-AI/generative-models),在非0的参考图向上添加Noise以破环原图,追求更大的运动幅度。 + +我们5b模型中添加了Noise,2b模型仅使用了新数据进行了finetune,因为我们在2b模型中尝试添加Noise之后,生成的视频运动幅度过大导致结果变形,破坏了生成结果,而5b模型因为更为的强大生成能力,在运动中也保持了较为稳定的输出。 + +另外,提示词对生成结果影响较大,请尽量描写动作以增加动态性。如果不知道怎么写正向提示词,可以使用smooth motion or in the wind来增加动态性。并且尽量避免在负向提示词中出现motion等表示动态的词汇。 + +## 添加控制信号的CogVideoX-Fun +在原本CogVideoX-FUN V1.0的基础上,我们使用Pose控制信号替代了原本的mask信号,将控制信号使用VAE编码后作为Guidance与latent一起进入patch流程, + +我们在0.48m数据中进行了筛选,选择出大约20000包含人像的视频与图片进行pose提取,作为condition控制信号进行训练。 + +在进行训练时,我们根据不同Token长度,对视频进行缩放后进行训练。整个训练过程分为两个阶段,每个阶段的13312(对应512x512x49的视频),53248(对应1024x1024x49的视频)。 + +以CogVideoX-Fun-V1.1-5b-Pose为例子,其中: +- 13312阶段,Batch size为128,训练步数为2.4k +- 53248阶段,Batch size为128,训练步数为1.2k。 + +工作原理图如下: +ui diff --git a/scripts/README_DEMO.md b/scripts/README_DEMO.md new file mode 100644 index 0000000..0d0032c --- /dev/null +++ b/scripts/README_DEMO.md @@ -0,0 +1,37 @@ +## Demo + +Image generation video corresponding images and prompts. + +If you don't know how to write positive prompts, you can use "smooth motion" or "in the wind" to add dynamism. + +| Image | Prompt | +|--|--| +| ![1.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/1.png) | closeup face photo of man is smiling in black clothes, night city street, bokeh, fireworks in background | +| ![2.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/2.png) | sunset, orange sky, warm lighting, fishing boats, ocean waves, seagulls, rippling water, wharf, silhouette, serene atmosphere, dusk, evening glow, golden hour, coastal landscape, seaside scenery | +| ![3.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/3.png) | a man in an astronaut suit playing a guitar | +| ![4.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/4.png) | time-lapse of a blooming flower with leaves and a stem, blossom | +| ![5.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/5.png) | fireworks display over night city | +| ![6.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/6.png) | a beautiful woman with long hair and a dress blowing in the wind | +| ![7.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/7.png) | the dog is shaking head | +| ![8.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/8.png) | a robot is walking through a destroyed city | +| ![9.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/9.png) | a group of penguins walking on a beach | +| ![10.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/10.png) | a bonfire is lit in the middle of a field | +| ![11.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/11.png) | a boat traveling on the ocean | +| ![12.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/12.png) | pouring honey onto some slices of bread | +| ![13.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/13.png) | a sailboat sailing in rough seas with a dramatic sunset | +| ![14.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/14.png) | a boat traveling on the ocean | +| ![15.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/15.png) | a scenic view of a lake with several seagulls flying above the water. In the foreground, there is a person wearing a red garment, possibly a jacket or a shawl, observing the scenery. The lake has clear blue water, and there's a structure that appears to be a wooden pavilion or boathouse on stilts situated in the water. In the background, hills or mountains can be seen under a clear blue sky, enhancing the tranquil and picturesque setting | +| ![16.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/16.png) | A man's body shimmered with golden light in the wind | +| ![17.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/17.png) | a buried broken emerald cross glazed by the sun emitting smoke, backlit, forgotten, atmospheric AF, detailed, 8k | +| ![18.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/18.png) | A beautiful woman is smiling in the wind | +| ![19.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/19.png) | A beautiful woman is smiling in the wind | +| ![20.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/20.png) | A beautiful woman is smiling in the wind | +| ![21.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/21.png) | A beautiful woman smiles in the heavy snow | +| ![22.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/22.png) | cats smiling taking a selfie with a super wide angle lenses, opening mouth. | +| ![23.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/23.png) | The sturdy sailboat in 'Temperamental Tides', masterfully navigating the restless, pulsating waves of the deep navy sea, maintaining balance on the surging storm grey crests | +| ![24.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/24.png) | The sturdy sailboat in 'Temperamental Tides', masterfully navigating the restless, pulsating waves of the deep navy sea, maintaining balance on the surging storm grey crests | +| ![25.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/25.png) | a beach with waves crashing against it and a sunset in the background a brigantine, a sailboat in the distance, 4k | +| ![26.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/26.png) | Create an illustration that captures the essence of water. The scene should be a tranquil beach at sunrise, with the calm ocean stretching out to the horizon. The sky is painted in soft hues of pink and orange as the sun begins to rise. Gentle waves lap against the sandy shore, creating delicate ripples. The water is crystal clear, reflecting the colors of the sky, and small, glistening seashells are scattered along the shoreline. In the distance, a small sailboat with white sails drifts peacefully on the water. The overall mood of the illustration should be serene and calming, emphasizing the fluid and reflective nature of water.glowneon, glowing, sparks, lightning | +| ![27.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/27.png) | a Lighthouse battered by high winds, huge crashing waves, realistic northern lights, behind lighthouse, realistic stormy seas, high quality image, photographic, mist, and sea spray, storm clouds, angry sky, dusk, peninsula, winter, almost dark, storm, gales, elevated view point, high up perspective, night time, lighthouse light beams, position lighthouse to left of image, view from on high | +| ![28.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/28.png) | Two racing cars racing towards the camera, desert dune in the background, hyperrealistic, driver turning the wheel, more details, speed of light, a trail of intense light follows the cars, image evokes the sensation of speed, frozen movement, insane intricate detail, (masterpiece, best quality), high resolution, (ultra detailed), | +| ![29.png](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/i2v_images/29.png) | one long eared dog, beagle, making goofy faces under water, lie the wind is blowing in his open mouth bubbles wide. an annoyed goldfish swims by | \ No newline at end of file diff --git a/scripts/README_TRAIN.md b/scripts/README_TRAIN.md index 93b6fbc..b109de5 100644 --- a/scripts/README_TRAIN.md +++ b/scripts/README_TRAIN.md @@ -4,6 +4,18 @@ The default training commands for the different versions are as follows: We can choose whether to use deep speed in CogVideoX-Fun, which can save a lot of video memory. +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution. +- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts. +- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When random_hw_adapt is enabled, the training images will have their height and width set to image_sample_size as the maximum and video_sample_size as the minimum. For training videos, the height and width will be set to video_sample_size as the maximum and min(video_sample_size, 512) as the minimum. +- `training_with_video_token_length` specifies training the model according to token length. The token length for a video with dimensions 512x512 and 49 frames is 13,312. + - At 512x512 resolution, the number of video frames is 49; + - At 768x768 resolution, the number of video frames is 21; + - At 1024x1024 resolution, the number of video frames is 9; + - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. +- `train_mode` is used to specify the training mode, which can be either normal or inpaint. Since CogVideoX-Fun uses the Inpaint model to achieve text-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. + CogVideoX-Fun without deepspeed: ```sh export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" @@ -18,7 +30,7 @@ accelerate launch --mixed_precision="bf16" scripts/train.py \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ --train_data_meta=$DATASET_META_NAME \ - --image_sample_size=1280 \ + --image_sample_size=1024 \ --video_sample_size=256 \ --token_sample_size=512 \ --video_sample_stride=3 \ @@ -45,7 +57,6 @@ accelerate launch --mixed_precision="bf16" scripts/train.py \ --random_frame_crop \ --enable_bucket \ --use_came \ - --use_ema \ --train_mode="inpaint" \ --resume_from_checkpoint="latest" \ --trainable_modules "." @@ -64,7 +75,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ --train_data_meta=$DATASET_META_NAME \ - --image_sample_size=1280 \ + --image_sample_size=1024 \ --video_sample_size=256 \ --token_sample_size=512 \ --video_sample_stride=3 \ @@ -92,7 +103,6 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --enable_bucket \ --use_came \ --use_deepspeed \ - --use_ema \ --train_mode="inpaint" \ --resume_from_checkpoint="latest" \ --trainable_modules "." diff --git a/scripts/README_TRAIN_CONTROL.md b/scripts/README_TRAIN_CONTROL.md new file mode 100644 index 0000000..14a1119 --- /dev/null +++ b/scripts/README_TRAIN_CONTROL.md @@ -0,0 +1,126 @@ +## Training Code + +The default training commands for the different versions are as follows: + +We can choose whether to use deep speed in CogVideoX-Fun, which can save a lot of video memory. + +The metadata_control.json is a little different from normal json in CogVideoX-Fun, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file. + +```json +[ + { + "file_path": "train/00000001.mp4", + "control_file_path": "control/00000001.mp4", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "video" + }, + { + "file_path": "train/00000002.jpg", + "control_file_path": "control/00000002.jpg", + "text": "A group of young men in suits and sunglasses are walking down a city street.", + "type": "image" + }, + ..... +] +``` + +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution. +- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts. +- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When random_hw_adapt is enabled, the training images will have their height and width set to image_sample_size as the maximum and video_sample_size as the minimum. For training videos, the height and width will be set to video_sample_size as the maximum and min(video_sample_size, 512) as the minimum. +- `training_with_video_token_length` specifies training the model according to token length. The token length for a video with dimensions 512x512 and 49 frames is 13,312. + - At 512x512 resolution, the number of video frames is 49; + - At 768x768 resolution, the number of video frames is 21; + - At 1024x1024 resolution, the number of video frames is 9; + - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. + +CogVideoX-Fun without deepspeed: +```sh +export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-Pose" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json" +export NCCL_IB_DISABLE=1 +export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'". +accelerate launch --mixed_precision="bf16" scripts/train_control.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --image_sample_size=1024 \ + --video_sample_size=256 \ + --token_sample_size=512 \ + --video_sample_stride=3 \ + --video_sample_n_frames=49 \ + --train_batch_size=4 \ + --video_repeat=1 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=50 \ + --seed=43 \ + --output_dir="output_dir" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --random_hw_adapt \ + --training_with_video_token_length \ + --random_frame_crop \ + --enable_bucket \ + --use_came \ + --resume_from_checkpoint="latest" \ + --trainable_modules "." +``` + +CogVideoX-Fun with deepspeed: +```sh +export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-Pose" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json" +export NCCL_IB_DISABLE=1 +export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/train.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --image_sample_size=1024 \ + --video_sample_size=256 \ + --token_sample_size=512 \ + --video_sample_stride=3 \ + --video_sample_n_frames=49 \ + --train_batch_size=4 \ + --video_repeat=1 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=50 \ + --seed=43 \ + --output_dir="output_dir" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --random_hw_adapt \ + --training_with_video_token_length \ + --random_frame_crop \ + --enable_bucket \ + --use_came \ + --use_deepspeed \ + --resume_from_checkpoint="latest" \ + --trainable_modules "." +``` diff --git a/scripts/README_TRAIN_LORA.md b/scripts/README_TRAIN_LORA.md index 157c3fb..9ee5dfe 100644 --- a/scripts/README_TRAIN_LORA.md +++ b/scripts/README_TRAIN_LORA.md @@ -2,6 +2,18 @@ We can choose whether to use deep speed in CogVideoX-Fun, which can save a lot of video memory. +Some parameters in the sh file can be confusing, and they are explained in this document: + +- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images and videos at the center, but instead, it trains the entire images and videos after grouping them into buckets based on resolution. +- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts. +- `random_hw_adapt` is used to enable automatic height and width scaling for images and videos. When random_hw_adapt is enabled, the training images will have their height and width set to image_sample_size as the maximum and video_sample_size as the minimum. For training videos, the height and width will be set to video_sample_size as the maximum and min(video_sample_size, 512) as the minimum. +- `training_with_video_token_length` specifies training the model according to token length. The token length for a video with dimensions 512x512 and 49 frames is 13,312. + - At 512x512 resolution, the number of video frames is 49; + - At 768x768 resolution, the number of video frames is 21; + - At 1024x1024 resolution, the number of video frames is 9; + - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. +- `train_mode` is used to specify the training mode, which can be either normal or inpaint. Since CogVideoX-Fun uses the Inpaint model to achieve text-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. + CogVideoX-Fun without deepspeed: ```sh @@ -17,7 +29,7 @@ accelerate launch --mixed_precision="bf16" scripts/train_lora.py \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ --train_data_meta=$DATASET_META_NAME \ - --image_sample_size=1280 \ + --image_sample_size=1024 \ --video_sample_size=256 \ --token_sample_size=512 \ --video_sample_stride=3 \ @@ -58,7 +70,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ --train_data_meta=$DATASET_META_NAME \ - --image_sample_size=1280 \ + --image_sample_size=1024 \ --video_sample_size=256 \ --token_sample_size=512 \ --video_sample_stride=3 \ diff --git a/scripts/train.py b/scripts/train.py index 3e7dd5e..bf9a0f1 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -75,7 +75,7 @@ from cogvideox.models.autoencoder_magvit import AutoencoderKLCogVideoX from cogvideox.models.transformer3d import CogVideoXTransformer3DModel from cogvideox.pipeline.pipeline_cogvideox import CogVideoX_Fun_Pipeline from cogvideox.pipeline.pipeline_cogvideox_inpaint import \ - CogVideoX_Fun_Pipeline_Inpaint + CogVideoX_Fun_Pipeline_Inpaint, add_noise_to_reference_video from cogvideox.utils.utils import get_image_to_video_latent, save_videos_grid if is_wandb_available(): @@ -1433,6 +1433,8 @@ def main(): mask = 1 - mask mask = resize_mask(mask, latents) + if unwrap_model(transformer3d).config.add_noise_in_inpaint_model: + mask_pixel_values = add_noise_to_reference_video(mask_pixel_values) mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> b c f h w") bs = args.vae_mini_batch new_mask_pixel_values = [] diff --git a/scripts/train_control.py b/scripts/train_control.py new file mode 100644 index 0000000..217a765 --- /dev/null +++ b/scripts/train_control.py @@ -0,0 +1,1653 @@ +"""Modified from https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +import argparse +import gc +import logging +import math +import os +import pickle +import shutil +import sys + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +import transformers +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import AutoencoderKL, DDPMScheduler +from diffusers.models.embeddings import get_3d_rotary_pos_embed +from diffusers.optimization import get_scheduler +from diffusers.training_utils import EMAModel +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.import_utils import is_xformers_available +from diffusers.utils.torch_utils import is_compiled_module +from einops import rearrange +from huggingface_hub import create_repo, upload_folder +from omegaconf import OmegaConf +from packaging import version +from PIL import Image +from torch.utils.data import RandomSampler +from torch.utils.tensorboard import SummaryWriter +from torchvision import transforms +from tqdm.auto import tqdm +from transformers import (BertModel, BertTokenizer, CLIPImageProcessor, + CLIPVisionModelWithProjection, MT5Tokenizer, + T5EncoderModel, T5Tokenizer) +from transformers.utils import ContextManagers + +import datasets + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from cogvideox.data.bucket_sampler import (ASPECT_RATIO_512, + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) +from cogvideox.data.dataset_image_video import (ImageVideoDataset, ImageVideoControlDataset, + ImageVideoSampler, + get_random_mask) +from cogvideox.models.autoencoder_magvit import AutoencoderKLCogVideoX +from cogvideox.models.transformer3d import CogVideoXTransformer3DModel +from cogvideox.pipeline.pipeline_cogvideox import CogVideoX_Fun_Pipeline +from cogvideox.pipeline.pipeline_cogvideox_inpaint import \ + CogVideoX_Fun_Pipeline_Inpaint, add_noise_to_reference_video +from cogvideox.utils.utils import get_image_to_video_latent, save_videos_grid + +if is_wandb_available(): + import wandb + + +def get_random_downsample_ratio(sample_size, image_ratio=[], + all_choices=False, rng=None): + def _create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + first_element = 0.75 + remaining_sum = 1.0 - first_element + other_elements_value = remaining_sum / (length - 1) + special_list = [first_element] + [other_elements_value] * (length - 1) + return special_list + + if sample_size >= 1536: + number_list = [1, 1.25, 1.5, 2, 2.5, 3] + image_ratio + elif sample_size >= 1024: + number_list = [1, 1.25, 1.5, 2] + image_ratio + elif sample_size >= 768: + number_list = [1, 1.25, 1.5] + image_ratio + elif sample_size >= 512: + number_list = [1] + image_ratio + else: + number_list = [1] + + if all_choices: + return number_list + + number_list_prob = np.array(_create_special_list(len(number_list))) + if rng is None: + return np.random.choice(number_list, p = number_list_prob) + else: + return rng.choice(number_list, p = number_list_prob) + +def resize_mask(mask, latent, process_first_frame_only=True): + latent_size = latent.size() + batch_size, channels, num_frames, height, width = mask.shape + + if process_first_frame_only: + target_size = list(latent_size[2:]) + target_size[0] = 1 + first_frame_resized = F.interpolate( + mask[:, :, 0:1, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + + target_size = list(latent_size[2:]) + target_size[0] = target_size[0] - 1 + if target_size[0] != 0: + remaining_frames_resized = F.interpolate( + mask[:, :, 1:, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + resized_mask = torch.cat([first_frame_resized, remaining_frames_resized], dim=2) + else: + resized_mask = first_frame_resized + else: + target_size = list(latent_size[2:]) + resized_mask = F.interpolate( + mask, + size=target_size, + mode='trilinear', + align_corners=False + ) + return resized_mask + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): + try: + logger.info("Running validation... ") + + transformer3d_val = CogVideoXTransformer3DModel.from_pretrained_2d( + args.pretrained_model_name_or_path, subfolder="transformer" + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + if args.train_mode != "normal": + pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( + args.pretrained_model_name_or_path, + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + torch_dtype=weight_dtype + ) + else: + pipeline = CogVideoX_Fun_Pipeline.from_pretrained( + args.pretrained_model_name_or_path, + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + torch_dtype=weight_dtype + ) + pipeline = pipeline.to(accelerator.device) + + if args.seed is None: + generator = None + else: + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + + images = [] + for i in range(len(args.validation_prompts)): + with torch.no_grad(): + if args.train_mode != "normal": + with torch.autocast("cuda", dtype=weight_dtype): + video_length = int(args.video_sample_n_frames // vae.mini_batch_encoder * vae.mini_batch_encoder) if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + sample = pipeline( + args.validation_prompts[i], + video_length = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + guidance_scale = 6.0, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + + video_length = 1 + input_video, input_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + sample = pipeline( + args.validation_prompts[i], + video_length = 1, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + guidance_scale = 6.0, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + else: + with torch.autocast("cuda", dtype=weight_dtype): + sample = pipeline( + args.validation_prompts[i], + video_length = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + + sample = pipeline( + args.validation_prompts[i], + video_length = 1, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + + del pipeline + del transformer3d_val + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + + return images + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error with info {e}") + return None + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + current_value = initial_value + step_size * current_step + return current_value + +def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): + u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) + t = 1 / (1 + torch.exp(-u)) * (high - low) + low + return torch.clip(t.to(torch.int32), low, high - 1) + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--input_perturbation", type=float, default=0, help="The scale of input perturbation. Recommended 0.1." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help=( + "A folder containing the training data. " + ), + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help=( + "A csv containing the training data. " + ), + ) + parser.add_argument( + "--max_train_samples", + type=int, + default=None, + help=( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--multi_stream", + action="store_true", + help="whether to use cuda multi-stream", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--vae_mini_batch", type=int, default=32, help="mini batch size for vae." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.") + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument( + "--enable_xformers_memory_efficient_attention", action="store_true", help="Whether or not to use xformers." + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--snr_loss", action="store_true", help="Whether or not to use snr_loss." + ) + parser.add_argument( + "--not_sigma_loss", action="store_true", help="Whether or not to not use sigma_loss." + ) + parser.add_argument( + "--enable_text_encoder_in_dataloader", action="store_true", help="Whether or not to use text encoder in dataloader." + ) + parser.add_argument( + "--enable_bucket", action="store_true", help="Whether enable bucket sample in datasets." + ) + parser.add_argument( + "--random_ratio_crop", action="store_true", help="Whether enable random ratio crop sample in datasets." + ) + parser.add_argument( + "--random_frame_crop", action="store_true", help="Whether enable random frame crop sample in datasets." + ) + parser.add_argument( + "--random_hw_adapt", action="store_true", help="Whether enable random adapt height and width in datasets." + ) + parser.add_argument( + "--training_with_video_token_length", action="store_true", help="The training stage of the model in training.", + ) + parser.add_argument( + "--motion_sub_loss", action="store_true", help="Whether enable motion sub loss." + ) + parser.add_argument( + "--motion_sub_loss_ratio", type=float, default=0.25, help="The ratio of motion sub loss." + ) + parser.add_argument( + "--train_sampling_steps", + type=int, + default=1000, + help="Run train_sampling_steps.", + ) + parser.add_argument( + "--keep_all_node_same_token_length", + action="store_true", + help="Reference of the length token.", + ) + parser.add_argument( + "--token_sample_size", + type=int, + default=512, + help="Sample size of the token.", + ) + parser.add_argument( + "--video_sample_size", + type=int, + default=512, + help="Sample size of the video.", + ) + parser.add_argument( + "--image_sample_size", + type=int, + default=512, + help="Sample size of the video.", + ) + parser.add_argument( + "--video_sample_stride", + type=int, + default=4, + help="Sample stride of the video.", + ) + parser.add_argument( + "--video_sample_n_frames", + type=int, + default=17, + help="Num frame of video.", + ) + parser.add_argument( + "--video_repeat", + type=int, + default=0, + help="Num of repeat video.", + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + + parser.add_argument( + '--trainable_modules', + nargs='+', + help='Enter a list of trainable modules' + ) + parser.add_argument( + '--trainable_modules_low_learning_rate', + nargs='+', + default=[], + help='Enter a list of trainable modules with lower learning rate' + ) + parser.add_argument( + '--tokenizer_max_length', + type=int, + default=226, + help='Max length of tokenizer' + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--low_vram", action="store_true", help="Whether enable low_vram mode." + ) + parser.add_argument( + "--abnormal_norm_clip_start", + type=int, + default=1000, + help=( + 'When do we start doing additional processing on abnormal gradients. ' + ), + ) + parser.add_argument( + "--initial_grad_norm_ratio", + type=int, + default=5, + help=( + 'The initial gradient is relative to the multiple of the max_grad_norm. ' + ), + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + if accelerator.is_main_process: + writer = SummaryWriter(log_dir=logging_dir) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + rng = np.random.default_rng(np.random.PCG64(args.seed + accelerator.process_index)) + torch_rng = torch.Generator(accelerator.device).manual_seed(args.seed + accelerator.process_index) + else: + rng = None + torch_rng = None + index_rng = np.random.default_rng(np.random.PCG64(43)) + print(f"Init rng with seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + noise_scheduler = DDPMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + + tokenizer = T5Tokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer", revision=args.revision + ) + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So CLIPTextModel and AutoencoderKL will not enjoy the parameter sharding + # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + text_encoder = T5EncoderModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision, variant=args.variant, + torch_dtype=weight_dtype + ) + + vae = AutoencoderKLCogVideoX.from_pretrained( + args.pretrained_model_name_or_path, subfolder="vae", revision=args.revision, variant=args.variant + ) + + transformer3d = CogVideoXTransformer3DModel.from_pretrained_2d( + args.pretrained_model_name_or_path, subfolder="transformer" + ) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + # A good trainable modules is showed below now. + # For 3D Patch: trainable_modules = ['ff.net', 'pos_embed', 'attn2', 'proj_out', 'timepositionalencoding', 'h_position', 'w_position'] + # For 2D Patch: trainable_modules = ['ff.net', 'attn2', 'timepositionalencoding', 'h_position', 'w_position'] + transformer3d.train() + if accelerator.is_main_process: + accelerator.print( + f"Trainable modules '{args.trainable_modules}'." + ) + for name, param in transformer3d.named_parameters(): + for trainable_module_name in args.trainable_modules + args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + param.requires_grad = True + break + + # Create EMA for the transformer3d. + if args.use_ema: + ema_transformer3d = CogVideoXTransformer3DModel.from_pretrained_2d( + args.pretrained_model_name_or_path, subfolder="transformer" + ) + ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=CogVideoXTransformer3DModel, model_config=ema_transformer3d.config) + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + if args.use_ema: + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + weights.pop() + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + if args.use_ema: + ema_path = os.path.join(input_dir, "transformer_ema") + _, ema_kwargs = CogVideoXTransformer3DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = CogVideoXTransformer3DModel.from_pretrained_2d( + input_dir, subfolder="transformer_ema" + ) + load_model = EMAModel(load_model.parameters(), model_cls=CogVideoXTransformer3DModel, model_config=load_model.config) + load_model.load_state_dict(ema_kwargs) + + ema_transformer3d.load_state_dict(load_model.state_dict()) + ema_transformer3d.to(accelerator.device) + del load_model + + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() + + # load diffusers style into model + load_model = CogVideoXTransformer3DModel.from_pretrained_2d( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) + + model.load_state_dict(load_model.state_dict()) + del load_model + + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = [ + {'params': [], 'lr': args.learning_rate}, + {'params': [], 'lr': args.learning_rate / 2}, + ] + in_already = [] + for name, param in transformer3d.named_parameters(): + high_lr_flag = False + if name in in_already: + continue + for trainable_module_name in args.trainable_modules: + if trainable_module_name in name: + in_already.append(name) + high_lr_flag = True + trainable_params_optim[0]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate}") + break + if high_lr_flag: + continue + for trainable_module_name in args.trainable_modules_low_learning_rate: + if trainable_module_name in name: + in_already.append(name) + trainable_params_optim[1]['params'].append(param) + if accelerator.is_main_process: + print(f"Set {name} to lr : {args.learning_rate / 2}") + break + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + # weight_decay=args.adam_weight_decay, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # Get the training dataset + sample_n_frames_bucket_interval = 4 + + train_dataset = ImageVideoControlDataset( + args.train_data_meta, args.train_data_dir, + video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, + video_repeat=args.video_repeat, + image_sample_size=args.image_sample_size, + enable_bucket=args.enable_bucket, enable_inpaint=False, + ) + + if args.enable_bucket: + aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = AspectRatioBatchImageVideoSampler( + sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset, + batch_size=args.train_batch_size, train_folder = args.train_data_dir, drop_last=True, + aspect_ratios=aspect_ratio_sample_size, + ) + if args.keep_all_node_same_token_length: + if args.token_sample_size > 256: + numbers_list = list(range(256, args.token_sample_size + 1, 128)) + + if numbers_list[-1] != args.token_sample_size: + numbers_list.append(args.token_sample_size) + else: + numbers_list = [256] + numbers_list = [_number * _number * args.video_sample_n_frames for _number in numbers_list] + else: + numbers_list = None + + def get_length_to_frame_num(token_length): + if args.image_sample_size > args.video_sample_size: + sample_sizes = list(range(256, args.image_sample_size + 1, 128)) + + if sample_sizes[-1] != args.image_sample_size: + sample_sizes.append(args.image_sample_size) + else: + sample_sizes = [256] + + length_to_frame_num = { + sample_size: min(token_length / sample_size / sample_size, args.video_sample_n_frames) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 for sample_size in sample_sizes + } + + return length_to_frame_num + + def collate_fn(examples): + target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size + length_to_frame_num = get_length_to_frame_num( + target_token_length, + ) + + # Create new output + new_examples = {} + new_examples["target_token_length"] = target_token_length + new_examples["pixel_values"] = [] + new_examples["text"] = [] + new_examples["control_pixel_values"] = [] + + # Get ratio + pixel_value = examples[0]["pixel_values"] + data_type = examples[0]["data_type"] + f, h, w, c = np.shape(pixel_value) + if data_type == 'image': + random_downsample_ratio = 1 if not args.random_hw_adapt else get_random_downsample_ratio(args.image_sample_size, image_ratio=[args.image_sample_size / args.video_sample_size], rng=rng) + + aspect_ratio_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.image_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + else: + if args.random_hw_adapt: + if args.training_with_video_token_length: + local_min_size = np.min(np.array([np.mean(np.array([np.shape(example["pixel_values"])[1], np.shape(example["pixel_values"])[2]])) for example in examples])) + choice_list = [length for length in list(length_to_frame_num.keys()) if length < local_min_size * 1.25] + if len(choice_list) == 0: + choice_list = list(length_to_frame_num.keys()) + if rng is None: + local_video_sample_size = np.random.choice(choice_list) + else: + local_video_sample_size = rng.choice(choice_list) + batch_video_length = length_to_frame_num[local_video_sample_size] + random_downsample_ratio = args.video_sample_size / local_video_sample_size + else: + random_downsample_ratio = get_random_downsample_ratio( + args.video_sample_size, rng=rng) + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + else: + random_downsample_ratio = 1 + batch_video_length = args.video_sample_n_frames + sample_n_frames_bucket_interval + + aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} + aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} + + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] + if args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + + for example in examples: + if args.random_ratio_crop: + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() + control_pixel_values = control_pixel_values / 255. + + # Get adapt hw for resize + b, c, h, w = pixel_values.size() + th, tw = random_sample_size + if th / tw > h / w: + nh = int(th) + nw = int(w / h * nh) + else: + nw = int(tw) + nh = int(h / w * nw) + + transform = transforms.Compose([ + transforms.Resize([nh, nw]), + transforms.CenterCrop([int(x) for x in random_sample_size]), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + else: + closest_size = list(map(lambda x: int(x), closest_size)) + if closest_size[0] / h > closest_size[1] / w: + resize_size = closest_size[0], int(w * closest_size[0] / h) + else: + resize_size = int(h * closest_size[1] / w), closest_size[1] + + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() + control_pixel_values = control_pixel_values / 255. + + transform = transforms.Compose([ + transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(closest_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + new_examples["pixel_values"].append(transform(pixel_values)) + new_examples["control_pixel_values"].append(transform(control_pixel_values)) + new_examples["text"].append(example["text"]) + batch_video_length = int( + min( + batch_video_length, + (len(pixel_values) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1, + ) + ) + + if batch_video_length == 0: + batch_video_length = 1 + + new_examples["pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["pixel_values"]]) + new_examples["control_pixel_values"] = torch.stack([example[:batch_video_length] for example in new_examples["control_pixel_values"]]) + + if args.enable_text_encoder_in_dataloader: + prompt_ids = tokenizer( + new_examples['text'], + max_length=args.tokenizer_max_length, + padding="max_length", + add_special_tokens=True, + truncation=True, + return_tensors="pt" + ) + encoder_hidden_states = text_encoder( + prompt_ids.input_ids, + return_dict=False + )[0] + new_examples['encoder_attention_mask'] = prompt_ids.attention_mask + new_examples['encoder_hidden_states'] = encoder_hidden_states + + return new_examples + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=collate_fn, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + ) + else: + # DataLoaders creation: + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + persistent_workers=True if args.dataloader_num_workers != 0 else False, + num_workers=args.dataloader_num_workers, + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + + if args.use_ema: + ema_transformer3d.to(accelerator.device) + + # Move text_encode and vae to gpu and cast to weight_dtype + vae.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + tracker_config.pop("validation_prompts") + tracker_config.pop("trainable_modules") + tracker_config.pop("trainable_modules_low_learning_rate") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + if args.multi_stream: + # create extra cuda streams to speedup inpaint vae computation + vae_stream_1 = torch.cuda.Stream() + vae_stream_2 = torch.cuda.Stream() + else: + vae_stream_1 = None + vae_stream_2 = None + + for epoch in range(first_epoch, args.num_train_epochs): + train_loss = 0.0 + batch_sampler.sampler.generator = torch.Generator().manual_seed(args.seed + epoch) + for step, batch in enumerate(train_dataloader): + # Data batch sanity check + if epoch == first_epoch and step == 0: + pixel_values, texts = batch['pixel_values'].cpu(), batch['text'] + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + os.makedirs(os.path.join(args.output_dir, "sanity_check"), exist_ok=True) + for idx, (pixel_value, text) in enumerate(zip(pixel_values, texts)): + pixel_value = pixel_value[None, ...] + gif_name = '-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'{global_step}-{idx}' + save_videos_grid(pixel_value, f"{args.output_dir}/sanity_check/{gif_name[:10]}.gif", rescale=True) + + with accelerator.accumulate(transformer3d): + # Convert images to latent space + pixel_values = batch["pixel_values"].to(weight_dtype) + control_pixel_values = batch["control_pixel_values"].to(weight_dtype) + if args.training_with_video_token_length: + if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1)) + control_pixel_values = torch.tile(control_pixel_values, (4, 1, 1, 1, 1)) + if args.enable_text_encoder_in_dataloader: + batch['encoder_hidden_states'] = torch.tile(batch['encoder_hidden_states'], (4, 1, 1)) + batch['encoder_attention_mask'] = torch.tile(batch['encoder_attention_mask'], (4, 1)) + else: + batch['text'] = batch['text'] * 4 + elif args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 4 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: + pixel_values = torch.tile(pixel_values, (2, 1, 1, 1, 1)) + control_pixel_values = torch.tile(control_pixel_values, (2, 1, 1, 1, 1)) + if args.enable_text_encoder_in_dataloader: + batch['encoder_hidden_states'] = torch.tile(batch['encoder_hidden_states'], (2, 1, 1)) + batch['encoder_attention_mask'] = torch.tile(batch['encoder_attention_mask'], (2, 1)) + else: + batch['text'] = batch['text'] * 2 + + def create_special_list(length): + if length == 1: + return [1.0] + if length >= 2: + last_element = 0.90 + remaining_sum = 1.0 - last_element + other_elements_value = remaining_sum / (length - 1) + special_list = [other_elements_value] * (length - 1) + [last_element] + return special_list + + if args.keep_all_node_same_token_length: + actual_token_length = index_rng.choice(numbers_list) + + actual_video_length = (min( + actual_token_length / pixel_values.size()[-1] / pixel_values.size()[-2], args.video_sample_n_frames + ) - 1) // sample_n_frames_bucket_interval * sample_n_frames_bucket_interval + 1 + actual_video_length = int(max(actual_video_length, 1)) + else: + actual_video_length = None + + if args.random_frame_crop: + select_frames = [_tmp for _tmp in list(range(sample_n_frames_bucket_interval + 1, args.video_sample_n_frames + sample_n_frames_bucket_interval, sample_n_frames_bucket_interval))] + select_frames_prob = np.array(create_special_list(len(select_frames))) + + if rng is None: + temp_n_frames = np.random.choice(select_frames, p = select_frames_prob) + else: + temp_n_frames = rng.choice(select_frames, p = select_frames_prob) + if args.keep_all_node_same_token_length: + temp_n_frames = min(actual_video_length, temp_n_frames) + + pixel_values = pixel_values[:, :temp_n_frames, :, :] + control_pixel_values = control_pixel_values[:, :temp_n_frames, :, :] + + if args.low_vram: + torch.cuda.empty_cache() + vae.to(accelerator.device) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device) + + with torch.no_grad(): + # This way is quicker when batch grows up + def _slice_vae(pixel_values): + pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w") + bs = args.vae_mini_batch + new_pixel_values = [] + for i in range(0, pixel_values.shape[0], bs): + pixel_values_bs = pixel_values[i : i + bs] + pixel_values_bs = vae.encode(pixel_values_bs)[0] + pixel_values_bs = pixel_values_bs.sample() + new_pixel_values.append(pixel_values_bs) + vae._clear_fake_context_parallel_cache() + return torch.cat(new_pixel_values, dim = 0) + if vae_stream_1 is not None: + vae_stream_1.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(vae_stream_1): + latents = _slice_vae(pixel_values) + else: + latents = _slice_vae(pixel_values) + latents = latents * vae.config.scaling_factor + + control_latents = _slice_vae(control_pixel_values) + control_latents = control_latents * vae.config.scaling_factor + control_latents = rearrange(control_latents, "b c f h w -> b f c h w") + + latents = rearrange(latents, "b c f h w -> b f c h w") + + # wait for latents = vae.encode(pixel_values) to complete + if vae_stream_1 is not None: + torch.cuda.current_stream().wait_stream(vae_stream_1) + + if args.low_vram: + vae.to('cpu') + torch.cuda.empty_cache() + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device) + + if args.enable_text_encoder_in_dataloader: + prompt_embeds = batch['encoder_hidden_states'].to(device=latents.device) + else: + with torch.no_grad(): + prompt_ids = tokenizer( + batch['text'], + max_length=args.tokenizer_max_length, + padding="max_length", + add_special_tokens=True, + truncation=True, + return_tensors="pt" + ) + prompt_embeds = text_encoder( + prompt_ids.input_ids.to(latents.device), + return_dict=False + )[0] + + if args.low_vram and not args.enable_text_encoder_in_dataloader: + text_encoder.to('cpu') + torch.cuda.empty_cache() + + bsz = latents.shape[0] + noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype) + # Sample a random timestep for each image + # timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + timesteps = torch.randint(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng) + timesteps = timesteps.long() + + # Similar to diffusers.pipelines.hunyuandit.pipeline_hunyuandit.get_resize_crop_region_for_grid + def get_resize_crop_region_for_grid(src, tgt_width, tgt_height): + tw = tgt_width + th = tgt_height + h, w = src + r = h / w + if r > (th / tw): + resize_height = th + resize_width = int(round(th / h * w)) + else: + resize_width = tw + resize_height = int(round(tw / w * h)) + + crop_top = int(round((th - resize_height) / 2.0)) + crop_left = int(round((tw - resize_width) / 2.0)) + + return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width) + + def _prepare_rotary_positional_embeddings( + height: int, + width: int, + num_frames: int, + device: torch.device + ): + vae_scale_factor_spatial = ( + 2 ** (len(vae.config.block_out_channels) - 1) if vae is not None else 8 + ) + grid_height = height // (vae_scale_factor_spatial * unwrap_model(transformer3d).config.patch_size) + grid_width = width // (vae_scale_factor_spatial * unwrap_model(transformer3d).config.patch_size) + base_size_width = 720 // (vae_scale_factor_spatial * unwrap_model(transformer3d).config.patch_size) + base_size_height = 480 // (vae_scale_factor_spatial * unwrap_model(transformer3d).config.patch_size) + + grid_crops_coords = get_resize_crop_region_for_grid( + (grid_height, grid_width), base_size_width, base_size_height + ) + freqs_cos, freqs_sin = get_3d_rotary_pos_embed( + embed_dim=unwrap_model(transformer3d).config.attention_head_dim, + crops_coords=grid_crops_coords, + grid_size=(grid_height, grid_width), + temporal_size=num_frames, + use_real=True, + ) + freqs_cos = freqs_cos.to(device=device) + freqs_sin = freqs_sin.to(device=device) + return freqs_cos, freqs_sin + + height, width = batch["pixel_values"].size()[-2], batch["pixel_values"].size()[-1] + # 7. Create rotary embeds if required + image_rotary_emb = ( + _prepare_rotary_positional_embeddings(height, width, latents.size(1), latents.device) + if unwrap_model(transformer3d).config.use_rotary_positional_embeddings + else None + ) + prompt_embeds = prompt_embeds.to(device=latents.device) + + noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps) + if noise_scheduler.config.prediction_type == "epsilon": + target = noise + elif noise_scheduler.config.prediction_type == "v_prediction": + target = noise_scheduler.get_velocity(latents, noise, timesteps) + else: + raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}") + + # predict the noise residual + noise_pred = transformer3d( + hidden_states=noisy_latents, + encoder_hidden_states=prompt_embeds, + timestep=timesteps, + image_rotary_emb=image_rotary_emb, + return_dict=False, + control_latents=control_latents, + )[0] + print(noise_pred.size(), noisy_latents.size(), latents.size(), pixel_values.size()) + loss = F.mse_loss(noise_pred.float(), target.float(), reduction="mean") + + if args.motion_sub_loss and noise_pred.size()[1] > 2: + gt_sub_noise = noise_pred[:, 1:, :].float() - noise_pred[:, :-1, :].float() + pre_sub_noise = target[:, 1:, :].float() - target[:, :-1, :].float() + sub_loss = F.mse_loss(gt_sub_noise, pre_sub_noise, reduction="mean") + loss = loss * (1 - args.motion_sub_loss_ratio) + sub_loss * args.motion_sub_loss_ratio + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + if not args.use_deepspeed: + trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] + trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) + max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) + if trainable_params_total_norm / max_grad_norm > 5 and global_step > args.abnormal_norm_clip_start: + actual_max_grad_norm = max_grad_norm / min((trainable_params_total_norm / max_grad_norm), 10) + else: + actual_max_grad_norm = max_grad_norm + else: + actual_max_grad_norm = args.max_grad_norm + + if not args.use_deepspeed and args.report_model_info and accelerator.is_main_process: + if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: + for name, param in transformer3d.named_parameters(): + if param.requires_grad: + writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) + + norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) + if not args.use_deepspeed and args.report_model_info and accelerator.is_main_process: + writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) + writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + + if args.use_ema: + ema_transformer3d.step(transformer3d.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + if accelerator.is_main_process: + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + if accelerator.is_main_process: + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if accelerator.is_main_process: + transformer3d = unwrap_model(transformer3d) + if args.use_ema: + ema_transformer3d.copy_to(transformer3d.parameters()) + + if args.use_deepspeed or accelerator.is_main_process: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/train_control.sh b/scripts/train_control.sh new file mode 100644 index 0000000..1d31d07 --- /dev/null +++ b/scripts/train_control.sh @@ -0,0 +1,41 @@ +export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-Pose" +export DATASET_NAME="datasets/internal_datasets/" +export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json" +export NCCL_IB_DISABLE=1 +export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'". +accelerate launch --mixed_precision="bf16" scripts/train_control.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --train_data_dir=$DATASET_NAME \ + --train_data_meta=$DATASET_META_NAME \ + --image_sample_size=1024 \ + --video_sample_size=256 \ + --token_sample_size=512 \ + --video_sample_stride=3 \ + --video_sample_n_frames=49 \ + --train_batch_size=4 \ + --video_repeat=1 \ + --gradient_accumulation_steps=1 \ + --dataloader_num_workers=8 \ + --num_train_epochs=100 \ + --checkpointing_steps=50 \ + --learning_rate=2e-05 \ + --lr_scheduler="constant_with_warmup" \ + --lr_warmup_steps=50 \ + --seed=43 \ + --output_dir="output_dir" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --vae_mini_batch=1 \ + --max_grad_norm=0.05 \ + --random_hw_adapt \ + --training_with_video_token_length \ + --random_frame_crop \ + --enable_bucket \ + --use_came \ + --resume_from_checkpoint="latest" \ + --trainable_modules "." \ No newline at end of file diff --git a/scripts/train_lora.py b/scripts/train_lora.py index e23074e..ab733a7 100644 --- a/scripts/train_lora.py +++ b/scripts/train_lora.py @@ -79,7 +79,7 @@ from cogvideox.models.autoencoder_magvit import AutoencoderKLCogVideoX from cogvideox.models.transformer3d import CogVideoXTransformer3DModel from cogvideox.pipeline.pipeline_cogvideox import CogVideoX_Fun_Pipeline from cogvideox.pipeline.pipeline_cogvideox_inpaint import \ - CogVideoX_Fun_Pipeline_Inpaint + CogVideoX_Fun_Pipeline_Inpaint, add_noise_to_reference_video from cogvideox.utils.lora_utils import create_network, merge_lora, unmerge_lora from cogvideox.utils.utils import get_image_to_video_latent, save_videos_grid @@ -1387,6 +1387,8 @@ def main(): mask = 1 - mask mask = resize_mask(mask, latents) + if unwrap_model(transformer3d).config.add_noise_in_inpaint_model: + mask_pixel_values = add_noise_to_reference_video(mask_pixel_values) mask_pixel_values = rearrange(mask_pixel_values, "b f c h w -> b c f h w") bs = args.vae_mini_batch new_mask_pixel_values = []