# Copyright 2026 The Alibaba Qwen team. # SPDX-License-Identifier: Apache-2.0 """Qwen3-TTS 12.5 Hz waveform decoder for the EdgeInstant model.""" from __future__ import annotations import json import math from pathlib import Path import torch from torch import nn from torch.nn import functional as F from transformers.models.qwen3_omni_moe.configuration_qwen3_omni_moe import ( Qwen3OmniMoeCode2WavConfig, ) from transformers.models.qwen3_omni_moe.modeling_qwen3_omni_moe import ( Qwen3OmniMoeCausalConvNet, Qwen3OmniMoeCode2WavDecoderResidualUnit, Qwen3OmniMoeCode2WavTransformerModel, Qwen3OmniMoeConvNeXtBlock, SnakeBeta, ) class _CausalTransConvNet(nn.Module): def __init__(self, in_channels, out_channels, kernel_size, stride=1): super().__init__() self.conv = nn.ConvTranspose1d(in_channels, out_channels, kernel_size, stride=stride) self.right_pad = kernel_size - stride def forward(self, hidden): hidden = self.conv(hidden) if self.right_pad: hidden = hidden[..., : -self.right_pad] return hidden.contiguous() class _Codebook(nn.Module): def __init__(self, dim, size): super().__init__() self.cluster_usage = nn.Parameter(torch.ones(size)) self.embedding_sum = nn.Parameter(torch.zeros(size, dim)) def forward(self, codes): embedding = self.embedding_sum / self.cluster_usage.clamp(min=1e-5)[:, None] return F.embedding(codes, embedding) class _VectorQuantization(nn.Module): def __init__(self, dim, size): super().__init__() self._codebook = _Codebook(dim, size) def forward(self, codes): return self._codebook(codes).transpose(1, 2) class _ResidualVectorQuantization(nn.Module): def __init__(self, count, dim, size): super().__init__() self.layers = nn.ModuleList([_VectorQuantization(dim, size) for _ in range(count)]) def forward(self, codes): quantized = torch.zeros([1], device=codes.device)[0] for index, layer in enumerate(self.layers): quantized = quantized + layer(codes[:, index]) return quantized class _ResidualVectorQuantizer(nn.Module): def __init__(self, count, config, *, inference_only=False): super().__init__() dim = config.codebook_dim // 2 self.input_proj = ( None if inference_only else nn.Conv1d(config.codebook_dim, dim, 1, bias=False) ) self.output_proj = nn.Conv1d(dim, config.codebook_dim, 1, bias=False) self.vq = _ResidualVectorQuantization(count, dim, config.codebook_size) def forward(self, codes): return self.output_proj(self.vq(codes)) class _SplitResidualVectorQuantizer(nn.Module): def __init__(self, config, *, inference_only=False): super().__init__() self.rvq_first = _ResidualVectorQuantizer(1, config, inference_only=inference_only) self.rvq_rest = _ResidualVectorQuantizer( config.num_quantizers - 1, config, inference_only=inference_only, ) def forward(self, codes): quantized = self.rvq_first(codes[:, :1]) quantized += self.rvq_rest(codes[:, 1:]) return quantized class _ProjectedTransformer(Qwen3OmniMoeCode2WavTransformerModel): def __init__(self, config): super().__init__(config) self.input_proj = nn.Linear(config.latent_dim, config.hidden_size) self.output_proj = nn.Linear(config.hidden_size, config.latent_dim) def forward(self, hidden): positions = torch.arange(hidden.shape[1], device=hidden.device) distances = positions[:, None] - positions[None, :] mask = ((distances >= 0) & (distances < self.config.sliding_window))[None, None] output = super().forward( inputs_embeds=self.input_proj(hidden), attention_mask={"sliding_attention": mask.expand(hidden.shape[0], -1, -1, -1)}, use_cache=False, ) return self.output_proj(output.last_hidden_state) class _DecoderBlock(nn.Module): def __init__(self, config, index): super().__init__() input_dim = config.decoder_dim // 2**index output_dim = input_dim // 2 rate = config.upsample_rates[index] self.block = nn.ModuleList( [SnakeBeta(input_dim), _CausalTransConvNet(input_dim, output_dim, 2 * rate, rate)] + [Qwen3OmniMoeCode2WavDecoderResidualUnit(output_dim, dilation) for dilation in (1, 3, 9)] ) def forward(self, hidden): for block in self.block: hidden = block(hidden) return hidden class EdgeInstantCodecDecoder(nn.Module): """Decode integer codes ``[batch, frames, 16]`` to float audio ``[batch, samples]``. ``config`` is the complete Qwen3-TTS speech tokenizer configuration. Weights match its ``decoder.*`` tensors after removing that prefix. Setting its top-level ``inference_only`` field omits the unused quantizer input projections. """ def __init__(self, config: dict): super().__init__() self.sample_rate = int(config["output_sample_rate"]) self.samples_per_frame = int(config["decode_upsample_rate"]) decoder_config = dict(config["decoder_config"]) decoder_config["rope_parameters"] = { "rope_type": "default", "rope_theta": decoder_config.pop("rope_theta", 10000), } self.config = Qwen3OmniMoeCode2WavConfig(**decoder_config) self.config._attn_implementation = "sdpa" cfg = self.config if math.prod([*cfg.upsample_rates, *cfg.upsampling_ratios]) != self.samples_per_frame: raise ValueError("decode_upsample_rate must match the waveform decoder upsampling factors") self.pre_transformer = _ProjectedTransformer(cfg) self.quantizer = _SplitResidualVectorQuantizer( cfg, inference_only=bool(config.get("inference_only", False)), ) self.pre_conv = Qwen3OmniMoeCausalConvNet(cfg.codebook_dim, cfg.latent_dim, 3) self.upsample = nn.ModuleList( [ nn.ModuleList( [ _CausalTransConvNet(cfg.latent_dim, cfg.latent_dim, factor, factor), Qwen3OmniMoeConvNeXtBlock(cfg.latent_dim), ] ) for factor in cfg.upsampling_ratios ] ) output_dim = cfg.decoder_dim // 2 ** len(cfg.upsample_rates) self.decoder = nn.ModuleList( [Qwen3OmniMoeCausalConvNet(cfg.latent_dim, cfg.decoder_dim, 7)] + [_DecoderBlock(cfg, index) for index in range(len(cfg.upsample_rates))] + [SnakeBeta(output_dim), Qwen3OmniMoeCausalConvNet(output_dim, 1, 7)] ) def forward(self, codes: torch.Tensor) -> torch.Tensor: if codes.ndim != 3 or codes.shape[-1] != self.config.num_quantizers: raise ValueError(f"codes must be [batch, frames, {self.config.num_quantizers}]") if codes.shape[1] == 0: return self.pre_conv.conv.weight.new_empty((codes.shape[0], 0), dtype=torch.float32) codes = codes.to(device=self.pre_conv.conv.weight.device, dtype=torch.long) hidden = self.quantizer(codes.transpose(1, 2)) hidden = self.pre_conv(hidden).transpose(1, 2) hidden = self.pre_transformer(hidden).transpose(1, 2) for blocks in self.upsample: for block in blocks: hidden = block(hidden) for block in self.decoder: hidden = block(hidden) return hidden.squeeze(1).clamp(min=-1, max=1).float() def decode( self, codes: torch.Tensor, *, chunk_size: int = 300, left_context_frames: int = 25 ) -> torch.Tensor: """Decode with the speech tokenizer's chunking and left context convention.""" if chunk_size <= 0 or left_context_frames < 0: raise ValueError("chunk_size must be positive and left_context_frames nonnegative") if codes.shape[1] == 0: return self(codes) waves = [] for start in range(0, codes.shape[1], chunk_size): context = min(start, left_context_frames) wave = self(codes[:, start - context : start + chunk_size]) waves.append(wave[:, context * self.samples_per_frame :]) return torch.cat(waves, dim=-1) def load_codec_decoder( directory: str | Path, *, dtype: torch.dtype = torch.float32, device: str = "cpu" ) -> EdgeInstantCodecDecoder: """Load the decoder tensors from a local Qwen3-TTS speech tokenizer directory.""" from safetensors import safe_open directory = Path(directory) with (directory / "config.json").open() as handle: config = json.load(handle) decoder = EdgeInstantCodecDecoder(config).to(dtype=dtype, device=device) rotary = decoder.pre_transformer.rotary_emb rotary.inv_freq, rotary.attention_scaling = rotary.compute_default_rope_parameters( rotary.config, torch.device(device) ) rotary.original_inv_freq = rotary.inv_freq.clone() with safe_open(directory / "model.safetensors", framework="pt", device="cpu") as checkpoint: state = { key.removeprefix("decoder."): checkpoint.get_tensor(key) for key in checkpoint.keys() if key.startswith("decoder.") } decoder.load_state_dict(state, strict=True) return decoder.eval()