Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions padertorch/contrib/mk/modules/activations.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
from torch import Tensor
from torch.nn import GELU as _GELU


class GELU(_GELU):
"""Magnitude-preserving GELU activation function."""
scale: float = 0.653

def __init__(
self, approximate: str = 'none', magnitude_preserving: bool = False
):
super().__init__(approximate=approximate)
self.magnitude_preserving = magnitude_preserving

def forward(self, input: Tensor) -> Tensor:
output = super().forward(input)
if self.magnitude_preserving:
return output / self.scale
return output
28 changes: 10 additions & 18 deletions padertorch/contrib/mk/modules/features/ssl/hubert.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import torchaudio
from transformers.models.hubert.modeling_hubert import HubertModel

from .wav2vec2 import Wav2Vec2
from .wav2vec2 import Wav2Vec2, SAMPLING_RATE

# See https://ieeexplore.ieee.org/abstract/document/9814838, Fig. 2
PR_BASE_LAYER = 11
Expand Down Expand Up @@ -61,13 +61,13 @@ def _init_model(self, model_name):
self.model = HubertModel.from_pretrained(
model_name, cache_dir=self.cache_dir
).to(self.device)
self.sampling_rate = 16_000
self.sampling_rate = SAMPLING_RATE
elif self.backend == "torchaudio":
bundle = getattr(torchaudio.pipelines, model_name)
self.model = bundle.get_model().to(self.device)
self.sampling_rate = bundle.sample_rate
if self.layer == -1:
self.layer = len(self.model.encoder.transformer.layers)-1
self.layer = self.num_layers
else:
raise ValueError(f'Unknown backend: {self.backend}')

Expand All @@ -80,20 +80,18 @@ def extract_features_from_latents(
if isinstance(sequence_lengths, np.ndarray):
sequence_lengths = torch.from_numpy(sequence_lengths).long()\
.to(latents.device)

with self.context:
if self.backend == "torchaudio":
num_layers = None if isinstance(self.layer, str) else self.layer
x = self.model.encoder.extract_features(
latents,
lengths=sequence_lengths,
num_layers=self.layer,
num_layers=num_layers,
)
if isinstance(self.layer, int):
x = x[-1]
elif self.layer is None:
return x
else:
raise NotImplementedError(self.layer)
return x
return self.extract_layer(x)

# hf backend
hidden_states = self.model.feature_projection(latents)
hidden_states = self.model._mask_hidden_states(hidden_states)
encoder_outputs = self.model.encoder(
Expand All @@ -102,13 +100,7 @@ def extract_features_from_latents(
return_dict=True,
)
x = encoder_outputs.hidden_states
if isinstance(self.layer, int):
x = x[self.layer]
elif self.layer is None:
return x
else:
raise NotImplementedError(self.layer)
return x
return self.extract_layer(x)

def forward(
self,
Expand Down
Original file line number Diff line number Diff line change
@@ -1 +1 @@
from ._wav2vec2 import Wav2Vec2
from ._wav2vec2 import Wav2Vec2, SAMPLING_RATE
134 changes: 93 additions & 41 deletions padertorch/contrib/mk/modules/features/ssl/wav2vec2/_wav2vec2.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,8 @@
STFT, stft_frame_index_to_sample_index
)
import padertorch as pt
from padertorch.contrib.je.modules.conv_utils import compute_conv_output_sequence_lengths
from padertorch.contrib.je.modules.conv_utils import\
compute_conv_output_sequence_lengths
from padertorch.contrib.mk.typing import TSeqLen
from padertorch.contrib.mk.utils import compute_receptive_field_1d
from padertorch.ops.sequence.mask import compute_mask
Expand All @@ -22,6 +23,8 @@
import torchaudio
from transformers.models.wav2vec2.modeling_wav2vec2 import Wav2Vec2Model

SAMPLING_RATE = 16_000


def tuple_to_int(sequence) -> list:
return list(map(lambda t: t[0], sequence))
Expand Down Expand Up @@ -50,6 +53,26 @@ class Wav2Vec2(pt.Module):
loaded from huggingface.co. Defaults to "torchaudio".
pad (bool):
fading (str, bool, optional):

>>> wav2vec2 = Wav2Vec2()
>>> signal = torch.zeros((2, 99919))
>>> num_samples = [99919, 99840]
>>> x, seq_len_x = wav2vec2(signal, num_samples)
>>> x.shape
torch.Size([2, 313, 768])
>>> seq_len_x
[313, 312]

>>> wav2vec = Wav2Vec2(layer=None)
>>> x, seq_len_x = wav2vec(signal, num_samples)
>>> len(x)
12

>>> wav2vec2 = Wav2Vec2(layer=13)
>>> x, seq_len_x = wav2vec2(signal, num_samples)
Traceback (most recent call last):
...
ValueError: `num_layers` must be between [1, 12]
"""
def __init__(
self,
Expand Down Expand Up @@ -109,6 +132,24 @@ def __init__(
pad=self.pad, fading=self.fading,
)

@property
def d_model(self):
if self.backend == "torchaudio":
return self.model.encoder.transformer.layers[-1]\
.feed_forward.output_dense.out_features
if self.backend == "hf":
return self.model.encoder.layers[-1].feed_forward.output_dense\
.out_features
raise ValueError(f'Unknown backend: {self.backend}')

@property
def num_layers(self):
if self.backend == "torchaudio":
return len(self.model.encoder.transformer.layers)
if self.backend == "hf":
return len(self.model.encoder.layers)
raise ValueError(f'Unknown backend: {self.backend}')

@property
def frame_rate(self):
return int(self.sampling_rate / self.downsample_factor)
Expand All @@ -126,6 +167,9 @@ def context(self):
return torch.no_grad()
return nullcontext()

def __getitem__(self, item):
return self.layers[item]

def _init_model(self, model_name):
if "wav2vec2" not in model_name.lower():
raise ValueError(
Expand All @@ -140,12 +184,13 @@ def _init_model(self, model_name):
self.model = Wav2Vec2Model.from_pretrained(
model_name, cache_dir=self.cache_dir, from_tf=False,
).to(self.device)
self.sampling_rate = SAMPLING_RATE
elif self.backend == "torchaudio":
bundle = getattr(torchaudio.pipelines, model_name)
self.model = bundle.get_model().to(self.device)
self.sampling_rate = bundle.sample_rate
if self.layer == -1:
self.layer = len(self.model.encoder.transformer.layers)-1
self.layer = self.num_layers
if self.attention_fn is not None:
for layer in self.model.encoder.transformer.layers:
named_params = dict(layer.attention.named_parameters())
Expand Down Expand Up @@ -173,9 +218,10 @@ def _get_conv_params(self):
))

def _forward(self, time_signal: Tensor, sequence_lengths: TSeqLen):
num_layers = None if isinstance(self.layer, str) else self.layer
x, _ = self.model.extract_features(
time_signal, lengths=sequence_lengths,
num_layers=self.layer,
num_layers=num_layers,
)
return x

Expand All @@ -188,7 +234,7 @@ def _check_shape(self, x: Tensor, sequence_lengths: TSeqLen):
f"Output shape: {x.shape}\n"
f"Expected sequence lengths: {sequence_lengths}\n"
f"Padded sequence lengths: {sequence_lengths}\n"
"Setting fading=half will likely fix this issue."
"Setting pad=True or fading=half will likely fix this issue."
)

def remove_weight_norm(self):
Expand Down Expand Up @@ -374,27 +420,50 @@ def to_frames(self, samples, num_frames):
last_frame_index = frame_index
return y

def extract_layer(
self, hidden: tp.Union[Tensor, tp.List[Tensor]],
sequence_lengths: TSeqLen = None,
):
if sequence_lengths is not None:
hidden = [hi[..., :sequence_lengths.max(), :] for hi in hidden]
if self.backend == "torchaudio":
if isinstance(self.layer, int):
return hidden[-1]
if self.layer is None:
return hidden
raise NotImplementedError(self.layer)

if isinstance(self.layer, int):
try:
hidden = hidden[self.layer]
except IndexError as exc:
raise ValueError(
f"`layer` must be between [1, {self.num_layers}]"
) from exc
return hidden
if self.layer is None:
return hidden[1:] # Drop input of first Transformer layer
raise NotImplementedError(self.layer)

def extract_features_from_latents(
self, latents: Tensor, sequence_lengths: TSeqLen
):
self.maybe_eval()
if isinstance(sequence_lengths, np.ndarray):
sequence_lengths = torch.from_numpy(sequence_lengths).long()\
.to(latents.device)

with self.context:
if self.backend == "torchaudio":
num_layers = None if isinstance(self.layer, str) else self.layer
x = self.model.encoder.extract_features(
latents,
lengths=sequence_lengths,
num_layers=self.layer,
num_layers=num_layers,
)
if isinstance(self.layer, int):
x = x[-1]
elif self.layer is None:
return x
else:
raise NotImplementedError(self.layer)
return x
return self.extract_layer(x)

# hf backend
hidden_states, latents = self.model.feature_projection(
latents
)
Expand All @@ -404,13 +473,7 @@ def extract_features_from_latents(
return_dict=True,
)
x = encoder_outputs.hidden_states
if isinstance(self.layer, int):
x = x[self.layer]
elif self.layer is None:
return x
else:
raise NotImplementedError(self.layer)
return x
return self.extract_layer(x)

def forward(
self,
Expand Down Expand Up @@ -483,16 +546,12 @@ def forward(
)
if self.detach:
x = list(map(torch.detach, x))
if isinstance(self.layer, int):
x = x[-1]
if out_sequence_lengths is not None:
x = x[..., :out_sequence_lengths.max(), :]
return x, out_sequence_lengths
if self.layer is None:
x = [xi[..., :out_sequence_lengths.max(), :] for xi in x]
return x, out_sequence_lengths
raise NotImplementedError(self.layer)
return (
self.extract_layer(x, out_sequence_lengths),
out_sequence_lengths
)

# hf backend
self.model: Wav2Vec2Model
out_sequence_lengths = self.compute_output_lengths(sequence_lengths)
z = self.model.feature_extractor(time_signal.float())\
Expand All @@ -508,16 +567,9 @@ def forward(
output_hidden_states=True,
return_dict=True,
)
if isinstance(self.layer, int):
x = outputs.hidden_states[self.layer]
if self.detach:
x = x.detach()
elif self.layer is None:
x = outputs.hidden_states
if self.detach:
x = [h.detach() for h in x]
return x, out_sequence_lengths
else:
raise ValueError(f'Unknown layer: {self.layer}')

return x, out_sequence_lengths
if self.detach:
x = [h.detach() for h in x]
return (
self.extract_layer(outputs.hidden_states, out_sequence_lengths),
out_sequence_lengths
)
Loading