AudioDecisionModel / server_runtime /qwen_audio_encoder.py
gojiteji's picture
Audio Decision Model demo
148af80
Raw History Blame Contribute Delete
7.24 kB
"""Frozen, transcription-free Qwen3-ASR audio feature extraction.
This adapter deliberately loads no tokenizer and never calls ``generate`` or
``transcribe``. It runs only Qwen3-ASR's native audio tower and its official
final audio projection, then pools observed frames for the shared speech/sound
feature-cache contract.
"""
from __future__ import annotations
from dataclasses import dataclass
from math import ceil
from pathlib import Path
from typing import Sequence
import numpy as np
import torch
QWEN_SAMPLE_RATE = 16_000
QWEN_INPUT_HOP_SECONDS = 160 / QWEN_SAMPLE_RATE
QWEN_CONV_STRIDE = 2**3
QWEN_FRAME_SECONDS = QWEN_INPUT_HOP_SECONDS * QWEN_CONV_STRIDE
TAIL_SECONDS = 0.5
TAIL_FRAMES = ceil(TAIL_SECONDS / QWEN_FRAME_SECONDS)
SUPPORTED_MODEL_TYPES = {"qwen3_asr"}
SUPPORTED_AUDIO_MODEL_TYPES = {"qwen3_asr_audio_encoder"}
@dataclass(frozen=True)
class QwenEncoderMetadata:
"""Pinned-model and geometry information needed to reuse feature caches."""
model_id: str
revision: str
model_type: str
audio_model_type: str
input_hop_seconds: float
encoder_frame_seconds: float
tail_seconds: float
tail_frames: int
frame_dim: int
projected_dim: int
pooled_dim: int
class QwenAudioEncoder:
"""Qwen3-ASR native audio tower plus padding-safe mean/tail pooling.
``encode`` accepts finite, mono 16-kHz waveforms and returns one
``[mean(projected frames), mean(last 0.5 s projected frames)]`` feature
vector per input. The 0.5-second tail is exactly seven Qwen output frames:
10-ms log-mel hops and three stride-2 convolutional layers produce 80-ms
output steps, so ``ceil(0.5 / 0.08) == 7``.
"""
def __init__(self, model_path: str | Path, *, model_id: str, revision: str, device: str = "cpu"):
self.model_path = Path(model_path)
self.model_id = model_id
self.revision = revision
self.device = torch.device(device)
if self.device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is unavailable")
# The official qwen-asr 0.0.6 transformer backend is vendored beside
# this adapter. Its public checkpoint has a ``thinker.*`` state-dict
# layout that differs from Transformers' generic Qwen implementations.
try:
from .qwen_backend import Qwen3ASRConfig, Qwen3ASRForConditionalGeneration
from transformers import WhisperFeatureExtractor
except ImportError as exc:
raise RuntimeError(
"QwenAudioEncoder requires its vendored official Qwen3-ASR backend."
) from exc
config = Qwen3ASRConfig.from_pretrained(self.model_path, local_files_only=True)
audio_config = config.thinker_config.audio_config
if config.model_type not in SUPPORTED_MODEL_TYPES or audio_config.model_type not in SUPPORTED_AUDIO_MODEL_TYPES:
raise ValueError("checkpoint is not a compatible native Qwen3-ASR audio model")
if int(audio_config.num_mel_bins) != 128 or int(audio_config.n_window) != 50:
raise ValueError("unexpected Qwen3-ASR audio encoder frontend/window configuration")
self.feature_extractor = WhisperFeatureExtractor.from_pretrained(self.model_path, local_files_only=True)
# The public checkpoint packages the audio tower under the complete
# conditional-generation model. We discard the text model immediately
# after loading and retain only audio_tower; no tokenizer,
# language model forward, or text decoding is reachable from this API.
dtype = torch.bfloat16 if self.device.type == "cuda" else torch.float32
full_model = Qwen3ASRForConditionalGeneration.from_pretrained(
self.model_path, local_files_only=True, torch_dtype=dtype, attn_implementation="sdpa"
).to(self.device).eval()
self.audio_tower = full_model.thinker.audio_tower
del full_model
self.audio_tower.requires_grad_(False)
frame_dim = int(audio_config.d_model)
projected_dim = int(audio_config.output_dim)
self.metadata = QwenEncoderMetadata(
model_id=model_id, revision=revision, model_type=config.model_type,
audio_model_type=audio_config.model_type, input_hop_seconds=QWEN_INPUT_HOP_SECONDS,
encoder_frame_seconds=QWEN_FRAME_SECONDS, tail_seconds=TAIL_SECONDS, tail_frames=TAIL_FRAMES,
frame_dim=frame_dim, projected_dim=projected_dim, pooled_dim=2 * projected_dim,
)
@staticmethod
def _prepare(waves16k: Sequence[np.ndarray]) -> list[np.ndarray]:
if not waves16k:
raise ValueError("at least one waveform is required")
prepared = []
max_samples = QWEN_SAMPLE_RATE * 30
for wave in waves16k:
value = np.asarray(wave, dtype=np.float32)
if value.ndim != 1 or value.size == 0 or not np.isfinite(value).all():
raise ValueError("each waveform must be finite, nonempty, mono 16-kHz audio")
if value.size > max_samples:
raise ValueError("Qwen3-ASR feature adapter accepts at most 30 seconds per waveform")
prepared.append(np.ascontiguousarray(np.clip(value, -1.0, 1.0)))
return prepared
@torch.inference_mode()
def encode(self, waves16k: Sequence[np.ndarray]) -> torch.Tensor:
"""Return padding-excluded projected mean/tail features of shape ``[B, 2 * projected_dim]``."""
waves = self._prepare(waves16k)
inputs = self.feature_extractor(waves, sampling_rate=QWEN_SAMPLE_RATE, padding=True,
return_attention_mask=True, return_tensors="pt")
if "attention_mask" not in inputs or "input_features" not in inputs:
raise RuntimeError("Qwen audio feature extractor did not provide features and an attention mask")
mask = inputs["attention_mask"].to(self.device)
input_features = inputs["input_features"].to(device=self.device, dtype=next(self.audio_tower.parameters()).dtype)
# The official tower intentionally processes one waveform at a time to
# preserve precision. It already includes proj1/GELU/proj2, yielding
# 1024-dimensional audio features without any language-model call.
input_lengths = mask.sum(dim=-1).to(torch.long)
pooled = []
for feature, length in zip(input_features, input_lengths):
length = int(length)
observed = self.audio_tower(feature[:, :length], feature_lens=torch.tensor([length], device=self.device)).last_hidden_state.float()
if observed.ndim != 2 or observed.shape[-1] != self.metadata.projected_dim:
raise RuntimeError("Qwen audio tower returned an unexpected feature shape")
# The official tower uses three stride-2 convolutions. Its output
# is intrinsically padding-free for this one waveform.
if length < 1:
raise RuntimeError("Qwen audio encoder returned no observed frames")
pooled.append(torch.cat((observed.mean(dim=0), observed[-min(TAIL_FRAMES, observed.shape[0]):].mean(dim=0))))
return torch.stack(pooled).cpu()