"""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) """ # Force transformers to stay PyTorch-only (avoid importing TensorFlow/Flax) 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") # Local import so env vars above take effect before transformers loads from transformers import AutoFeatureExtractor, WavLMModel # Set device if device is None: self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") else: self.device = torch.device(device) # Load feature extractor and model self.feature_extractor = AutoFeatureExtractor.from_pretrained(model_name) self.model = WavLMModel.from_pretrained(model_name) # Freeze model parameters (no fine-tuning) for param in self.model.parameters(): param.requires_grad = False # Move model to device and set to eval mode 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. """ # Ensure correct shape and convert directly to 1D numpy array 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}") # Extract features using feature extractor inputs = self.feature_extractor( waveform_np, sampling_rate=16000, return_tensors="pt", ) # Move inputs to device input_values = inputs.input_values.to(self.device) # Extract features (no gradient computation) 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: # Get the specified hidden state layer: (batch_size, num_frames, hidden_size) hidden_states = outputs.hidden_states[self.layer_idx] # Remove batch dimension and return # Shape: (num_frames, hidden_size) 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