Release best Gemini-tuned Whisper Base and Small with code, normalization and evaluation
cd9b2d8 verified Download training/embedding_probe_study/backbones.py from laion/humaneness-ears-base-medium: direct link, hf CLI and curl.
- Browser
- Download file 9.87 kB
-
https://huggingface.co/laion/humaneness-ears-base-medium/resolve/main/training/embedding_probe_study/backbones.py
- Command line
-
hf download hf://laion/humaneness-ears-base-medium/training/embedding_probe_study/backbones.py
-
curl -L -o backbones.py https://huggingface.co/laion/humaneness-ears-base-medium/resolve/main/training/embedding_probe_study/backbones.py
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])) | |
| 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 | |