From 67ef062b64fb50070edcb096ad8f4d7e7e261563 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Mon, 30 Sep 2024 16:45:48 +0800 Subject: [PATCH] update bug in control training (#26) --- cogvideox/data/dataset_image_video.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/cogvideox/data/dataset_image_video.py b/cogvideox/data/dataset_image_video.py index 032ae9e..f714d69 100644 --- a/cogvideox/data/dataset_image_video.py +++ b/cogvideox/data/dataset_image_video.py @@ -393,7 +393,7 @@ class ImageVideoControlDataset(Dataset): 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'] + video_id, text = data_info['file_path'], data_info['text'] if data_info.get('type', 'image')=='video': if self.data_root is None: @@ -444,6 +444,8 @@ class ImageVideoControlDataset(Dataset): if random.random() < self.text_drop_ratio: text = '' + control_video_id = data_info['control_file_path'] + if self.data_root is None: control_video_id = control_video_id else: @@ -489,6 +491,8 @@ class ImageVideoControlDataset(Dataset): if random.random() < self.text_drop_ratio: text = '' + control_image_id = data_info['control_file_path'] + if self.data_root is None: control_image_id = control_image_id else: