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