Get multimodal inputs for training.
(
self,
images: list["ImageInput"],
videos: list["VideoInput"],
audios: list["AudioInput"],
imglens: list[int],
vidlens: list[int],
audlens: list[int],
batch_ids: list[list[int]],
processor: Optional["MMProcessor"],
)
| 207 | |
| 208 | @override |
| 209 | def get_mm_inputs( |
| 210 | self, |
| 211 | images: list["ImageInput"], |
| 212 | videos: list["VideoInput"], |
| 213 | audios: list["AudioInput"], |
| 214 | imglens: list[int], |
| 215 | vidlens: list[int], |
| 216 | audlens: list[int], |
| 217 | batch_ids: list[list[int]], |
| 218 | processor: Optional["MMProcessor"], |
| 219 | ) -> dict[str, Union[list[int], "torch.Tensor"]]: |
| 220 | """Get multimodal inputs for training.""" |
| 221 | mm_inputs = {} |
| 222 | |
| 223 | if audios is not None: |
| 224 | # Parse JSON data |
| 225 | parsed_audios = [json.loads(audio) for audio in audios] |
| 226 | |
| 227 | # Filter assistant wav based on <|audio_bos|> |
| 228 | for i in range(len(parsed_audios)): |
| 229 | if parsed_audios[i]['token'].startswith('<|audio_bos|>') and parsed_audios[i]['token'].count('AU') > 0: |
| 230 | if 'wav_path' in parsed_audios[i]: |
| 231 | parsed_audios[i]['wav_path'] = '' |
| 232 | if 'path' in parsed_audios[i]: |
| 233 | parsed_audios[i]['path'] = '' |
| 234 | |
| 235 | audio_paths = [ |
| 236 | _audio_data.get('path', _audio_data.get('wav_path', '')) |
| 237 | for _audio_data in parsed_audios |
| 238 | if _audio_data.get('path', _audio_data.get('wav_path', '')) != '' |
| 239 | ] |
| 240 | audio_tokens = [_audio_data['token'] for _audio_data in parsed_audios] |
| 241 | audio_texts = [_audio_data['text'] for _audio_data in parsed_audios] |
| 242 | |
| 243 | audio_inputs = self._get_mm_inputs(images, videos, audio_paths, processor) |
| 244 | audio_inputs['feature_exist_mask'] = torch.tensor( |
| 245 | [_audio_data.get('path', _audio_data.get('wav_path', '')) != '' for _audio_data in parsed_audios], |
| 246 | dtype=torch.bool |
| 247 | ) |
| 248 | |
| 249 | if audios is not None and len(audios) != 0: |
| 250 | _audio_inputs = processor.speech_tokenizer( |
| 251 | audio_tokens, |
| 252 | return_attention_mask=True, |
| 253 | return_token_type_ids=False, |
| 254 | padding=True, |
| 255 | return_tensors="pt" |
| 256 | ) |
| 257 | audio_inputs["speech_ids"] = _audio_inputs.pop("input_ids") |
| 258 | audio_inputs["speech_attention_mask"] = _audio_inputs.pop("attention_mask") |
| 259 | |
| 260 | if audio_texts is not None and len(audio_texts) != 0: |
| 261 | audio_inputs['text_ids'] = processor.tokenizer( |
| 262 | audio_texts, |
| 263 | return_attention_mask=False, |
| 264 | return_token_type_ids=False, |
| 265 | padding=True, |
| 266 | return_tensors="pt" |
nothing calls this directly
no test coverage detected