MCPcopy Create free account
hub / github.com/JavisVerse/JavisGPT / load_data

Method load_data

interface/javisdit_interface.py:78–133  ·  view source on GitHub ↗
(self, meta_info: dict)

Source from the content-addressed store, hash-verified

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):

Callers 1

_get_itemMethod · 0.80

Calls 1

Tested by

no test coverage detected