Audio-Text-to-Text
Transformers
Safetensors
Chinese
English
edgeinstant
feature-extraction
audio
speech-recognition
speech-translation
audio-question-answering
custom_code
Instructions to use chenjz24/EdgeIn-v1 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use chenjz24/EdgeIn-v1 with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("chenjz24/EdgeIn-v1", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download native_clock_talker.py from chenjz24/EdgeIn-v1: direct link, hf CLI and curl.
- Browser
- Download file 12.8 kB
-
https://huggingface.co/chenjz24/EdgeIn-v1/resolve/f74eb655f4afe32e298a9f0f657f002b47755995/native_clock_talker.py
- Command line
-
hf download hf://chenjz24/EdgeIn-v1@f74eb655f4afe32e298a9f0f657f002b47755995/native_clock_talker.py
-
curl -L -o native_clock_talker.py https://huggingface.co/chenjz24/EdgeIn-v1/resolve/f74eb655f4afe32e298a9f0f657f002b47755995/native_clock_talker.py
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") | |
| 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)) | |
| 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 | |
| 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) | |
| 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) | |