ChristophSchuhmann's picture
Release best Gemini-tuned Whisper Base and Small with code, normalization and evaluation
cd9b2d8 verified
Raw History Blame Contribute Delete
9.87 kB
"""Frozen native audio encoders, pooled embeddings and aligned temporal features."""
from __future__ import annotations
import sys
from pathlib import Path
import numpy as np
import torch
from torch import nn
from torch.nn import functional as F
from study_paths import ROOT, CLAP_CODE, SEED
class FrozenEncoder:
def __init__(self, spec, device):
self.spec, self.device = spec, device
self.kind = spec['backend']
self.frame_projection = None
local = Path(spec.get('local_path') or '')
if self.kind == 'clap':
sys.path.insert(0, str(CLAP_CODE / 'src'))
import open_clip
from open_clip.audio.naflex_audio import AudioNaFlexCfg, AudioNaFlexPatchify
self.model = open_clip.create_model(spec['model_name'], output_dict=True)
ck = torch.load(spec['checkpoint'], map_location='cpu', weights_only=False)
state = {k.removeprefix('module.'): v for k, v in ck.get('state_dict', ck).items()}
self.model.load_state_dict(state, strict=True)
del ck, state
self.audio_cfg = AudioNaFlexCfg.from_clip_audio_cfg(self.model.audio.cfg)
self.patchify = AudioNaFlexPatchify(self.audio_cfg, random_crop=False)
self.tokens = None
self.model.audio.encoder.vit.norm.register_forward_hook(self._capture)
elif self.kind == 'commercial':
from transformers import AutoModel
self.model = AutoModel.from_pretrained(local, local_files_only=True, trust_remote_code=True)
self.tokens = None
self.model.audio_encoder.register_forward_hook(self._capture)
elif self.kind == 'gemma':
from transformers import AutoModel, AutoProcessor
self.processor = AutoProcessor.from_pretrained(local, local_files_only=True)
self.model = AutoModel.from_pretrained(local, local_files_only=True, torch_dtype=torch.bfloat16,
vision_config=None, attn_implementation='sdpa')
elif self.kind == 'omni':
from transformers import Qwen2_5OmniProcessor, Qwen2_5OmniThinkerConfig, Qwen2_5OmniThinkerForConditionalGeneration
import json
local = local / spec.get('subfolder', '')
config = json.loads((local / 'config.json').read_text())
thinker = config.get('thinker_config', config)
self.processor = Qwen2_5OmniProcessor.from_pretrained(local, local_files_only=True)
self.model, loading = Qwen2_5OmniThinkerForConditionalGeneration.from_pretrained(
local, config=Qwen2_5OmniThinkerConfig(**thinker), local_files_only=True,
torch_dtype=torch.bfloat16, attn_implementation='sdpa', output_loading_info=True)
missing = [k for k in loading.get('missing_keys', []) if not k.startswith(('visual.', 'lm_head.'))]
if missing:
raise RuntimeError('Missing trained audio/Thinker weights: ' + str(missing[:8]))
self.model.lm_head = nn.Identity() # No vocabulary logits or generation.
self.tokens = None
self.model.model.register_forward_hook(self._capture)
else:
raise ValueError(self.kind)
self.model.to(device).eval()
self.model.requires_grad_(False)
def _capture(self, _module, _args, result):
self.tokens = result.last_hidden_state if hasattr(result, 'last_hidden_state') else result
if isinstance(self.tokens, tuple):
self.tokens = self.tokens[0]
def _compress_frames(self, tensor):
width = tensor.shape[-1]
if self.frame_projection is None:
generator = torch.Generator(device='cpu').manual_seed(SEED + width)
matrix = torch.randn(width, min(64, width), generator=generator)
self.frame_projection = torch.linalg.qr(matrix, mode='reduced').Q.to(self.device)
result = tensor.float() @ self.frame_projection
return F.pad(result, (0, 64 - result.shape[-1]))
@torch.inference_mode()
def encode(self, waves):
durations = [len(w) / 16000 for w in waves]
frames, times = [], []
with torch.autocast('cuda', dtype=torch.bfloat16):
if self.kind == 'clap':
patches = [self.patchify((torch.from_numpy(w).float()[None], 16000)) for w in waves]
width = max(len(p['patches']) for p in patches)
inputs = {}
for name in ('patches', 'patch_coord', 'patch_valid'):
shape = (len(waves), width, *patches[0][name].shape[1:])
inputs[name] = torch.zeros(shape, dtype=patches[0][name].dtype, device=self.device)
for i, p in enumerate(patches):
inputs[name][i, :len(p[name])] = p[name].to(self.device)
pooled = self.model.encode_audio(inputs, normalize=True)
if self.tokens is None or self.tokens.ndim != 3 or self.tokens.shape[1] < width:
raise RuntimeError('NaFlex temporal hook did not expose the patch sequence')
tokens = self.tokens[:, -width:]
dt = self.audio_cfg.patch_time * self.audio_cfg.hop_size / self.audio_cfg.sample_rate
for i, p in enumerate(patches):
coords = inputs['patch_coord'][i, :len(p['patches']), 1]
valid = inputs['patch_valid'][i, :len(p['patches'])].bool()
columns = coords[valid].unique(sorted=True)
frame = torch.stack([tokens[i, :len(coords)][valid & (coords == t)].mean(0) for t in columns])
frames.append(self._compress_frames(frame))
times.append((columns.float() + .5) * dt)
elif self.kind == 'commercial':
wave = torch.zeros(len(waves), 480000, device=self.device)
for i, w in enumerate(waves):
wave[i, :len(w)] = torch.from_numpy(w).to(self.device)
pooled = self.model.encode_waveform(wave)
if self.tokens is None:
raise RuntimeError('Commercial audio encoder hook failed')
for i, duration in enumerate(durations):
count = min(self.tokens.shape[1], int(np.ceil(duration * 50)))
frames.append(self._compress_frames(self.tokens[i, :count]))
times.append((torch.arange(count, device=self.device) + .5) / 50)
else:
if self.kind == 'gemma':
inputs = self.processor(audio=waves, padding=True, return_tensors='pt')
else:
conversation = [{'role': 'user', 'content': [{'type': 'audio', 'audio': 'unused'}]}]
text = self.processor.apply_chat_template(conversation, tokenize=False, add_generation_prompt=False)
inputs = self.processor(text=[text] * len(waves), audio=waves, padding=True,
sampling_rate=16000, return_tensors='pt')
inputs = {k: v.to(self.device) if torch.is_tensor(v) else v for k, v in inputs.items()}
output = self.model(**inputs, use_cache=False) if self.kind == 'omni' else self.model(**inputs)
hidden = output.last_hidden_state if self.kind == 'gemma' else self.tokens
if hidden is None:
raise RuntimeError('Native hidden state extraction failed')
mask = inputs['attention_mask'].bool()
if self.kind == 'gemma':
pooled = (hidden.float() * mask[:, :, None]).sum(1) / mask.sum(1)[:, None].clamp_min(1)
else:
last = (mask * torch.arange(mask.shape[1], device=self.device)[None]).max(1).values
pooled = hidden[torch.arange(len(waves), device=self.device), last]
config = self.model.config
token = getattr(config, 'audio_token_id', None)
if token is None:
token = getattr(config, 'audio_token_index', None)
if token is None:
raise RuntimeError('Native audio token ID is unavailable')
for i, duration in enumerate(durations):
audio_mask = (inputs['input_ids'][i] == token) & mask[i]
frame = hidden[i, audio_mask]
if not len(frame):
raise RuntimeError('No aligned audio-token features')
frames.append(self._compress_frames(frame))
# One contiguous complete audio input. Native token count
# defines the interval grid; no invented frame resolution.
times.append((torch.arange(len(frame), device=self.device) + .5) * duration / len(frame))
pooled = F.normalize(pooled.float(), dim=-1)
if pooled.shape != (len(waves), self.spec['native_dim']):
raise RuntimeError('Native pooled embedding dimension mismatch')
outputs = []
for i, duration in enumerate(durations):
embedding = pooled[i].float().cpu().numpy()
frame = frames[i].float().cpu().numpy()
time = times[i].float().cpu().numpy()
keep = time < duration
frame, time = frame[keep], time[keep]
if not len(frame) or not np.isfinite(embedding).all() or not np.isfinite(frame).all():
raise RuntimeError('Empty or non-finite frozen features')
outputs.append({'embedding': embedding.astype(np.float16), 'frame_features': frame.astype(np.float32),
'frame_times_s': time.astype(np.float32), 'duration_s': np.float32(duration)})
self.tokens = None
return outputs