(
self,
data_in,
data_lengths=None,
key: list = None,
tokenizer=None,
frontend=None,
**kwargs,
)
| 459 | return output |
| 460 | |
| 461 | def inference_prepare( |
| 462 | self, |
| 463 | data_in, |
| 464 | data_lengths=None, |
| 465 | key: list = None, |
| 466 | tokenizer=None, |
| 467 | frontend=None, |
| 468 | **kwargs, |
| 469 | ): |
| 470 | meta_data = {} |
| 471 | |
| 472 | if kwargs.get("batch_size", 1) > 1: |
| 473 | raise NotImplementedError("batch decoding is not implemented") |
| 474 | |
| 475 | contents = self.data_template(data_in[0]) |
| 476 | output = self.data_load_speech(contents, tokenizer, frontend, meta_data=meta_data, **kwargs) |
| 477 | batch = to_device(output, kwargs["device"]) |
| 478 | |
| 479 | # audio encoder |
| 480 | speech = batch["speech"] |
| 481 | |
| 482 | if len(speech) > 0: |
| 483 | if "audio_embedding" in kwargs and "audio_embedding_lens" in kwargs: |
| 484 | encoder_out = kwargs["audio_embedding"] |
| 485 | encoder_out_lens = kwargs["audio_embedding_lens"] |
| 486 | else: |
| 487 | speech_lengths = batch["speech_lengths"][:, 0] |
| 488 | # fp16 |
| 489 | if kwargs.get("fp16", False): |
| 490 | speech = speech.to(torch.float16) |
| 491 | elif kwargs.get("bf16", False): |
| 492 | speech = speech.to(torch.bfloat16) |
| 493 | # audio encoder |
| 494 | encoder_out, encoder_out_lens = self.encode(speech, speech_lengths) |
| 495 | |
| 496 | # audio_adaptor |
| 497 | adaptor_out, adaptor_out_lens = self.audio_adaptor(encoder_out, encoder_out_lens) |
| 498 | meta_data["encoder_out"] = encoder_out |
| 499 | meta_data["encoder_out_lens"] = encoder_out_lens |
| 500 | meta_data["audio_adaptor_out"] = adaptor_out |
| 501 | meta_data["audio_adaptor_out_lens"] = adaptor_out_lens |
| 502 | |
| 503 | input_ids = batch["input_ids"] |
| 504 | source_ids = batch["source_ids"] |
| 505 | fbank_beg = batch["fbank_beg"] |
| 506 | fake_token_len = batch["fake_token_len"] |
| 507 | |
| 508 | if not kwargs.get("teacherforcing", False): |
| 509 | input_ids = source_ids |
| 510 | |
| 511 | input_ids[input_ids < 0] = 0 |
| 512 | inputs_embeds = self.llm.model.get_input_embeddings()(input_ids) |
| 513 | |
| 514 | batch_size, token_num, dims = inputs_embeds.shape |
| 515 | |
| 516 | fake_token_len[fake_token_len < 0] = 0 |
| 517 | fbank_beg[fbank_beg < 0] = 0 |
| 518 |
no test coverage detected