from __future__ import annotations from collections import deque from dataclasses import dataclass, field import json from pathlib import Path import torch from safetensors import safe_open from torch import nn from .teacher import TeacherAcousticModel from .sequence import Action from .sampling import sample_token from .generation import AdvanceStatus, GenerationResult class NativeClockTalker(nn.Module): """Teacher-initialized acoustic model with an incremental native text clock.""" codebook_size = 2048 hidden_size = 1024 eos_class = 2048 supported_languages = ("chinese", "english", "auto") @classmethod @torch.no_grad() def from_teacher( cls, teacher_path: str | Path, native_tokenizer_path: str | Path, speaker_vector_path: str | Path, *, dtype: torch.dtype = torch.float32, attn_implementation: str = "eager", ) -> "NativeClockTalker": teacher_path = Path(teacher_path) source = TeacherAcousticModel.from_teacher( teacher_path, native_tokenizer_path, speaker_vector_path, dtype=dtype, attn_implementation=attn_implementation, ) model = cls() for name in ("backbone", "text_embedding", "text_projection", "codec_history_embeddings", "residual_predictor", "q0_head"): setattr(model, name, getattr(source, name)) model.codec_bos = source.codec_bos model.native_teacher_ids = source.native_teacher_ids.tolist() model.native_teacher_offsets = source.native_teacher_offsets.tolist() model.codec_eos_head = nn.Linear(model.hidden_size, 1, bias=False, dtype=dtype) config = json.loads((teacher_path / "config.json").read_text()) talker = config["talker_config"] vocabulary = json.loads((teacher_path / "vocab.json").read_text()) def text(ids): return model.text_projection(model.text_embedding(torch.tensor(ids))) tts_bos, tts_eos, tts_pad = text([config["tts_bos_token_id"], config["tts_eos_token_id"], config["tts_pad_token_id"]]) role = text([config["im_start_token_id"], config["assistant_token_id"], vocabulary["Ċ"]]) with safe_open(teacher_path / "model.safetensors", framework="pt", device="cpu") as weights: codec = weights.get_tensor("talker.model.codec_embedding.weight").to(dtype=dtype) language = [talker["codec_think_id"], talker["codec_think_bos_id"], talker["codec_language_id"]["chinese"], talker["codec_think_eos_id"]] prefix = torch.cat((role, codec[language] + tts_pad, (source.speaker_vector + tts_pad).unsqueeze(0), (codec[talker["codec_pad_id"]] + tts_bos).unsqueeze(0)), dim=0) english = prefix.clone() english[5] = codec[talker["codec_language_id"]["english"]] + tts_pad auto = torch.cat((role, codec[[talker["codec_nothink_id"], talker["codec_think_bos_id"], talker["codec_think_eos_id"]]] + tts_pad, prefix[7:]), dim=0) model.codec_eos_head.weight.copy_( weights.get_tensor("talker.codec_head.weight")[talker["codec_eos_token_id"]].unsqueeze(0) ) model.register_buffer("prefix", prefix.unsqueeze(0), persistent=False) model.register_buffer("prefix_english", english.unsqueeze(0), persistent=False) model.register_buffer("prefix_auto", auto.unsqueeze(0), persistent=False) model.register_buffer("tts_eos", tts_eos.view(1, 1, -1), persistent=False) model.register_buffer("tts_pad", tts_pad.view(1, 1, -1), persistent=False) return model def new_state(self, **kwargs) -> "NativeClockState": return NativeClockState(model=self, **kwargs) def prefix_for_language(self, language: str) -> torch.Tensor: if language == "chinese": return self.prefix if language == "english": return self.prefix_english if language == "auto": return self.prefix_auto raise ValueError(f"Unsupported speech language: {language}") def mapped_teacher_ids(self, native_id: int) -> list[int]: if not 0 <= native_id < len(self.native_teacher_offsets) - 1: return [] start, end = self.native_teacher_offsets[native_id:native_id + 2] return self.native_teacher_ids[start:end] def history_features(self, codes: torch.Tensor) -> torch.Tensor: return sum(embedding(codes[..., group]) for group, embedding in enumerate(self.codec_history_embeddings)) @dataclass class NativeClockState: model: NativeClockTalker language: str = "chinese" max_frames: int = 1536 max_events: int = 4096 do_sample: bool = False top_k: int = 50 top_p: float = 1.0 temperature: float = 0.9 repetition_penalty: float = 1.05 sampling_generator: torch.Generator | None = None pending: deque = field(default_factory=deque) native_token_ids: list[int] = field(default_factory=list) native_remaining: list[int] = field(default_factory=list) teacher_token_ids: list[int] = field(default_factory=list) teacher_token_owners: list[int] = field(default_factory=list) consumed_teacher_tokens: int = 0 consumed_tokens: int = 0 frames: list[torch.Tensor] = field(default_factory=list) actions: list[int] = field(default_factory=list) frame_conditioned_text_tokens: list[int] = field(default_factory=list) frame_arrived_text_tokens: list[int] = field(default_factory=list) frame_available_teacher_tokens: list[int] = field(default_factory=list) frame_consumed_teacher_tokens: list[int] = field(default_factory=list) cache: object | None = None previous_codes: torch.Tensor | None = None input_ended: bool = False text_end_consumed: bool = False text_end_at_codec_frame: int | None = None ended: bool = False cancelled: bool = False event_count: int = 0 suppressed_early_eos: int = 0 q0_eos_generated: bool = False status: AdvanceStatus = AdvanceStatus.WAIT_INPUT @property def consumed_native_tokens(self) -> int: return self.consumed_tokens def _update_consumed_native(self) -> None: while self.consumed_tokens < len(self.native_remaining) and self.native_remaining[self.consumed_tokens] == 0: self.consumed_tokens += 1 def push_token(self, native_id: int) -> None: if self.input_ended or self.ended: raise RuntimeError("Cannot append native text after the input has ended") pieces = self.model.mapped_teacher_ids(int(native_id)) index = len(self.native_token_ids) self.native_token_ids.append(int(native_id)) self.native_remaining.append(len(pieces)) self.teacher_token_ids.extend(pieces) self.teacher_token_owners.extend([index] * len(pieces)) self.pending.extend((piece, index) for piece in pieces) self._update_consumed_native() def end_input(self) -> None: self.input_ended = True finish = end_input def cancel(self) -> None: self.cancelled = True self.ended = True self.status = AdvanceStatus.FORCED_STOP self.cache = None self.previous_codes = None self.pending.clear() def _result(self, status: AdvanceStatus) -> GenerationResult: self.status = status return GenerationResult(self.frames, status, self.actions, self.consumed_tokens, self.text_end_consumed) @torch.inference_mode() def advance(self, max_new_frames: int = 4) -> GenerationResult: if self.ended: return self._result(self.status) produced = 0 while produced < max_new_frames: if len(self.frames) >= self.max_frames or self.event_count >= self.max_events: self.ended = True return self._result(AdvanceStatus.FORCED_STOP) if self.pending: teacher_id, native_index = self.pending.popleft() ids = torch.tensor([[teacher_id]], device=self.model.q0_head.weight.device) text = self.model.text_projection(self.model.text_embedding(ids)) self.native_remaining[native_index] -= 1 self._update_consumed_native() self.consumed_teacher_tokens += 1 elif not self.input_ended: return self._result(AdvanceStatus.WAIT_INPUT) elif not self.text_end_consumed: text = self.model.tts_eos self.text_end_consumed = True self.text_end_at_codec_frame = len(self.frames) else: text = self.model.tts_pad if self.cache is None: inputs = torch.cat((self.model.prefix_for_language(self.language), self.model.codec_bos.view(1, 1, -1) + text), dim=1) else: inputs = self.model.history_features(self.previous_codes).unsqueeze(1) + text output = self.model.backbone(inputs_embeds=inputs, past_key_values=self.cache, use_cache=True) self.cache = output.past_key_values hidden = output.last_hidden_state[:, -1] logits = torch.cat((self.model.q0_head(hidden), self.model.codec_eos_head(hidden)), dim=-1).float() if self.frames and self.repetition_penalty != 1.0: seen = torch.tensor(list({int(frame[0]) for frame in self.frames}), device=logits.device) scores = logits[:, seen] logits[:, seen] = torch.where(scores < 0, scores * self.repetition_penalty, scores / self.repetition_penalty) eos_allowed = self.input_ended and not self.pending and self.text_end_consumed if not eos_allowed: self.suppressed_early_eos += int(logits.argmax(-1).item() == self.model.eos_class) logits[:, self.model.eos_class] = -torch.inf q0 = sample_token(logits, do_sample=self.do_sample, top_k=self.top_k, top_p=self.top_p, temperature=self.temperature, generator=self.sampling_generator) self.event_count += 1 if int(q0.item()) == self.model.eos_class: self.ended = True self.q0_eos_generated = True self.actions.append(int(Action.END)) return self._result(AdvanceStatus.END_AUDIO) codes = self.model.residual_predictor.generate(hidden, q0, do_sample=self.do_sample, top_k=self.top_k, top_p=self.top_p, temperature=self.temperature, generator=self.sampling_generator) self.previous_codes = codes self.frames.append(codes[0].detach().cpu()) self.actions.append(int(Action.EMIT)) self.frame_conditioned_text_tokens.append(self.consumed_tokens) self.frame_arrived_text_tokens.append(len(self.native_token_ids)) self.frame_available_teacher_tokens.append(len(self.teacher_token_ids)) self.frame_consumed_teacher_tokens.append(self.consumed_teacher_tokens) produced += 1 return self._result(AdvanceStatus.PRODUCED) def speech_unit_report(self) -> dict: return dict(native_token_ids=self.native_token_ids, native_token_count=len(self.native_token_ids), teacher_token_ids=self.teacher_token_ids, teacher_token_count=len(self.teacher_token_ids), teacher_token_owners=self.teacher_token_owners, consumed_native_tokens=self.consumed_tokens, consumed_teacher_tokens=self.consumed_teacher_tokens) def report(self) -> dict: return dict(**self.speech_unit_report(), model_variant="teacher_initialized_native_clock_talker_v1", language=self.language, terminal_status=self.status.value, q0_eos_generated=self.q0_eos_generated, q0_eos_teacher_id=2150, suppressed_early_eos=self.suppressed_early_eos, text_end_consumed=self.text_end_consumed, text_end_at_codec_frame=self.text_end_at_codec_frame, frame_arrived_text_tokens=self.frame_arrived_text_tokens, frame_conditioned_text_tokens=self.frame_conditioned_text_tokens, frame_available_teacher_tokens=self.frame_available_teacher_tokens, frame_consumed_teacher_tokens=self.frame_consumed_teacher_tokens, generated_codec_frames=len(self.frames), cancelled=self.cancelled)