| """WavLM Encoder Module |
| |
| Extracts frame-level features using pretrained WavLM model (frozen). |
| |
| Note: |
| - Some Python environments have an incompatible TensorFlow install (commonly due |
| to NumPy 2.x ABI changes). HuggingFace `transformers` may try to import TF as |
| an optional dependency; we explicitly disable TF/Flax to keep this pipeline |
| PyTorch-only. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import os |
| from typing import Optional |
|
|
| import torch |
|
|
|
|
| class WavLMEncoder: |
| """Extracts features from audio using pretrained WavLM model.""" |
| |
| def __init__( |
| self, |
| model_name: str = "./wavlm-base-plus", |
| device: Optional[str] = None, |
| layer_idx: int = -1 |
| ): |
| """ |
| Initialize WavLM encoder. |
| |
| Args: |
| model_name: HuggingFace model identifier |
| device: Device to run model on ('cuda', 'cpu', or None for auto) |
| layer_idx: Which transformer layer to extract features from (-1 for last layer) |
| """ |
| |
| os.environ.setdefault("TRANSFORMERS_NO_TF", "1") |
| os.environ.setdefault("TRANSFORMERS_NO_FLAX", "1") |
| os.environ.setdefault("USE_TF", "0") |
| os.environ.setdefault("USE_FLAX", "0") |
|
|
| |
| from transformers import AutoFeatureExtractor, WavLMModel |
|
|
| |
| if device is None: |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| else: |
| self.device = torch.device(device) |
| |
| |
| self.feature_extractor = AutoFeatureExtractor.from_pretrained(model_name) |
| self.model = WavLMModel.from_pretrained(model_name) |
| |
| |
| for param in self.model.parameters(): |
| param.requires_grad = False |
| |
| |
| self.model = self.model.to(self.device) |
| self.model.eval() |
| self.layer_idx = layer_idx |
| |
| print(f"WavLM model loaded on {self.device}") |
| |
| def encode(self, waveform: torch.Tensor, extract_layers: Optional[list[int]] = None) -> torch.Tensor | dict[int, torch.Tensor]: |
| """ |
| Extract frame-level features from audio waveform. |
| |
| Args: |
| waveform: Audio waveform tensor, shape (1, num_samples) or (num_samples,) |
| extract_layers: Optional list of layer indices to extract. If provided, returns a dict of layer representations. |
| |
| Returns: |
| If extract_layers is None, returns frame-level features for self.layer_idx, shape (num_frames, hidden_size) |
| If extract_layers is a list, returns a dictionary mapping layer index to feature tensor. |
| """ |
| |
| if waveform.dim() == 1: |
| waveform_np = waveform.detach().cpu().numpy() |
| elif waveform.dim() == 2 and waveform.shape[0] == 1: |
| waveform_np = waveform.squeeze(0).detach().cpu().numpy() |
| else: |
| raise ValueError(f"Expected waveform shape (num_samples,) or (1, num_samples), got {waveform.shape}") |
| |
| |
| inputs = self.feature_extractor( |
| waveform_np, |
| sampling_rate=16000, |
| return_tensors="pt", |
| ) |
| |
| |
| input_values = inputs.input_values.to(self.device) |
| |
| |
| with torch.no_grad(): |
| outputs = self.model(input_values, output_hidden_states=True) |
| |
| if extract_layers is not None: |
| features_dict = {} |
| for idx in extract_layers: |
| hidden_states = outputs.hidden_states[idx] |
| features_dict[idx] = hidden_states.squeeze(0) |
| return features_dict |
| else: |
| |
| hidden_states = outputs.hidden_states[self.layer_idx] |
| |
| |
| |
| features = hidden_states.squeeze(0) |
| |
| return features |
| |
| def encode_batch(self, waveforms: list) -> list: |
| """ |
| Encode multiple waveforms. |
| |
| Args: |
| waveforms: List of waveform tensors |
| |
| Returns: |
| List of feature tensors |
| """ |
| features_list = [] |
| for waveform in waveforms: |
| features = self.encode(waveform) |
| features_list.append(features) |
| return features_list |
| |
| def get_feature_dim(self) -> int: |
| """ |
| Get the dimensionality of the feature vectors. |
| |
| Returns: |
| Feature dimension size |
| """ |
| return self.model.config.hidden_size |
|
|