processing_moss_tts.py
34.7 KB · 950 lines · python Raw
1 # coding=utf-8
2 # Copyright 2026 OpenMOSS and the HuggingFace Inc. team. All rights reserved.
3 #
4 # Licensed under the Apache License, Version 2.0 (the "License");
5 # you may not use this file except in compliance with the License.
6 # You may obtain a copy of the License at
7 #
8 # http://www.apache.org/licenses/LICENSE-2.0
9 #
10 # Unless required by applicable law or agreed to in writing, software
11 # distributed under the License is distributed on an "AS IS" BASIS,
12 # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 # See the License for the specific language governing permissions and
14 # limitations under the License.
15
16 import os
17 from typing import Any, Dict, List, Optional, Tuple, Type, Union, Literal, Final, cast
18 from dataclasses import dataclass
19 from pathlib import Path
20 import re
21 import torchaudio
22
23 from transformers import processing_utils
24
25 processing_utils.MODALITY_TO_BASE_CLASS_MAPPING["audio_tokenizer"] = "PreTrainedModel"
26
27 import torch
28 from transformers import (
29 PreTrainedTokenizerBase,
30 BatchFeature,
31 ProcessorMixin,
32 logging,
33 AutoConfig,
34 AutoModel,
35 AutoTokenizer,
36 )
37
38 from .configuration_moss_tts import MossTTSDelayConfig
39 from .tts_robust_normalizer_single_script import normalize_tts_text
40
41
42 logger = logging.get_logger(__name__)
43
44
45 AUDIO_PLACEHOLDER = "<|audio|>"
46
47
48 @dataclass
49 class Message:
50 def to_dict(self) -> Dict[str, Any]:
51 raise NotImplementedError
52
53
54 @dataclass
55 class UserMessage(Message):
56 text: Optional[str] = None
57 reference: Optional[List[Optional[Union[str, torch.Tensor]]]] = None
58 instruction: Optional[str] = None
59 tokens: Optional[int] = None
60 quality: Optional[str] = None
61 sound_event: Optional[str] = None
62 ambient_sound: Optional[str] = None
63 language: Optional[str] = None
64
65 def __post_init__(self):
66 template = """<user_inst>
67 - Reference(s):
68 {reference}
69 - Instruction:
70 {instruction}
71 - Tokens:
72 {tokens}
73 - Quality:
74 {quality}
75 - Sound Event:
76 {sound_event}
77 - Ambient Sound:
78 {ambient_sound}
79 - Language:
80 {language}
81 - Text:
82 {text}
83 </user_inst>"""
84
85 audio_codes_list = []
86 if self.reference is None:
87 reference = "None"
88 elif isinstance(self.reference, List):
89 reference = []
90 for speaker_idx, speaker_reference in enumerate(self.reference):
91 if speaker_reference is not None:
92 reference.append(f"[S{speaker_idx+1}]:\n{AUDIO_PLACEHOLDER}")
93 reference = "\n".join(reference)
94 audio_codes_list = [
95 speaker_reference
96 for speaker_reference in self.reference
97 if speaker_reference is not None
98 ]
99 else:
100 raise TypeError("`reference` should be exactly a list when it is not None.")
101
102 content = (
103 template.replace("{reference}", str(reference))
104 .replace("{instruction}", str(self.instruction))
105 .replace("{tokens}", str(self.tokens))
106 .replace("{quality}", str(self.quality))
107 .replace("{sound_event}", str(self.sound_event))
108 .replace("{ambient_sound}", str(self.ambient_sound))
109 .replace("{language}", str(self.language))
110 .replace("{text}", str(self.text))
111 )
112
113 self._content = content
114 self._audio_codes_list = audio_codes_list
115
116 def to_dict(self):
117 return {
118 "role": "user",
119 "content": self._content,
120 "audio_codes_list": self._audio_codes_list,
121 }
122
123
124 @dataclass
125 class AssistantMessage(Message):
126 audio_codes_list: List[Union[str, torch.Tensor]]
127 content: str = AUDIO_PLACEHOLDER
128
129 def to_dict(self):
130 return {
131 "role": "assistant",
132 "content": self.content,
133 "audio_codes_list": self.audio_codes_list,
134 }
135
136
137 USER_MESSAGE_FIELDS = (
138 "text",
139 "reference",
140 "instruction",
141 "tokens",
142 "quality",
143 "sound_event",
144 "ambient_sound",
145 "language",
146 )
147
148
149 class MossTTSDelayProcessor(ProcessorMixin):
150 tokenizer_class = "AutoTokenizer"
151 audio_tokenizer_class = "AutoModel"
152
153 tokenizer: PreTrainedTokenizerBase
154 audio_tokenizer: Any
155
156 def __init__(
157 self,
158 tokenizer: PreTrainedTokenizerBase,
159 audio_tokenizer: Any = None,
160 model_config: Optional[MossTTSDelayConfig] = None,
161 **kwargs,
162 ):
163 super().__init__(tokenizer=tokenizer, audio_tokenizer=audio_tokenizer, **kwargs)
164
165 # Explicit assignments for type-checkers; ProcessorMixin sets these too.
166 self.tokenizer = tokenizer
167 self.audio_tokenizer = audio_tokenizer
168 if model_config is None:
169 model_config = MossTTSDelayConfig()
170 self.model_config = model_config
171
172 self.imstart_token_id = tokenizer.convert_tokens_to_ids("<|im_start|>")
173 self.imend_token_id = tokenizer.convert_tokens_to_ids("<|im_end|>")
174 self.newline_token_id = 198
175
176 def _id_to_token(token_id: int) -> str:
177 tok = tokenizer.convert_ids_to_tokens(int(token_id))
178 if isinstance(tok, list):
179 return tok[0] if len(tok) > 0 else ""
180 return cast(str, tok)
181
182 self.audio_user_slot_token = _id_to_token(
183 self.model_config.audio_user_slot_token_id
184 )
185 self.audio_assistant_gen_slot_token = _id_to_token(
186 self.model_config.audio_assistant_gen_slot_token_id
187 )
188 self.audio_assistant_delay_slot_token = _id_to_token(
189 self.model_config.audio_assistant_delay_slot_token_id
190 )
191 self.audio_start_token = _id_to_token(self.model_config.audio_start_token_id)
192 self.audio_end_token = _id_to_token(self.model_config.audio_end_token_id)
193
194 @classmethod
195 def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):
196 trust_remote_code = kwargs.pop("trust_remote_code", True)
197 kwargs.pop("_from_auto", None)
198
199 audio_tokenizer_name_or_path = kwargs.pop("codec_path", None)
200 if audio_tokenizer_name_or_path is None:
201 processor_lookup_kwargs = dict(kwargs)
202 try:
203 processor_dict, _ = cls.get_processor_dict(
204 pretrained_model_name_or_path,
205 **processor_lookup_kwargs,
206 )
207 audio_tokenizer_name_or_path = processor_dict.get(
208 "audio_tokenizer_name_or_path"
209 )
210 audio_tokenizer_dict = processor_dict.get("audio_tokenizer", {})
211 if isinstance(audio_tokenizer_dict, dict):
212 audio_tokenizer_name_or_path = audio_tokenizer_dict.get(
213 "audio_tokenizer_name_or_path"
214 ) or audio_tokenizer_name_or_path
215 except Exception:
216 audio_tokenizer_name_or_path = None
217 if audio_tokenizer_name_or_path is None:
218 audio_tokenizer_name_or_path = "OpenMOSS-Team/MOSS-Audio-Tokenizer"
219
220 pretrained_model_name_or_path = Path(pretrained_model_name_or_path)
221 model_config = cast(
222 MossTTSDelayConfig,
223 AutoConfig.from_pretrained(
224 pretrained_model_name_or_path,
225 *args,
226 trust_remote_code=trust_remote_code,
227 **kwargs,
228 ),
229 )
230 tokenizer = AutoTokenizer.from_pretrained(
231 pretrained_model_name_or_path,
232 *args,
233 trust_remote_code=trust_remote_code,
234 **kwargs,
235 )
236 audio_tokenizer = AutoModel.from_pretrained(
237 audio_tokenizer_name_or_path,
238 trust_remote_code=trust_remote_code,
239 **kwargs,
240 )
241
242 return cls(
243 tokenizer=tokenizer,
244 audio_tokenizer=audio_tokenizer,
245 model_config=model_config,
246 **kwargs,
247 )
248
249 def __call__(self, *args, **kwargs) -> BatchFeature:
250 conversations = args[0] if len(args) > 0 else kwargs.pop("conversations")
251 mode: str = kwargs.pop("mode", "generation")
252 apply_chat_template: bool = kwargs.pop("apply_chat_template", True)
253 n_vq: Optional[int] = kwargs.pop("n_vq", None)
254
255 # Common ProcessorMixin kwargs that we ignore because we always return torch tensors.
256 kwargs.pop("return_tensors", None)
257 kwargs.pop("padding", None)
258 kwargs.pop("truncation", None)
259
260 """
261 mode only works when a Message is converted to a dict.
262 """
263
264 if mode not in {"generation", "continuation", "computing_loss"}:
265 raise RuntimeError
266
267 if isinstance(conversations, (Message, Dict)):
268 conversations = [conversations]
269
270 truncation = False
271 if mode == "continuation":
272 truncation = True
273
274 input_ids_list = []
275 for conversation in conversations:
276 if isinstance(conversation, (Message, Dict)):
277 conversation = [conversation]
278
279 # Normalize early so downstream logic always deals with dict messages.
280 conversation = [self._normalize_message(m) for m in conversation]
281
282 if (mode == "generation") ^ (len(conversation) % 2 != 0):
283 raise ValueError
284
285 if (mode == "generation") ^ (conversation[-1]["role"] == "user"):
286 raise ValueError
287
288 unified_codes = []
289 for message_idx, message in enumerate(conversation):
290 if apply_chat_template:
291 add_generation_prompt = (
292 mode == "generation" and message_idx == len(conversation) - 1
293 )
294 try:
295 content = self.tokenizer.apply_chat_template(
296 [{"role": message["role"], "content": message["content"]}],
297 add_generation_prompt=add_generation_prompt,
298 tokenize=False,
299 )
300 except TypeError:
301 try:
302 content = self.tokenizer.apply_chat_template(
303 [
304 {
305 "role": message["role"],
306 "content": message["content"],
307 }
308 ],
309 add_generation_prompt=add_generation_prompt,
310 )
311 except Exception:
312 logger.warning(
313 "apply_chat_template failed; fallback to raw content."
314 )
315 content = message["content"]
316 else:
317 content = message["content"]
318
319 if not isinstance(content, str):
320 content = str(content)
321
322 # Batch-encode all path-based references in one call when possible.
323 # This ensures we actually exercise audio_tokenizer.batch_encode for multi-reference prompts,
324 # instead of repeatedly calling it with batch=1.
325 raw_audio_items = message.get("audio_codes_list", [])
326
327 audio_codes_list: List[torch.Tensor] = []
328 if len(raw_audio_items) > 0:
329 encoded_items: List[Optional[torch.Tensor]] = [None] * len(
330 raw_audio_items
331 )
332 paths: List[str] = []
333 path_positions: List[int] = []
334
335 for idx, item in enumerate(raw_audio_items):
336 if isinstance(item, torch.Tensor):
337 if n_vq is not None and item.shape[1] != n_vq:
338 raise RuntimeError(
339 "audio_codes's n_vq is not equal to the parameter `n_vq`. Your can set the parameter `n_vq` as None if you have already tokenzied the wavs."
340 )
341 encoded_items[idx] = item
342 continue
343
344 if isinstance(item, (str, os.PathLike)):
345 paths.append(str(item))
346 path_positions.append(idx)
347 continue
348
349 raise TypeError(
350 "Each audio item must be a torch.Tensor of codes or a path-like string."
351 )
352
353 if len(paths) > 0:
354 encoded_from_paths = self.encode_audios_from_path(paths, n_vq)
355 if len(encoded_from_paths) != len(paths):
356 raise RuntimeError(
357 "encode_audios_from_path returned an unexpected number of items."
358 )
359 for pos, codes in zip(path_positions, encoded_from_paths):
360 encoded_items[pos] = codes
361
362 audio_codes_list = [cast(torch.Tensor, t) for t in encoded_items]
363 unified_codes.append(
364 self._get_unified_codes(
365 message["role"], content, audio_codes_list, truncation
366 )
367 )
368
369 unified_codes = torch.cat(unified_codes)
370 input_ids_list.append(unified_codes)
371
372 return BatchFeature(data=self._pad(input_ids_list))
373
374 @staticmethod
375 def build_user_message(
376 text: Optional[str] = None,
377 reference: Optional[List[Optional[Union[str, torch.Tensor]]]] = None,
378 instruction: Optional[str] = None,
379 tokens: Optional[int] = None,
380 quality: Optional[str] = None,
381 sound_event: Optional[str] = None,
382 ambient_sound: Optional[str] = None,
383 language: Optional[str] = None,
384 ) -> Dict:
385 if reference is not None and not isinstance(reference, list):
386 reference = [reference]
387 text = normalize_tts_text(text)
388 return UserMessage(
389 text=text,
390 reference=reference,
391 instruction=instruction,
392 tokens=tokens,
393 quality=quality,
394 sound_event=sound_event,
395 ambient_sound=ambient_sound,
396 language=language,
397 ).to_dict()
398
399 @staticmethod
400 def build_assistant_message(
401 audio_codes_list: List[Union[str, torch.Tensor]],
402 content: str = AUDIO_PLACEHOLDER,
403 ) -> Dict:
404 return AssistantMessage(
405 audio_codes_list=audio_codes_list,
406 content=content,
407 ).to_dict()
408
409 def _normalize_message(self, message: Union[Message, Dict]) -> Dict:
410 if isinstance(message, Message):
411 return message.to_dict()
412 if not isinstance(message, dict):
413 raise TypeError("Each message must be a Message or dict.")
414 if "role" not in message:
415 raise ValueError("Message dict must include a 'role' field.")
416 if "content" in message and "audio_codes_list" in message:
417 return message
418 role = message["role"]
419 if role == "user":
420 kwargs = {key: message.get(key) for key in USER_MESSAGE_FIELDS}
421 return self.build_user_message(**kwargs)
422 if role == "assistant":
423 return self.build_assistant_message(
424 audio_codes_list=message.get("audio_codes_list", []),
425 content=message.get("content", AUDIO_PLACEHOLDER),
426 )
427 raise ValueError(f"Unsupported role: {role}")
428
429 def _pad(self, input_ids_list: List[torch.Tensor]):
430 device = input_ids_list[0].device
431 lengths = torch.tensor([w.shape[0] for w in input_ids_list], device=device)
432 pad_input_ids = torch.nn.utils.rnn.pad_sequence(
433 input_ids_list,
434 batch_first=True,
435 padding_value=self.model_config.audio_pad_code,
436 padding_side="left",
437 )
438 other_channel_mask = (pad_input_ids.shape[1] - lengths).unsqueeze(
439 1
440 ) > torch.arange(pad_input_ids.shape[1], device=device).unsqueeze(0)
441 pad_input_ids[..., 0][other_channel_mask] = self.model_config.pad_token_id
442 attention_mask = torch.zeros(
443 pad_input_ids.shape[0], pad_input_ids.shape[1], device=device
444 )
445 attention_mask[~other_channel_mask] = 1
446 attention_mask = attention_mask.bool()
447 return {
448 "input_ids": pad_input_ids, # [batch_size, seqlen, n_vq]
449 "attention_mask": attention_mask,
450 }
451
452 @staticmethod
453 def _replace_audio_placeholders(
454 content: str,
455 lengths: List[int],
456 n_vq: int,
457 gen_slot_token: str,
458 delay_slot_token: str,
459 audio_start_token: str,
460 audio_end_token: str,
461 ) -> str:
462 if n_vq < 1:
463 raise ValueError(f"n_vq must be >= 1, got {n_vq}")
464
465 num_placeholders = content.count(AUDIO_PLACEHOLDER)
466 if num_placeholders != len(lengths):
467 raise ValueError(
468 f"Number of {AUDIO_PLACEHOLDER} ({num_placeholders}) "
469 f"does not match lengths ({len(lengths)})"
470 )
471
472 def build_audio_block(length: int) -> str:
473 if length < 0:
474 raise ValueError(f"length must be >= 0, got {length}")
475
476 if length == 0:
477 return f"{audio_start_token}{audio_end_token}"
478
479 step_tokens = gen_slot_token * length + (delay_slot_token * (n_vq - 1))
480 return f"{audio_start_token}{step_tokens}{audio_end_token}"
481
482 lengths_iter = iter(lengths)
483
484 def replacer(match: re.Match) -> str:
485 length = next(lengths_iter)
486 return build_audio_block(length)
487
488 result = re.sub(re.escape(AUDIO_PLACEHOLDER), replacer, content)
489
490 return result
491
492 @staticmethod
493 def _merge_consecutive_audio_placeholders(
494 content: str,
495 audio_codes_list: List[torch.Tensor],
496 ) -> Tuple[str, List[torch.Tensor]]:
497 matches = list(re.finditer(re.escape(AUDIO_PLACEHOLDER), content))
498 if len(matches) <= 1:
499 return content, audio_codes_list
500
501 if len(matches) != len(audio_codes_list):
502 raise ValueError(
503 "Audio placeholders do not match the provided audio codes list."
504 )
505
506 new_audio_codes_list = []
507 new_parts = []
508 last_pos = 0
509 i = 0
510 while i < len(matches):
511 j = i
512 while (
513 j + 1 < len(matches)
514 and content[matches[j].end() : matches[j + 1].start()].strip() == ""
515 ):
516 j += 1
517
518 new_parts.append(content[last_pos : matches[i].start()])
519 new_parts.append(AUDIO_PLACEHOLDER)
520 last_pos = matches[j].end()
521
522 if j == i:
523 new_audio_codes_list.append(audio_codes_list[i])
524 else:
525 new_audio_codes_list.append(
526 torch.cat(audio_codes_list[i : j + 1], dim=0)
527 )
528
529 i = j + 1
530
531 new_parts.append(content[last_pos:])
532 return "".join(new_parts), new_audio_codes_list
533
534 @staticmethod
535 def apply_delay_pattern(codes: torch.Tensor, pad_code: int) -> torch.Tensor:
536 delayed_tokens = torch.full(
537 (codes.shape[0] + codes.shape[1] - 1, codes.shape[1]),
538 pad_code,
539 device=codes.device,
540 dtype=codes.dtype,
541 )
542 for i in range(codes.shape[1]):
543 delayed_tokens[i : i + codes.shape[0], i] = codes[:, i]
544 return delayed_tokens
545
546 @staticmethod
547 def apply_de_delay_pattern(delay_codes: torch.Tensor) -> torch.Tensor:
548 tokens = torch.full(
549 (delay_codes.shape[0] - delay_codes.shape[1] + 1, delay_codes.shape[1]),
550 0,
551 device=delay_codes.device,
552 dtype=delay_codes.dtype,
553 )
554 for i in range(delay_codes.shape[1]):
555 tokens[:, i] = delay_codes[i : i + tokens.shape[0], i]
556 return tokens
557
558 def _get_unified_codes(
559 self,
560 role: str,
561 content: str,
562 audio_codes_list: List[torch.Tensor],
563 truncation: bool,
564 ) -> torch.Tensor:
565 """
566 此时的 content 已经是带上了对话格式
567 """
568 if role == "user":
569 audio_gen_slot_token = audio_delay_slot_token = self.audio_user_slot_token
570 truncation = False
571 else:
572 audio_gen_slot_token = self.audio_assistant_gen_slot_token
573 audio_delay_slot_token = self.audio_assistant_delay_slot_token
574
575 if len(audio_codes_list):
576 n_vq = audio_codes_list[0].shape[1]
577 else:
578 n_vq = self.model_config.n_vq
579
580 if len(audio_codes_list) > 1 and AUDIO_PLACEHOLDER in content:
581 content, audio_codes_list = self._merge_consecutive_audio_placeholders(
582 content, audio_codes_list
583 )
584 content = self._replace_audio_placeholders(
585 content=content,
586 lengths=[len(audio_codes) for audio_codes in audio_codes_list],
587 n_vq=n_vq,
588 gen_slot_token=audio_gen_slot_token,
589 delay_slot_token=audio_delay_slot_token,
590 audio_start_token=self.audio_start_token,
591 audio_end_token=self.audio_end_token,
592 )
593 text_codes = torch.tensor(
594 self.tokenizer.encode(content),
595 device=audio_codes_list[0].device if audio_codes_list else None,
596 )
597
598 audio_start_indices = torch.where(
599 text_codes == self.model_config.audio_start_token_id
600 )[0]
601 audio_end_indices = torch.where(
602 text_codes == self.model_config.audio_end_token_id
603 )[0]
604 if len(audio_start_indices) != len(audio_codes_list) or len(
605 audio_end_indices
606 ) != len(audio_codes_list):
607 raise ValueError(
608 "Audio placeholders do not match the provided audio codes list."
609 )
610
611 delay_audio_codes_list = []
612 if len(audio_codes_list) == 0:
613 delay_audio_codes_list = torch.full(
614 (len(text_codes), n_vq),
615 self.model_config.audio_pad_code,
616 device=text_codes.device,
617 dtype=text_codes.dtype,
618 )
619 else:
620 prefix_idx = 0
621 for audio_start_idx_t, audio_end_idx_t, audio_codes in zip(
622 audio_start_indices, audio_end_indices, audio_codes_list
623 ):
624 audio_start_idx = int(audio_start_idx_t.item())
625 audio_end_idx = int(audio_end_idx_t.item())
626 delay_audio_codes = self.apply_delay_pattern(
627 audio_codes, self.model_config.audio_pad_code
628 )
629 pad_codes = torch.full(
630 (audio_start_idx - prefix_idx + 1, n_vq),
631 self.model_config.audio_pad_code,
632 device=audio_codes.device,
633 dtype=audio_codes.dtype,
634 )
635 delay_audio_codes_list.extend([pad_codes, delay_audio_codes])
636 prefix_idx = audio_end_idx
637
638 if truncation:
639 delay_audio_codes_list[-1] = delay_audio_codes_list[-1][
640 : -(n_vq - 1), :
641 ]
642 else:
643 last_audio_end_idx = int(audio_end_indices[-1].item())
644 pad_codes = torch.full(
645 (len(text_codes) - last_audio_end_idx, n_vq),
646 self.model_config.audio_pad_code,
647 device=audio_codes_list[0].device,
648 dtype=audio_codes_list[0].dtype,
649 )
650 delay_audio_codes_list.append(pad_codes)
651
652 delay_audio_codes_list = torch.cat(delay_audio_codes_list)
653
654 if text_codes.shape[0] != delay_audio_codes_list.shape[0]:
655 text_codes = text_codes[: delay_audio_codes_list.shape[0]]
656
657 unified_codes = torch.cat(
658 [text_codes.unsqueeze(1), delay_audio_codes_list], dim=1
659 )
660 return unified_codes
661
662 def _parse_text_codes(self, start_length, text_codes):
663 text = cast(str, self.tokenizer.decode(text_codes))
664 prefix = cast(str, self.tokenizer.decode(text_codes[:start_length]))
665 text = text[len(prefix) :]
666
667 AUDIO_PATTERN = re.compile(
668 rf"(?:{self.audio_start_token})?"
669 rf"(?:{self.audio_assistant_gen_slot_token})*"
670 rf"(?:{self.audio_assistant_delay_slot_token})*"
671 rf"{self.audio_end_token}"
672 )
673
674 def normalize_audio_segments(text: str) -> str:
675 def repl(match: re.Match) -> str:
676 seg = match.group(0)
677 # Replace with <|audio|> if gen_slot is present in the segment;
678 if self.audio_assistant_gen_slot_token in seg:
679 return AUDIO_PLACEHOLDER
680 # Otherwise, remove it.
681 return ""
682
683 return AUDIO_PATTERN.sub(repl, text)
684
685 return normalize_audio_segments(text)
686
687 def _parse_audio_codes(self, start_length, audio_codes):
688 # De-delay back to [T', n_vq]
689 audio_codes = self.apply_de_delay_pattern(audio_codes)
690
691 # Rows that are all pad are separators between real audio segments.
692 is_pad = (audio_codes == self.model_config.audio_pad_code).all(dim=1)
693 non_pad = ~is_pad
694 if not non_pad.any():
695 return []
696
697 idx = torch.nonzero(non_pad).squeeze(1)
698 breaks = torch.where(idx[1:] != idx[:-1] + 1)[0] + 1
699 if breaks.numel() == 0:
700 segments_idx = [idx]
701 else:
702 segments_idx = torch.split(idx, breaks.tolist())
703
704 audio_codes_list = [audio_codes[s] for s in segments_idx]
705
706 # Batch-decode all audio segments together.
707 decoded_audio_list = self.decode_audio_codes(audio_codes_list)
708
709 # Keep codec causal context by decoding the whole first segment first,
710 # then trim at waveform level according to start_length ratio.
711 if (
712 start_length > 0
713 and len(audio_codes_list) > 0
714 and len(decoded_audio_list) > 0
715 ):
716 first_codes_length = audio_codes_list[0].shape[0]
717 if first_codes_length > 0:
718 trim_ratio = max(
719 0.0, min(float(start_length) / float(first_codes_length), 1.0)
720 )
721 first_audio = decoded_audio_list[0]
722 if trim_ratio >= 1.0:
723 decoded_audio_list = decoded_audio_list[1:]
724 elif trim_ratio > 0.0:
725 trim_samples = int(first_audio.shape[-1] * trim_ratio)
726 decoded_audio_list[0] = first_audio[..., trim_samples:]
727
728 return decoded_audio_list
729
730 def decode(self, output: List[Tuple[int, torch.Tensor]]):
731 """
732 1. 这里不管怎样,都需要一个完整的 assistant generation ids;
733 2. 支持从任意位置进行截断;
734 """
735
736 genearted_messages = []
737 for start_length, generation_ids in output:
738 content = self._parse_text_codes(start_length, generation_ids[:, 0])
739 audio_codes_list = self._parse_audio_codes(
740 start_length, generation_ids[:, 1:]
741 )
742 if content == "":
743 message = None
744 else:
745 message = AssistantMessage(
746 content=content,
747 audio_codes_list=cast(
748 List[Union[str, torch.Tensor]], audio_codes_list
749 ),
750 )
751 genearted_messages.append(message)
752 return genearted_messages
753
754 @staticmethod
755 def loudness_normalize(
756 wav: torch.Tensor,
757 target_dbfs: float = -20,
758 gain_range: tuple[float, float] = (-3.0, 3.0),
759 ) -> torch.Tensor:
760 wav = wav.to(torch.float32)
761 if wav.numel() == 0:
762 return wav
763 current_dbfs = 10.0 * torch.log10(torch.mean(wav**2) + 1e-9)
764 gain = float(target_dbfs - current_dbfs)
765 gain = max(gain_range[0], min(gain, gain_range[1]))
766 factor = 10.0 ** (gain / 20.0)
767 return wav * factor
768
769 def _get_audio_tokenizer_device(self) -> torch.device:
770 """Best-effort device inference for `self.audio_tokenizer`.
771
772 Notes:
773 - Old TAC wrapper exposed `.device`, but standard `torch.nn.Module` does not.
774 - New MossAudioTokenizerModel is a `PreTrainedModel`; parameters define its device.
775 """
776
777 audio_tokenizer = getattr(self, "audio_tokenizer", None)
778 if audio_tokenizer is None:
779 logger.warning(
780 "audio_tokenizer is not set on processor. Using CPU as default."
781 )
782 return torch.device("cpu")
783
784 device_attr = getattr(audio_tokenizer, "device", None)
785 if isinstance(device_attr, torch.device):
786 return device_attr
787
788 try:
789 return next(audio_tokenizer.parameters()).device
790 except StopIteration:
791 # No parameters (shouldn't happen for real models); default to CPU.
792 logger.warning(
793 "No parameters found on audio_tokenizer. Using CPU as default."
794 )
795 return torch.device("cpu")
796
797 def encode_audios_from_wav(
798 self,
799 wav_list: List[torch.Tensor],
800 sampling_rate: int,
801 n_vq: Optional[int] = None,
802 ):
803 if self.audio_tokenizer is None:
804 raise RuntimeError("audio_tokenizer is not set on processor.")
805 audio_tokenizer = self.audio_tokenizer
806
807 if isinstance(wav_list, torch.Tensor):
808 wav_list = [wav_list]
809 wav_list_ = []
810 resample = False
811 if sampling_rate != self.model_config.sampling_rate:
812 resample = True
813 device = self._get_audio_tokenizer_device()
814 for wav in wav_list:
815 if wav.shape[0] > 1:
816 wav = torch.mean(wav, dim=0, keepdim=True)
817 if resample:
818 wav = torchaudio.functional.resample(
819 waveform=wav,
820 orig_freq=sampling_rate,
821 new_freq=self.model_config.sampling_rate,
822 )
823 wav = wav.to(device)
824 wav_list_.append(self.loudness_normalize(wav.squeeze(0)))
825
826 # New MossAudioTokenizerModel API: prefer batch_encode(list[wav])
827 if hasattr(audio_tokenizer, "batch_encode"):
828 enc = audio_tokenizer.batch_encode(wav_list_, num_quantizers=n_vq)
829 audio_codes = enc.audio_codes # (NQ, B, T)
830 audio_codes_lengths = enc.audio_codes_lengths # (B,)
831 else:
832 # Fallback: use encode() with explicit padding.
833 max_len = max(int(wav.shape[-1]) for wav in wav_list_)
834 input_values = torch.zeros(
835 len(wav_list_), 1, max_len, device=device, dtype=torch.float32
836 )
837 padding_mask = torch.zeros(
838 len(wav_list_), max_len, device=device, dtype=torch.bool
839 )
840 for i, wav in enumerate(wav_list_):
841 this_len = int(wav.shape[-1])
842 input_values[i, 0, :this_len] = wav
843 padding_mask[i, :this_len] = True
844 enc = audio_tokenizer.encode(
845 input_values,
846 padding_mask=padding_mask,
847 num_quantizers=n_vq,
848 return_dict=True,
849 )
850 audio_codes = enc.audio_codes
851 audio_codes_lengths = enc.audio_codes_lengths
852
853 if audio_codes is None or audio_codes_lengths is None:
854 raise RuntimeError(
855 "audio_tokenizer.encode() returned empty outputs (audio_codes/audio_codes_lengths)."
856 )
857
858 # Keep processor's historical contract: list[Tensor] with shape (T, NQ)
859 # and on CPU (so downstream text/audio packing remains device-agnostic).
860 codes_list: List[torch.Tensor] = []
861 for i in range(int(audio_codes.shape[1])):
862 length_i = int(audio_codes_lengths[i].item())
863 codes_i = (
864 audio_codes[:, i, :length_i]
865 .transpose(0, 1)
866 .contiguous()
867 .to(torch.long)
868 .cpu()
869 )
870 codes_list.append(codes_i)
871 return codes_list
872
873 def encode_audios_from_path(
874 self, wav_path_list: Union[str, List[str]], n_vq: Optional[int] = None
875 ):
876 if isinstance(wav_path_list, str):
877 wav_path_list = [wav_path_list]
878
879 if len(wav_path_list) == 0:
880 raise ValueError("Empty wav_path_list")
881
882 # Load + (if needed) resample each wav independently, so callers can
883 # pass a heterogeneous batch of files while still benefiting from
884 # audio_tokenizer.batch_encode.
885 target_sr = int(self.model_config.sampling_rate)
886 wav_list: List[torch.Tensor] = []
887 for wav_path in wav_path_list:
888 wav, sr = torchaudio.load(wav_path)
889 if int(sr) != target_sr:
890 wav = torchaudio.functional.resample(
891 waveform=wav,
892 orig_freq=int(sr),
893 new_freq=target_sr,
894 )
895 wav_list.append(wav)
896
897 return self.encode_audios_from_wav(wav_list, target_sr, n_vq)
898
899 def decode_audio_codes(
900 self, audio_tokens_list: Union[torch.Tensor, List[torch.Tensor]]
901 ):
902 if self.audio_tokenizer is None:
903 raise RuntimeError("audio_tokenizer is not set on processor.")
904 audio_tokenizer = self.audio_tokenizer
905
906 if isinstance(audio_tokens_list, torch.Tensor):
907 audio_tokens_list = [audio_tokens_list]
908 if len(audio_tokens_list) == 0:
909 return []
910
911 device = self._get_audio_tokenizer_device()
912
913 # Processor uses (T, NQ); MossAudioTokenizer expects (NQ, T) (or (NQ, B, T)).
914 codes_list = [
915 codes.transpose(0, 1).contiguous().to(device=device, dtype=torch.long)
916 for codes in audio_tokens_list
917 ]
918
919 # Fallback: pad to (NQ, B, T) + mask, then decode.
920 nq = int(codes_list[0].shape[0])
921 max_t = max(int(c.shape[1]) for c in codes_list)
922 audio_codes = torch.zeros(
923 nq, len(codes_list), max_t, device=device, dtype=torch.long
924 )
925 padding_mask = torch.zeros(
926 len(codes_list), max_t, device=device, dtype=torch.bool
927 )
928 for i, c in enumerate(codes_list):
929 t = int(c.shape[1])
930 audio_codes[:, i, :t] = c
931 padding_mask[i, :t] = True
932 dec = audio_tokenizer.decode(
933 audio_codes, padding_mask=padding_mask, return_dict=True, chunk_duration=8
934 )
935 audio = dec.audio
936 audio_lengths = dec.audio_lengths
937
938 if audio is None or audio_lengths is None:
939 raise RuntimeError(
940 "audio_tokenizer.decode() returned empty outputs (audio/audio_lengths)."
941 )
942
943 # Return historical contract: list of 1D waveforms (T,)
944 wav_list: List[torch.Tensor] = []
945 for i in range(int(audio.shape[0])):
946 length_i = int(audio_lengths[i].item())
947 wav = audio[i, 0, :length_i].contiguous().to(torch.float32).cpu()
948 wav_list.append(wav)
949 return wav_list
950