"""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()