EdgeIn-v1 / native_clock_talker.py
chenjz24's picture
Upload folder using huggingface_hub
f74eb65 verified
Raw History Blame
12.8 kB
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)