| |
| |
| |
| |
| from __future__ import annotations |
|
|
| import os |
| from functools import lru_cache |
| from pathlib import Path |
| from typing import Any, cast |
|
|
| import torch |
| from neucodec import NeuCodec |
| from torch import Tensor, nn |
| from transformers.modeling_attn_mask_utils import AttentionMaskConverter |
| from typing_extensions import Self |
|
|
| from qai_hub_models import SampleInputsType |
| from qai_hub_models.models._shared.llama3.model import Llama3Base |
| from qai_hub_models.models._shared.llm.common import LLMIOType |
| from qai_hub_models.utils.base_model import SerializationSettings |
| from qai_hub_models.utils.base_multi_graph_model import MultiGraphWorkbenchModel |
| from qai_hub_models.utils.input_spec import InputSpec, OutputSpec |
|
|
| MODEL_ID = __name__.split(".")[-2] |
| MODEL_ASSET_VERSION = 1 |
|
|
| BACKBONE_REPO = "neuphonic/neutts-nano" |
| CODEC_REPO = "neuphonic/neucodec" |
|
|
| |
| NUM_LAYERS = 24 |
| HIDDEN_SIZE = 576 |
| NUM_ATTN_HEADS = 9 |
| NUM_KEY_VALUE_HEADS = 3 |
| HEAD_DIM = 64 |
|
|
| CONTEXT_LENGTH = 2048 |
| PREFILL_SEQ_LEN = 128 |
| DECODE_SEQ_LEN = 1 |
|
|
| SAMPLE_RATE = 24_000 |
| HOP_LENGTH = 480 |
|
|
| |
| |
| SAMPLE_PROMPT = "My name is Jo and I am a text to speech model." |
|
|
| |
| |
| |
| MASK_MIN = -1e4 |
|
|
|
|
| def build_attention_mask( |
| query_len: int, num_valid_keys: int, context_length: int = CONTEXT_LENGTH |
| ) -> Tensor: |
| """4D additive causal mask over a right-aligned ``context_length`` key window.""" |
| mask_2d = torch.zeros((1, context_length)) |
| mask_2d[:, -num_valid_keys:] = 1.0 |
| return ( |
| AttentionMaskConverter(True) |
| .to_4d( |
| mask_2d, |
| query_length=query_len, |
| key_value_length=context_length, |
| dtype=torch.float32, |
| ) |
| .clip(MASK_MIN, 0) |
| ) |
|
|
|
|
| def empty_kv_cache(kv_seq_len: int) -> list[Tensor]: |
| """Zeroed per-layer key/value cache tensors in graph input order.""" |
| caches: list[Tensor] = [] |
| for _ in range(NUM_LAYERS): |
| caches.append(torch.zeros(NUM_KEY_VALUE_HEADS, 1, HEAD_DIM, kv_seq_len)) |
| caches.append(torch.zeros(NUM_KEY_VALUE_HEADS, 1, kv_seq_len, HEAD_DIM)) |
| return caches |
|
|
|
|
| @lru_cache(maxsize=1) |
| def load_codec() -> nn.Module: |
| |
| |
| |
| codec = cast(nn.Module, NeuCodec.from_pretrained(CODEC_REPO)) |
| codec.to("cpu") |
| codec.eval() |
| return codec |
|
|
|
|
| class Backbone(Llama3Base): |
| """NeuTTS-Nano causal LM with the KV cache exposed as graph I/O. |
| |
| ``neuphonic/neutts-nano`` is a ``LlamaForCausalLM``, so it inherits the |
| repo's Llama treatment: split-head attention, rank-4 RMS norm, rotary |
| embeddings supplied as ``position_ids_cos`` / ``position_ids_sin`` inputs, |
| and per-layer ``past_key_*`` / ``past_value_*`` in and out. Only the newly |
| computed cache entries are returned; the caller slides the window. |
| """ |
|
|
| min_memory_recommended = 0 |
|
|
| def __init__( |
| self, |
| sequence_length: int = PREFILL_SEQ_LEN, |
| context_length: int = CONTEXT_LENGTH, |
| **kwargs: Any, |
| ) -> None: |
| super().__init__( |
| checkpoint=BACKBONE_REPO, |
| sequence_length=sequence_length, |
| context_length=context_length, |
| **kwargs, |
| ) |
| |
| |
| self.serialization_settings = SerializationSettings(use_pt2=False) |
|
|
| @classmethod |
| def from_pretrained( |
| cls, |
| sequence_length: int = PREFILL_SEQ_LEN, |
| context_length: int = CONTEXT_LENGTH, |
| ) -> Self: |
| return cls(sequence_length=sequence_length, context_length=context_length) |
|
|
| def get_input_spec( |
| self, |
| llm_config: dict | None = None, |
| sequence_length: int = PREFILL_SEQ_LEN, |
| context_length: int = CONTEXT_LENGTH, |
| llm_io_type: LLMIOType = LLMIOType.genie_input_ids, |
| ) -> InputSpec: |
| return self._get_input_spec( |
| num_hidden_layers=NUM_LAYERS, |
| sequence_length=sequence_length, |
| context_length=context_length, |
| hidden_size=HIDDEN_SIZE, |
| num_key_value_heads=NUM_KEY_VALUE_HEADS, |
| num_attention_heads=NUM_ATTN_HEADS, |
| head_dim=HEAD_DIM, |
| llm_io_type=llm_io_type, |
| ) |
|
|
| def get_output_spec(self) -> OutputSpec: |
| return self._get_output_spec(NUM_LAYERS) |
|
|
| def sample_graph_inputs(self, sequence_length: int) -> SampleInputsType: |
| """Inputs for one graph, with every position carrying a real token. |
| |
| The prompt is tiled to fill the window rather than right-aligned in it. |
| Masked pad positions compute arbitrary finite values that the app throws |
| away, but an on-device comparison still scores them: right-aligning a |
| 14-token prompt in a 128-wide window reports 11.8 dB PSNR where the |
| positions that matter are at 65 dB. |
| """ |
| ids = self.tokenizer(SAMPLE_PROMPT, return_tensors="pt")["input_ids"] |
| repeats = -(-sequence_length // int(ids.shape[1])) |
| input_ids = ids.repeat(1, repeats)[:, :sequence_length].to(torch.int32) |
|
|
| position_ids = torch.arange(sequence_length, dtype=torch.long).reshape(1, -1) |
| cos, sin = self.embedding.get_embedding(position_ids) |
|
|
| inputs: SampleInputsType = { |
| "input_ids": [input_ids.numpy()], |
| "attention_mask": [ |
| build_attention_mask( |
| sequence_length, sequence_length, self.context_length |
| ).numpy() |
| ], |
| "position_ids_cos": [cos.numpy()], |
| "position_ids_sin": [sin.numpy()], |
| } |
| caches = empty_kv_cache(self.context_length - sequence_length) |
| for layer in range(NUM_LAYERS): |
| inputs[f"past_key_{layer}_in"] = [caches[2 * layer].numpy()] |
| inputs[f"past_value_{layer}_in"] = [caches[2 * layer + 1].numpy()] |
| return inputs |
|
|
| def _sample_inputs_impl( |
| self, input_spec: InputSpec | None = None |
| ) -> SampleInputsType: |
| return self.sample_graph_inputs(self.sequence_length) |
|
|
|
|
| class NeuTTSNano(MultiGraphWorkbenchModel): |
| """NeuTTS-Nano backbone as two graphs of one weight-shared context binary. |
| |
| Prefill and decode are the same 340M parameters at two sequence lengths. A |
| fixed-shape graph cannot serve both, but one traced source compiles to both, |
| so the weights are uploaded once and linked into a single binary rather than |
| carried twice. |
| |
| Everything outside the transformer -- phonemization, prompt assembly, |
| sampling, cache management, and NeuCodec encode/decode -- is CPU work in |
| ``NeuTTSApp``. The codec is this model's speech tokenizer and is inherently |
| variable-length, so it is not a compiled graph. |
| """ |
|
|
| def __init__(self, backbone: Backbone) -> None: |
| self.backbone = backbone |
| self._graph_sequence_lengths = { |
| f"prompt_ar{PREFILL_SEQ_LEN}_cl{CONTEXT_LENGTH}": PREFILL_SEQ_LEN, |
| f"token_ar{DECODE_SEQ_LEN}_cl{CONTEXT_LENGTH}": DECODE_SEQ_LEN, |
| } |
|
|
| @property |
| def graph_names(self) -> list[str]: |
| return list(self._graph_sequence_lengths) |
|
|
| @property |
| def prefill_graph(self) -> str: |
| return f"prompt_ar{PREFILL_SEQ_LEN}_cl{CONTEXT_LENGTH}" |
|
|
| @property |
| def decode_graph(self) -> str: |
| return f"token_ar{DECODE_SEQ_LEN}_cl{CONTEXT_LENGTH}" |
|
|
| @property |
| def shared_source_model(self) -> bool: |
| return True |
|
|
| def get_graph_input_spec(self, graph_name: str) -> InputSpec: |
| return self.backbone.get_input_spec( |
| sequence_length=self._graph_sequence_lengths[graph_name], |
| context_length=self.backbone.context_length, |
| ) |
|
|
| def get_graph_output_spec(self, graph_name: str) -> OutputSpec: |
| return self.backbone.get_output_spec() |
|
|
| def get_graph_sample_inputs( |
| self, |
| graph_name: str, |
| input_spec: InputSpec | None = None, |
| use_channel_last_format: bool = True, |
| ) -> SampleInputsType: |
| return self.backbone.sample_graph_inputs( |
| self._graph_sequence_lengths[graph_name] |
| ) |
|
|
| def serialize_graph( |
| self, |
| graph_name: str, |
| output_dir: str | os.PathLike, |
| input_spec: InputSpec | None = None, |
| ) -> Path: |
| """Trace the backbone once; the same source compiles at every graph shape. |
| |
| Traced at the longest sequence length, so the graph handed to the |
| converter covers the widest reshapes. |
| |
| Parameters |
| ---------- |
| graph_name |
| Unused -- one shared source serves every graph. |
| output_dir |
| Directory to write the traced module into. |
| input_spec |
| Unused -- the trace shape is fixed at the longest sequence length. |
| |
| Returns |
| ------- |
| Path |
| Path to the serialized TorchScript module. |
| """ |
| seq_len = max(self._graph_sequence_lengths.values()) |
| inputs = [ |
| torch.from_numpy(value[0]) |
| for value in self.backbone.sample_graph_inputs(seq_len).values() |
| ] |
| self.backbone.eval() |
| output_path = Path(output_dir) / f"{self.name}.pt" |
| with torch.no_grad(): |
| torch.jit.save( |
| torch.jit.trace(self.backbone, inputs, check_trace=False), output_path |
| ) |
| return output_path |
|
|
| @classmethod |
| def from_pretrained(cls) -> Self: |
| return cls(Backbone.from_pretrained()) |
|
|