File size: 5,121 Bytes
1a0e6e8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 | """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
|