(self, meta_info: dict)
| 76 | raise NotImplementedError(f'Unknown version: {self.version}') |
| 77 | |
| 78 | def load_data(self, meta_info: dict): |
| 79 | target_text = meta_info['target_text'] |
| 80 | target_audio = meta_info.get('target_audio', None) |
| 81 | target_video = meta_info.get('target_video', None) |
| 82 | |
| 83 | if self.pre_tokenize: |
| 84 | if self.version == 'v0.1': |
| 85 | target_text = self.jav_tokenizer( |
| 86 | meta_info['target_text'], |
| 87 | max_length=self.jav_ctx_maxlen, |
| 88 | padding="max_length", |
| 89 | truncation=True, |
| 90 | return_attention_mask=True, |
| 91 | add_special_tokens=True, |
| 92 | return_tensors="pt", |
| 93 | ) # 'input_ids', 'attention_mask' |
| 94 | target_text['ib_text'] = imagebind_data.load_and_transform_text( |
| 95 | [meta_info['target_text']], 'cpu' |
| 96 | ) |
| 97 | elif self.version == 'v1.0': |
| 98 | ids, mask = self.jav_tokenizer( |
| 99 | meta_info['target_text'], |
| 100 | return_mask=True, |
| 101 | add_special_tokens=True |
| 102 | ) |
| 103 | target_text = {'ids': ids, 'mask': mask} |
| 104 | else: |
| 105 | raise ValueError(f'Unknown version: {self.version}') |
| 106 | |
| 107 | task_type = meta_info.get('task_type', 'T2AV') |
| 108 | fix_start_frame = None |
| 109 | if task_type != 'T2AV': |
| 110 | fix_start_frame = 0 |
| 111 | if 'Exten' in task_type: |
| 112 | fix_start_frame = int(self.video_fps * meta_info['start_s']) |
| 113 | |
| 114 | if target_video is not None: |
| 115 | if isinstance(target_video, torch.Tensor): |
| 116 | pass |
| 117 | elif isinstance(target_video, str): |
| 118 | target_video = f'{self.video_folder}/{target_video}' |
| 119 | target_audio = f'{self.audio_folder}/{target_audio}' |
| 120 | assert os.path.exists(target_video) |
| 121 | if not os.path.exists(target_audio): |
| 122 | target_audio = target_video |
| 123 | if self.load_av_feat: |
| 124 | target_audio, target_video = torch.load(target_audio), torch.load(target_video) |
| 125 | else: |
| 126 | target_audio, target_video = self.load_transform_audio_video( |
| 127 | target_audio, target_video, fix_start_frame=fix_start_frame |
| 128 | ) |
| 129 | else: |
| 130 | raise ValueError(f"Unsupported target_video type: {type(target_video)}") |
| 131 | |
| 132 | # TODO: AV-Temporal Masks for X-Conditional Generation |
| 133 | return target_text, target_audio, target_video |
| 134 | |
| 135 | def load_transform_audio_video(self, audio_path, video_path, fix_start_frame=None): |
no test coverage detected