wavlm-dirosa / wavlm_encoder.py
AutoReXz's picture
Upload project with bundled WavLM model
1a0e6e8 verified
Raw
History Blame Contribute Delete
5.12 kB
"""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