MCPcopy Create free account
hub / github.com/FunAudioLLM/Fun-Audio-Chat / get_mm_inputs

Method get_mm_inputs

training/plugin/mm_plugins.py:209–276  ·  view source on GitHub ↗

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"],
    )

Source from the content-addressed store, hash-verified

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"

Callers

nothing calls this directly

Calls 1

_get_mm_inputsMethod · 0.95

Tested by

no test coverage detected