Parse the generated output into text content and audio codes.
(
tokenizer: Tokenizer,
generation_ids: np.ndarray,
prompt_len: int,
)
| 221 | |
| 222 | |
| 223 | def parse_generation_output( |
| 224 | tokenizer: Tokenizer, |
| 225 | generation_ids: np.ndarray, |
| 226 | prompt_len: int, |
| 227 | ) -> tuple[str, np.ndarray]: |
| 228 | """Parse the generated output into text content and audio codes.""" |
| 229 | gen_part = generation_ids[prompt_len:] |
| 230 | text_channel = gen_part[:, 0].tolist() |
| 231 | audio_channels = gen_part[:, 1:] |
| 232 | |
| 233 | audio_start_tok = _get_special_token_str(tokenizer, AUDIO_START_TOKEN_ID) |
| 234 | gen_slot_tok = _get_special_token_str(tokenizer, AUDIO_ASSISTANT_GEN_SLOT_TOKEN_ID) |
| 235 | delay_slot_tok = _get_special_token_str(tokenizer, AUDIO_ASSISTANT_DELAY_SLOT_TOKEN_ID) |
| 236 | audio_end_tok = _get_special_token_str(tokenizer, AUDIO_END_TOKEN_ID) |
| 237 | |
| 238 | raw_text = tokenizer.decode(text_channel) |
| 239 | |
| 240 | pattern = re.compile( |
| 241 | rf"(?:{re.escape(audio_start_tok)})?" |
| 242 | rf"(?:{re.escape(gen_slot_tok)})*" |
| 243 | rf"(?:{re.escape(delay_slot_tok)})*" |
| 244 | rf"{re.escape(audio_end_tok)}" |
| 245 | ) |
| 246 | |
| 247 | def repl(match: re.Match) -> str: |
| 248 | seg = match.group(0) |
| 249 | if gen_slot_tok in seg: |
| 250 | return AUDIO_PLACEHOLDER |
| 251 | return "" |
| 252 | |
| 253 | text = pattern.sub(repl, raw_text) |
| 254 | |
| 255 | segments = extract_audio_segments(audio_channels) |
| 256 | |
| 257 | if segments: |
| 258 | audio_codes = np.concatenate(segments, axis=0) |
| 259 | else: |
| 260 | audio_codes = np.zeros((0, N_VQ), dtype=np.int64) |
| 261 | |
| 262 | return text, audio_codes |
no test coverage detected