vocoder-small-speaker-encoder / speaker_embedding_architecture.py
mlr2000's picture
Upload folder using huggingface_hub
e8c522a verified
Raw
History Blame Contribute Delete
10.5 kB
from dataclasses import dataclass
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import Wav2Vec2Model, PreTrainedModel
from transformers.modeling_outputs import ModelOutput
from .ge2e_loss import GE2ELoss
from .speaker_embedding_config import SpeakerEmbeddingConfig
@dataclass
class EmbeddingOutput(ModelOutput):
loss: Optional[torch.FloatTensor] = None
embeddings: torch.Tensor = None # shape (batch_size, embedding_size)
dvecs: Optional[torch.Tensor] = None # shape (N, M, embedding_size) for analysis
# https://arxiv.org/pdf/1803.10963
class AttentiveStatisticsPooling(nn.Module):
def __init__(self, input_dim, hidden_dim=128):
super().__init__()
self.attention = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.Tanh(),
nn.Linear(hidden_dim, 1)
)
def forward(self, x, mask=None):
# x: (batch, time, features)
# Compute attention weights
attn_weights = self.attention(x) # (batch, time, 1)
if mask is not None:
# mask: (batch, time), True = valid, False = padding
attn_weights = attn_weights.masked_fill(~mask.unsqueeze(-1), float('-inf'))
attn_weights = F.softmax(attn_weights, dim=1) # (batch, time, 1)
# Weighted mean
mean = torch.sum(x * attn_weights, dim=1) # (batch, features)
# Weighted std
variance = torch.sum(attn_weights * (x - mean.unsqueeze(1)) ** 2, dim=1)
std = torch.sqrt(variance.clamp(min=1e-5))
# Concatenate mean and std
return torch.cat([mean, std], dim=-1) # (batch, features * 2)
class ZeroModule(nn.Module):
def forward(self, hidden_states):
return torch.zeros_like(hidden_states)
class SpeakerEmbeddingArchitecture(PreTrainedModel):
config_class = SpeakerEmbeddingConfig
def __init__(self, config):
super().__init__(config)
self.config = config
self.encoder = Wav2Vec2Model(config)
self.encoder.init_weights()
if getattr(config, 'disable_positional_embeddings', False):
self.encoder.encoder.pos_conv_embed = ZeroModule()
# Projection layer to get desired embedding dimension
self.pooling = AttentiveStatisticsPooling(config.hidden_size, hidden_dim=config.hidden_size)
self.embedding_size = getattr(config, 'embedding_size', 256)
n_projection_layers = getattr(config, 'n_projection_layers', 1)
self.projection_layers = nn.ModuleList([])
for l in range(n_projection_layers - 1):
self.projection_layers.append(
nn.Sequential(
nn.Linear(
config.hidden_size * 2 if l == 0 else self.embedding_size,
self.embedding_size
),
#nn.BatchNorm1d(self.embedding_size, momentum=0.01, eps=1e-5),
nn.LeakyReLU(0.01)
)
)
self.final_projection = nn.Sequential(
nn.Linear(
config.hidden_size * 2 if n_projection_layers == 1 else self.embedding_size,
self.embedding_size
),
#nn.BatchNorm1d(self.embedding_size, momentum=0.01, eps=1e-5),
nn.LeakyReLU(0.01)
)
self.layer_weights = None
if getattr(config, 'use_layer_weights', False):
num_layers = config.num_hidden_layers + 1
self.layer_weights = nn.Parameter(torch.ones(num_layers), requires_grad=True)
# Optional: Layer normalization before pooling
self.layer_norm = nn.LayerNorm(config.hidden_size)
self.margin = config.loss_margin
self.scale = config.loss_scale
for m in self.pooling.attention.modules():
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight, gain=0.2)
if m.bias is not None:
nn.init.zeros_(m.bias)
# Normal init for projection layers
self.projection_layers.apply(self._init_weights)
self.final_projection.apply(self._init_weights)
# GE2E Loss
# self.ge2e_loss = GE2ELoss(loss_method='softmax')
def _init_weights(self, module):
"""Initialize the weights"""
if isinstance(module, nn.Linear):
nn.init.xavier_uniform_(module.weight, gain=0.5)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.BatchNorm1d):
nn.init.ones_(module.weight)
nn.init.zeros_(module.bias)
def temporal_pooling(self, hidden_states, attention_mask=None):
"""
Pool over time dimension with optional attention mask.
Args:
hidden_states: (batch_size, seq_len, hidden_size)
attention_mask: (batch_size, seq_len) - optional
Returns:
pooled: (batch_size, hidden_size)
"""
if attention_mask is not None:
# Mask padding tokens before pooling
# attention_mask shape: (batch_size, seq_len)
mask_expanded = attention_mask.unsqueeze(-1).expand(hidden_states.size())
sum_hidden = torch.sum(hidden_states * mask_expanded, dim=1)
sum_mask = torch.clamp(mask_expanded.sum(dim=1), min=1e-9)
pooled = sum_hidden / sum_mask
else:
# Simple mean pooling
pooled = torch.mean(hidden_states, dim=1)
return pooled
def forward(
self,
input_values: Optional[torch.Tensor],
attention_mask: Optional[torch.Tensor] = None,
labels: Optional[torch.Tensor] = None,
mask_time_indices: Optional[torch.FloatTensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
cached_centroids: Optional[torch.Tensor] = None,
cached_labels: Optional[torch.Tensor] = None,
):
"""
Args:
input_values: (N*M, seq_len) audio input
attention_mask: (N*M, seq_len) optional
compute_loss: whether to compute GE2E loss (needs N and M)
speakers_per_batch: N - number of speakers in batch
utterances_per_speaker: M - utterances per speaker
Returns:
EmbeddingOutput with loss and embeddings
"""
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
# 1. Encode audio
encoder_outputs = self.encoder(
input_values=input_values,
attention_mask=attention_mask,
mask_time_indices=mask_time_indices,
output_attentions=output_attentions,
output_hidden_states=True,
return_dict=return_dict,
)
if self.layer_weights is None:
hidden_states = encoder_outputs.last_hidden_state
else:
# Weighted sum of hidden states from all layers
hidden_states = torch.stack(encoder_outputs.hidden_states, dim=0)
layer_weights = F.softmax(self.layer_weights, dim=0)
hidden_states = (hidden_states * layer_weights.view(-1, 1, 1, 1)).sum(dim=0)
# hidden_states = encoder_outputs.last_hidden_state # (N*M, seq_len, hidden_size)
sub_attention_mask = None
if attention_mask is not None:
# Use the built-in Wav2Vec2 function!
sub_attention_mask = self.encoder._get_feature_vector_attention_mask(
feature_vector_length=hidden_states.shape[1],
attention_mask=attention_mask,
add_adapter=False
)
# 2. Optional normalization
hidden_states = self.layer_norm(hidden_states)
# 3. Temporal pooling
pooled = self.pooling(hidden_states, mask=sub_attention_mask)
# pooled = self.temporal_pooling(hidden_states, sub_attention_mask) # (N*M, hidden_size)
# 4. Project to embedding space
# embeddings = self.projection(pooled) # (N*M, embedding_size)
embeddings = pooled
for projection_layer in self.projection_layers:
embeddings = projection_layer(embeddings)
embeddings = self.final_projection(embeddings)
# 5. L2 normalize embeddings (important for GE2E and cosine similarity!)
embeddings = nn.functional.normalize(embeddings, p=2, dim=1, eps=1e-8)
if cached_centroids is not None:
# We are inside the DDP Forward pass now!
# DDP will see 'w' and 'b' being used here.
# w = self.ge2e_loss.w
# b = self.ge2e_loss.b
# Calculate Similarity against CACHED centroids
# embeddings: (Batch, D), cached_centroids: (N, D)
# sim_matrix = torch.mm(embeddings, cached_centroids.transpose(0, 1))
sim_matrix = torch.einsum('bd,bnd->bn', embeddings, cached_centroids)
# sim_matrix = torch.clamp(sim_matrix, min=1e-6)
# sim_matrix = w * sim_matrix + b
# https://arxiv.org/pdf/1801.05599 this loss here
one_hot = torch.zeros_like(sim_matrix)
one_hot.scatter_(1, cached_labels.unsqueeze(1), 1.0)
sim_matrix = sim_matrix - one_hot * self.margin
sim_matrix = sim_matrix * self.scale
if cached_labels is not None:
# loss_fct = nn.CrossEntropyLoss()
loss = F.cross_entropy(sim_matrix, cached_labels)
# Return the Loss directly
return EmbeddingOutput(loss=loss, embeddings=embeddings)
else:
return EmbeddingOutput(loss=None, embeddings=embeddings)
return EmbeddingOutput(loss=None, embeddings=embeddings)
def get_speaker_embedding(self, audio, attention_mask=None):
"""
Convenience method to get embedding for a single audio sample.
Args:
audio: (1, seq_len) or (seq_len,)
attention_mask: optional
Returns:
embedding: (embedding_size,)
"""
if audio.dim() == 1:
audio = audio.unsqueeze(0)
output = self.forward(
input_values=audio,
attention_mask=attention_mask,
compute_loss=False,
return_dict=True
)
return output.embeddings.squeeze(0)