| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796 |
- """
- State-of-the-art architectures for voice recognition.
- This module implements ECAPA-TDNN, TitaNet-S, and SpeakerNet-M architectures
- optimized for speaker recognition and voice identification tasks.
- """
- import math
- from typing import Dict, List, Optional, Tuple, Union
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- class SEBlock(nn.Module):
- """Squeeze-and-Excitation block for channel attention."""
-
- def __init__(self, channels: int, reduction: int = 16):
- super().__init__()
- self.squeeze = nn.AdaptiveAvgPool1d(1)
- self.excitation = nn.Sequential(
- nn.Linear(channels, channels // reduction, bias=False),
- nn.ReLU(inplace=True),
- nn.Linear(channels // reduction, channels, bias=False),
- nn.Sigmoid()
- )
-
- def forward(self, x):
- b, c, _ = x.size()
- y = self.squeeze(x).view(b, c)
- y = self.excitation(y).view(b, c, 1)
- return x * y.expand_as(x)
- class Res2Conv1d(nn.Module):
- """
- Res2Conv1d implementation for multi-scale feature extraction.
-
- Based on Res2Net architecture adapted for 1D convolutions.
- """
-
- def __init__(self, in_channels: int, out_channels: int, kernel_size: int = 3,
- dilation: int = 1, scale: int = 4):
- super().__init__()
-
- assert out_channels % scale == 0, "out_channels must be divisible by scale"
-
- self.scale = scale
- self.out_channels = out_channels
- width = out_channels // scale
-
- self.convs = nn.ModuleList()
- for i in range(scale - 1):
- self.convs.append(
- nn.Conv1d(width, width, kernel_size, padding=kernel_size//2 * dilation,
- dilation=dilation, bias=False)
- )
-
- self.bn = nn.BatchNorm1d(out_channels)
- self.relu = nn.ReLU(inplace=True)
-
- def forward(self, x):
- batch_size = x.size(0)
-
- # Split input into scale parts
- spx = torch.split(x, self.out_channels // self.scale, dim=1)
-
- outputs = [spx[0]]
- for i in range(1, self.scale):
- if i == 1:
- sp = spx[i]
- else:
- sp = spx[i] + outputs[-1]
- sp = self.convs[i-1](sp)
- outputs.append(sp)
-
- out = torch.cat(outputs, dim=1)
- out = self.bn(out)
- out = self.relu(out)
-
- return out
- class TDNN(nn.Module):
- """Time Delay Neural Network layer."""
-
- def __init__(self, in_channels: int, out_channels: int, kernel_size: int = 5,
- dilation: int = 1, use_bn: bool = True, use_relu: bool = True):
- super().__init__()
-
- self.conv = nn.Conv1d(
- in_channels, out_channels, kernel_size,
- padding=kernel_size//2 * dilation, dilation=dilation, bias=not use_bn
- )
-
- self.bn = nn.BatchNorm1d(out_channels) if use_bn else None
- self.relu = nn.ReLU(inplace=True) if use_relu else None
-
- def forward(self, x):
- x = self.conv(x)
- if self.bn is not None:
- x = self.bn(x)
- if self.relu is not None:
- x = self.relu(x)
- return x
- class StatisticalPooling(nn.Module):
- """Statistical pooling layer that computes mean and standard deviation."""
-
- def __init__(self, input_dim: int, output_dim: Optional[int] = None):
- super().__init__()
- self.input_dim = input_dim
- self.output_dim = output_dim or input_dim * 2
-
- if output_dim and output_dim != input_dim * 2:
- self.projection = nn.Linear(input_dim * 2, output_dim)
- else:
- self.projection = None
-
- def forward(self, x):
- # x shape: (batch, channels, time)
- mean = torch.mean(x, dim=2)
- # Add epsilon for numerical stability to prevent NaN
- std = torch.std(x, dim=2) + 1e-8
-
- # Concatenate mean and std
- stats = torch.cat([mean, std], dim=1)
-
- if self.projection is not None:
- stats = self.projection(stats)
-
- return stats
- class AttentiveStatisticalPooling(nn.Module):
- """Attentive statistical pooling with learnable attention weights."""
-
- def __init__(self, input_dim: int, attention_dim: int = 128):
- super().__init__()
- self.input_dim = input_dim
- self.attention_dim = attention_dim
-
- self.attention = nn.Sequential(
- nn.Linear(input_dim, attention_dim),
- nn.Tanh(),
- nn.Linear(attention_dim, 1)
- )
-
- def forward(self, x):
- # x shape: (batch, channels, time)
- batch_size, channels, time_steps = x.shape
-
- # Transpose for attention computation
- x_t = x.transpose(1, 2) # (batch, time, channels)
-
- # Compute attention weights
- attention_weights = self.attention(x_t) # (batch, time, 1)
- attention_weights = F.softmax(attention_weights, dim=1)
-
- # Apply attention weights
- weighted_x = x_t * attention_weights # (batch, time, channels)
-
- # Compute weighted statistics
- mean = torch.sum(weighted_x, dim=1) # (batch, channels)
-
- # Compute weighted standard deviation with numerical stability
- diff = x_t - mean.unsqueeze(1)
- weighted_var = torch.sum(attention_weights * diff ** 2, dim=1)
- # Clamp variance to prevent negative values and add epsilon
- weighted_var = torch.clamp(weighted_var, min=1e-8)
- std = torch.sqrt(weighted_var)
-
- # Concatenate mean and std
- stats = torch.cat([mean, std], dim=1)
-
- return stats
- class ECAPA_TDNN(nn.Module):
- """
- ECAPA-TDNN: Emphasized Channel Attention, Propagation and Aggregation
- in TDNN based Speaker Verification.
-
- Reference: "ECAPA-TDNN: Emphasized Channel Attention, Propagation and
- Aggregation in TDNN based Speaker Verification" by Desplanques et al.
- """
-
- def __init__(self, input_dim: int = 40, channels: int = 512,
- embedding_dim: int = 192, num_speakers: int = None,
- use_attention_pooling: bool = True):
- """
- Initialize ECAPA-TDNN model.
-
- Args:
- input_dim: Input feature dimension (e.g., 40 for mel spectrograms)
- channels: Number of channels in TDNN layers
- embedding_dim: Dimension of speaker embeddings
- num_speakers: Number of speakers for classification (None for embeddings only)
- use_attention_pooling: Whether to use attentive statistical pooling
- """
- super().__init__()
-
- self.input_dim = input_dim
- self.channels = channels
- self.embedding_dim = embedding_dim
- self.num_speakers = num_speakers
-
- # Frame-level feature extraction
- self.frame_conv = nn.Conv1d(input_dim, channels, 5, padding=2)
-
- # ECAPA blocks
- self.ecapa1 = self._make_ecapa_block(channels, channels, kernel_size=3, dilation=2)
- self.ecapa2 = self._make_ecapa_block(channels, channels, kernel_size=3, dilation=3)
- self.ecapa3 = self._make_ecapa_block(channels, channels, kernel_size=3, dilation=4)
- self.ecapa4 = self._make_ecapa_block(channels, 3 * channels, kernel_size=1, dilation=1)
-
- # Aggregation
- self.aggregation_conv = nn.Conv1d(3 * channels, 1536, 1)
-
- # Pooling
- if use_attention_pooling:
- self.pooling = AttentiveStatisticalPooling(1536)
- pooling_output_dim = 1536 * 2
- else:
- self.pooling = StatisticalPooling(1536)
- pooling_output_dim = 1536 * 2
-
- # Batch normalization after pooling
- self.bn_pooling = nn.BatchNorm1d(pooling_output_dim)
-
- # Embedding layer
- self.embedding = nn.Linear(pooling_output_dim, embedding_dim)
- self.bn_embedding = nn.BatchNorm1d(embedding_dim)
-
- # Classification layer (optional)
- if num_speakers is not None:
- self.classifier = nn.Linear(embedding_dim, num_speakers)
- else:
- self.classifier = None
-
- self._initialize_weights()
-
- def _make_ecapa_block(self, in_channels: int, out_channels: int,
- kernel_size: int = 3, dilation: int = 1):
- """Create an ECAPA block with Res2Conv and SE."""
- return nn.Sequential(
- Res2Conv1d(in_channels, out_channels, kernel_size, dilation, scale=8),
- SEBlock(out_channels, reduction=16)
- )
-
- def forward(self, x, return_embedding: bool = False):
- """
- Forward pass through ECAPA-TDNN.
-
- Args:
- x: Input tensor (batch, features, time)
- return_embedding: Whether to return embeddings instead of classification
-
- Returns:
- If num_speakers is None or return_embedding is True: speaker embeddings
- Otherwise: classification logits
- """
- # Frame-level processing
- x = self.frame_conv(x)
-
- # ECAPA blocks with residual connections
- x1 = self.ecapa1(x)
- x = x + x1
-
- x2 = self.ecapa2(x)
- x = x + x2
-
- x3 = self.ecapa3(x)
- x = x + x3
-
- # Aggregation (concatenate all frame-level features)
- x_agg = torch.cat([x1, x2, x3], dim=1)
- x = self.ecapa4(x_agg)
-
- x = self.aggregation_conv(x)
-
- # Statistical pooling
- x = self.pooling(x)
- x = self.bn_pooling(x)
-
- # Embedding
- embedding = self.embedding(x)
- embedding = self.bn_embedding(embedding)
-
- if self.classifier is None or return_embedding:
- return F.normalize(embedding, p=2, dim=1)
- else:
- logits = self.classifier(embedding)
- return logits
-
- def get_embeddings(self, x):
- """Get speaker embeddings."""
- return self.forward(x, return_embedding=True)
-
- def _initialize_weights(self):
- """Initialize model weights."""
- for m in self.modules():
- if isinstance(m, nn.Conv1d):
- nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
- if m.bias is not None:
- nn.init.constant_(m.bias, 0)
- elif isinstance(m, nn.BatchNorm1d):
- nn.init.constant_(m.weight, 1)
- nn.init.constant_(m.bias, 0)
- elif isinstance(m, nn.Linear):
- nn.init.normal_(m.weight, 0, 0.01)
- if m.bias is not None:
- nn.init.constant_(m.bias, 0)
- class TitaNet_S(nn.Module):
- """
- TitaNet-S: A lightweight variant of TitaNet for speaker recognition.
-
- Based on the TitaNet architecture but with reduced parameters
- for efficient speaker recognition.
- """
-
- def __init__(self, input_dim: int = 40, channels: List[int] = None,
- embedding_dim: int = 192, num_speakers: int = None,
- dropout_rate: float = 0.1):
- """
- Initialize TitaNet-S model.
-
- Args:
- input_dim: Input feature dimension
- channels: List of channel dimensions for each block
- embedding_dim: Dimension of speaker embeddings
- num_speakers: Number of speakers for classification
- dropout_rate: Dropout rate for regularization
- """
- super().__init__()
-
- if channels is None:
- channels = [128, 256, 512, 1024]
-
- self.input_dim = input_dim
- self.channels = channels
- self.embedding_dim = embedding_dim
- self.num_speakers = num_speakers
-
- # Input projection
- self.input_conv = nn.Conv1d(input_dim, channels[0], 1)
- self.input_bn = nn.BatchNorm1d(channels[0])
-
- # TitaNet blocks
- self.blocks = nn.ModuleList()
- for i in range(len(channels) - 1):
- block = self._make_titanet_block(
- channels[i], channels[i+1], dropout_rate
- )
- self.blocks.append(block)
-
- # Statistical pooling
- self.pooling = AttentiveStatisticalPooling(channels[-1])
-
- # Embedding layers
- pooling_dim = channels[-1] * 2
- self.embedding_layers = nn.Sequential(
- nn.Linear(pooling_dim, channels[-1]),
- nn.BatchNorm1d(channels[-1]),
- nn.ReLU(inplace=True),
- nn.Dropout(dropout_rate),
- nn.Linear(channels[-1], embedding_dim),
- nn.BatchNorm1d(embedding_dim)
- )
-
- # Classification layer (optional)
- if num_speakers is not None:
- self.classifier = nn.Linear(embedding_dim, num_speakers)
- else:
- self.classifier = None
-
- self._initialize_weights()
-
- def _make_titanet_block(self, in_channels: int, out_channels: int,
- dropout_rate: float):
- """Create a TitaNet block."""
- return nn.Sequential(
- # Depthwise separable convolution
- nn.Conv1d(in_channels, in_channels, 3, padding=1, groups=in_channels),
- nn.Conv1d(in_channels, out_channels, 1),
- nn.BatchNorm1d(out_channels),
- nn.ReLU(inplace=True),
- nn.Dropout1d(dropout_rate),
-
- # Squeeze-and-Excitation
- SEBlock(out_channels, reduction=8),
-
- # Residual connection (if dimensions match)
- # Note: Actual residual connection would need dimension matching
- )
-
- def forward(self, x, return_embedding: bool = False):
- """Forward pass through TitaNet-S."""
- # Input projection
- x = self.input_conv(x)
- x = self.input_bn(x)
- x = F.relu(x, inplace=True)
-
- # TitaNet blocks
- for block in self.blocks:
- x = block(x)
-
- # Statistical pooling
- x = self.pooling(x)
-
- # Embedding
- embedding = self.embedding_layers(x)
-
- if self.classifier is None or return_embedding:
- return F.normalize(embedding, p=2, dim=1)
- else:
- logits = self.classifier(embedding)
- return logits
-
- def get_embeddings(self, x):
- """Get speaker embeddings."""
- return self.forward(x, return_embedding=True)
-
- def _initialize_weights(self):
- """Initialize model weights."""
- for m in self.modules():
- if isinstance(m, nn.Conv1d):
- nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
- if m.bias is not None:
- nn.init.constant_(m.bias, 0)
- elif isinstance(m, nn.BatchNorm1d):
- nn.init.constant_(m.weight, 1)
- nn.init.constant_(m.bias, 0)
- elif isinstance(m, nn.Linear):
- nn.init.normal_(m.weight, 0, 0.01)
- if m.bias is not None:
- nn.init.constant_(m.bias, 0)
- class SpeakerNet_M(nn.Module):
- """
- SpeakerNet-M: Medium-sized speaker recognition network.
-
- A balanced architecture for speaker recognition that provides
- good performance with moderate computational requirements.
- """
-
- def __init__(self, input_dim: int = 40, hidden_dim: int = 512,
- embedding_dim: int = 256, num_speakers: int = None,
- num_layers: int = 4, dropout_rate: float = 0.1):
- """
- Initialize SpeakerNet-M model.
-
- Args:
- input_dim: Input feature dimension
- hidden_dim: Hidden dimension for TDNN layers
- embedding_dim: Dimension of speaker embeddings
- num_speakers: Number of speakers for classification
- num_layers: Number of TDNN layers
- dropout_rate: Dropout rate for regularization
- """
- super().__init__()
-
- self.input_dim = input_dim
- self.hidden_dim = hidden_dim
- self.embedding_dim = embedding_dim
- self.num_speakers = num_speakers
- self.num_layers = num_layers
-
- # Input layer
- self.input_layer = TDNN(input_dim, hidden_dim, kernel_size=5)
-
- # TDNN layers with different dilations
- self.tdnn_layers = nn.ModuleList()
- dilations = [1, 2, 3, 4, 5][:num_layers]
-
- for i, dilation in enumerate(dilations):
- layer = TDNN(
- hidden_dim, hidden_dim,
- kernel_size=3, dilation=dilation
- )
- self.tdnn_layers.append(layer)
-
- # Context expansion
- self.context_conv = nn.Conv1d(hidden_dim, hidden_dim * 2, 1)
- self.context_bn = nn.BatchNorm1d(hidden_dim * 2)
-
- # Statistical pooling
- self.pooling = StatisticalPooling(hidden_dim * 2)
-
- # Embedding layers
- pooling_dim = hidden_dim * 4 # 2x for mean+std
- self.embedding_layers = nn.Sequential(
- nn.Linear(pooling_dim, hidden_dim),
- nn.BatchNorm1d(hidden_dim),
- nn.ReLU(inplace=True),
- nn.Dropout(dropout_rate),
- nn.Linear(hidden_dim, embedding_dim),
- nn.BatchNorm1d(embedding_dim)
- )
-
- # Classification layer (optional)
- if num_speakers is not None:
- self.classifier = nn.Sequential(
- nn.Dropout(dropout_rate),
- nn.Linear(embedding_dim, num_speakers)
- )
- else:
- self.classifier = None
-
- self._initialize_weights()
-
- def forward(self, x, return_embedding: bool = False):
- """Forward pass through SpeakerNet-M."""
- # Input processing
- x = self.input_layer(x)
-
- # TDNN layers with residual connections
- for layer in self.tdnn_layers:
- residual = x
- x = layer(x)
- x = x + residual # Residual connection
-
- # Context expansion
- x = self.context_conv(x)
- x = self.context_bn(x)
- x = F.relu(x, inplace=True)
-
- # Statistical pooling
- x = self.pooling(x)
-
- # Embedding
- embedding = self.embedding_layers(x)
-
- if self.classifier is None or return_embedding:
- return F.normalize(embedding, p=2, dim=1)
- else:
- logits = self.classifier(embedding)
- return logits
-
- def get_embeddings(self, x):
- """Get speaker embeddings."""
- return self.forward(x, return_embedding=True)
-
- def _initialize_weights(self):
- """Initialize model weights."""
- for m in self.modules():
- if isinstance(m, nn.Conv1d):
- nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
- if m.bias is not None:
- nn.init.constant_(m.bias, 0)
- elif isinstance(m, nn.BatchNorm1d):
- nn.init.constant_(m.weight, 1)
- nn.init.constant_(m.bias, 0)
- elif isinstance(m, nn.Linear):
- nn.init.normal_(m.weight, 0, 0.01)
- if m.bias is not None:
- nn.init.constant_(m.bias, 0)
- def create_voice_recognition_model(model_type: str = "ecapa_tdnn", **kwargs) -> nn.Module:
- """
- Factory function to create voice recognition models.
-
- Args:
- model_type: Type of model ('ecapa_tdnn', 'titanet_s', 'speakernet_m')
- **kwargs: Additional arguments for model initialization
-
- Returns:
- Voice recognition model instance
- """
- if model_type.lower() == "ecapa_tdnn":
- return ECAPA_TDNN(**kwargs)
- elif model_type.lower() == "titanet_s":
- return TitaNet_S(**kwargs)
- elif model_type.lower() == "speakernet_m":
- return SpeakerNet_M(**kwargs)
- else:
- raise ValueError(f"Unknown model type: {model_type}")
- def count_parameters(model: nn.Module) -> Dict[str, int]:
- """Count model parameters."""
- total_params = sum(p.numel() for p in model.parameters())
- trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
-
- return {
- "total_parameters": total_params,
- "trainable_parameters": trainable_params,
- "non_trainable_parameters": total_params - trainable_params
- }
- class AngularMarginLoss(nn.Module):
- """
- Angular Margin Loss (ArcFace) for speaker recognition.
-
- Reference: "ArcFace: Additive Angular Margin Loss for Deep Face Recognition"
- """
-
- def __init__(self, embedding_dim: int, num_speakers: int,
- margin: float = 0.5, scale: float = 64.0):
- """
- Initialize Angular Margin Loss.
-
- Args:
- embedding_dim: Dimension of input embeddings
- num_speakers: Number of speaker classes
- margin: Angular margin (m)
- scale: Scale factor (s)
- """
- super().__init__()
- self.embedding_dim = embedding_dim
- self.num_speakers = num_speakers
- self.margin = margin
- self.scale = scale
-
- # Weight matrix (represents speaker centroids)
- self.weight = nn.Parameter(torch.randn(num_speakers, embedding_dim))
- nn.init.xavier_uniform_(self.weight)
-
- self.cos_margin = math.cos(margin)
- self.sin_margin = math.sin(margin)
- self.threshold = math.cos(math.pi - margin)
- self.mm = math.sin(math.pi - margin) * margin
-
- def forward(self, embeddings: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
- """
- Compute Angular Margin Loss with enhanced numerical stability.
-
- Args:
- embeddings: Speaker embeddings (batch_size, embedding_dim)
- labels: Speaker labels (batch_size,)
-
- Returns:
- Loss value
- """
- # Input validation
- if embeddings.size(0) != labels.size(0):
- raise ValueError(f"Batch size mismatch: embeddings {embeddings.size(0)} vs labels {labels.size(0)}")
-
- if labels.max() >= self.num_speakers or labels.min() < 0:
- raise ValueError(f"Invalid labels: min={labels.min()}, max={labels.max()}, expected 0-{self.num_speakers-1}")
-
- # Check for NaN/Inf in inputs
- if torch.isnan(embeddings).any() or torch.isinf(embeddings).any():
- raise ValueError("NaN or Inf detected in embeddings")
-
- # Normalize embeddings and weights with enhanced stability
- embeddings_norm = F.normalize(embeddings, p=2, dim=1, eps=1e-8)
- weight_norm = F.normalize(self.weight, p=2, dim=1, eps=1e-8)
-
- # Compute cosine similarity with numerical stability
- cosine = F.linear(embeddings_norm, weight_norm) # (batch_size, num_speakers)
-
- # Enhanced clamping to prevent numerical issues
- eps = 1e-7
- cosine = torch.clamp(cosine, -1.0 + eps, 1.0 - eps)
-
- # Compute sine with enhanced stability
- sine = torch.sqrt(torch.clamp(1.0 - torch.pow(cosine, 2), min=eps))
-
- # Compute phi (cosine with margin) with stability checks
- phi = cosine * self.cos_margin - sine * self.sin_margin
-
- # Apply threshold with enhanced stability
- phi = torch.where(cosine > self.threshold, phi, cosine - self.mm)
-
- # Create one-hot labels
- one_hot = torch.zeros_like(cosine, dtype=cosine.dtype, device=cosine.device)
- one_hot.scatter_(1, labels.view(-1, 1).long(), 1)
-
- # Apply margin to target class
- output = (one_hot * phi) + ((1.0 - one_hot) * cosine)
-
- # Apply scale with gradient clipping
- output = output * self.scale
-
- # Final stability check
- if torch.isnan(output).any() or torch.isinf(output).any():
- # Fallback to standard cosine similarity
- output = cosine * self.scale
-
- # Additional check after fallback
- if torch.isnan(output).any() or torch.isinf(output).any():
- # Last resort: return a valid loss tensor
- return torch.tensor(1.0, device=output.device, requires_grad=True)
-
- # Compute cross entropy with additional stability
- try:
- loss = F.cross_entropy(output, labels)
- if torch.isnan(loss) or torch.isinf(loss):
- return torch.tensor(1.0, device=output.device, requires_grad=True)
- return loss
- except Exception:
- return torch.tensor(1.0, device=output.device, requires_grad=True)
- class GE2ELoss(nn.Module):
- """
- Generalized End-to-End Loss for speaker verification.
-
- Reference: "Generalized End-to-End Loss for Speaker Verification"
- """
-
- def __init__(self, init_w: float = 10.0, init_b: float = -5.0):
- """
- Initialize GE2E Loss.
-
- Args:
- init_w: Initial value for learnable weight
- init_b: Initial value for learnable bias
- """
- super().__init__()
- self.w = nn.Parameter(torch.tensor(init_w))
- self.b = nn.Parameter(torch.tensor(init_b))
-
- def forward(self, embeddings: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
- """
- Compute GE2E Loss with numerical stability.
-
- Args:
- embeddings: Speaker embeddings (batch_size, embedding_dim)
- labels: Speaker labels (batch_size,)
-
- Returns:
- Loss value
- """
- try:
- # Check for NaN/Inf in inputs
- if torch.isnan(embeddings).any() or torch.isinf(embeddings).any():
- return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
-
- # Normalize embeddings with epsilon for stability
- embeddings = F.normalize(embeddings, p=2, dim=1, eps=1e-8)
-
- # Compute centroids for each speaker
- unique_labels = torch.unique(labels)
-
- if len(unique_labels) == 0:
- return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
-
- centroids = []
-
- for label in unique_labels:
- mask = labels == label
- if mask.sum() == 0:
- continue
- centroid = torch.mean(embeddings[mask], dim=0)
-
- # Check centroid validity
- if torch.isnan(centroid).any() or torch.isinf(centroid).any():
- centroid = torch.zeros_like(centroid)
-
- centroids.append(centroid)
-
- if len(centroids) == 0:
- return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
-
- centroids = torch.stack(centroids)
- centroids = F.normalize(centroids, p=2, dim=1, eps=1e-8)
-
- # Compute similarities
- similarities = torch.mm(embeddings, centroids.t())
-
- # Check similarities for stability
- if torch.isnan(similarities).any() or torch.isinf(similarities).any():
- return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
-
- # Apply learnable parameters with clamping
- w_clamped = torch.clamp(torch.abs(self.w), min=0.1, max=100.0)
- b_clamped = torch.clamp(self.b, min=-50.0, max=50.0)
- similarities = w_clamped * similarities + b_clamped
-
- # Create target labels (map original labels to centroid indices)
- target_labels = torch.zeros_like(labels)
- for i, label in enumerate(unique_labels):
- mask = labels == label
- target_labels[mask] = i
-
- # Final stability check
- if torch.isnan(similarities).any() or torch.isinf(similarities).any():
- return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
-
- loss = F.cross_entropy(similarities, target_labels)
-
- if torch.isnan(loss) or torch.isinf(loss):
- return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
-
- return loss
-
- except Exception:
- return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
|