1. 这里不管怎样,都需要一个完整的 assistant generation ids; 2. 支持从任意位置进行截断;
(self, output: List[Tuple[int, torch.Tensor]])
| 729 | return decoded_audio_list |
| 730 | |
| 731 | def decode(self, output: List[Tuple[int, torch.Tensor]]): |
| 732 | """ |
| 733 | 1. 这里不管怎样,都需要一个完整的 assistant generation ids; |
| 734 | 2. 支持从任意位置进行截断; |
| 735 | """ |
| 736 | |
| 737 | genearted_messages = [] |
| 738 | for start_length, generation_ids in output: |
| 739 | content = self._parse_text_codes(start_length, generation_ids[:, 0]) |
| 740 | audio_codes_list = self._parse_audio_codes( |
| 741 | start_length, generation_ids[:, 1:] |
| 742 | ) |
| 743 | if content == "": |
| 744 | message = None |
| 745 | else: |
| 746 | message = AssistantMessage( |
| 747 | content=content, |
| 748 | audio_codes_list=cast( |
| 749 | List[Union[str, torch.Tensor]], audio_codes_list |
| 750 | ), |
| 751 | ) |
| 752 | genearted_messages.append(message) |
| 753 | return genearted_messages |
| 754 | |
| 755 | @staticmethod |
| 756 | def loudness_normalize( |
no test coverage detected