models.py 28 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796
  1. """
  2. State-of-the-art architectures for voice recognition.
  3. This module implements ECAPA-TDNN, TitaNet-S, and SpeakerNet-M architectures
  4. optimized for speaker recognition and voice identification tasks.
  5. """
  6. import math
  7. from typing import Dict, List, Optional, Tuple, Union
  8. import torch
  9. import torch.nn as nn
  10. import torch.nn.functional as F
  11. class SEBlock(nn.Module):
  12. """Squeeze-and-Excitation block for channel attention."""
  13. def __init__(self, channels: int, reduction: int = 16):
  14. super().__init__()
  15. self.squeeze = nn.AdaptiveAvgPool1d(1)
  16. self.excitation = nn.Sequential(
  17. nn.Linear(channels, channels // reduction, bias=False),
  18. nn.ReLU(inplace=True),
  19. nn.Linear(channels // reduction, channels, bias=False),
  20. nn.Sigmoid()
  21. )
  22. def forward(self, x):
  23. b, c, _ = x.size()
  24. y = self.squeeze(x).view(b, c)
  25. y = self.excitation(y).view(b, c, 1)
  26. return x * y.expand_as(x)
  27. class Res2Conv1d(nn.Module):
  28. """
  29. Res2Conv1d implementation for multi-scale feature extraction.
  30. Based on Res2Net architecture adapted for 1D convolutions.
  31. """
  32. def __init__(self, in_channels: int, out_channels: int, kernel_size: int = 3,
  33. dilation: int = 1, scale: int = 4):
  34. super().__init__()
  35. assert out_channels % scale == 0, "out_channels must be divisible by scale"
  36. self.scale = scale
  37. self.out_channels = out_channels
  38. width = out_channels // scale
  39. self.convs = nn.ModuleList()
  40. for i in range(scale - 1):
  41. self.convs.append(
  42. nn.Conv1d(width, width, kernel_size, padding=kernel_size//2 * dilation,
  43. dilation=dilation, bias=False)
  44. )
  45. self.bn = nn.BatchNorm1d(out_channels)
  46. self.relu = nn.ReLU(inplace=True)
  47. def forward(self, x):
  48. batch_size = x.size(0)
  49. # Split input into scale parts
  50. spx = torch.split(x, self.out_channels // self.scale, dim=1)
  51. outputs = [spx[0]]
  52. for i in range(1, self.scale):
  53. if i == 1:
  54. sp = spx[i]
  55. else:
  56. sp = spx[i] + outputs[-1]
  57. sp = self.convs[i-1](sp)
  58. outputs.append(sp)
  59. out = torch.cat(outputs, dim=1)
  60. out = self.bn(out)
  61. out = self.relu(out)
  62. return out
  63. class TDNN(nn.Module):
  64. """Time Delay Neural Network layer."""
  65. def __init__(self, in_channels: int, out_channels: int, kernel_size: int = 5,
  66. dilation: int = 1, use_bn: bool = True, use_relu: bool = True):
  67. super().__init__()
  68. self.conv = nn.Conv1d(
  69. in_channels, out_channels, kernel_size,
  70. padding=kernel_size//2 * dilation, dilation=dilation, bias=not use_bn
  71. )
  72. self.bn = nn.BatchNorm1d(out_channels) if use_bn else None
  73. self.relu = nn.ReLU(inplace=True) if use_relu else None
  74. def forward(self, x):
  75. x = self.conv(x)
  76. if self.bn is not None:
  77. x = self.bn(x)
  78. if self.relu is not None:
  79. x = self.relu(x)
  80. return x
  81. class StatisticalPooling(nn.Module):
  82. """Statistical pooling layer that computes mean and standard deviation."""
  83. def __init__(self, input_dim: int, output_dim: Optional[int] = None):
  84. super().__init__()
  85. self.input_dim = input_dim
  86. self.output_dim = output_dim or input_dim * 2
  87. if output_dim and output_dim != input_dim * 2:
  88. self.projection = nn.Linear(input_dim * 2, output_dim)
  89. else:
  90. self.projection = None
  91. def forward(self, x):
  92. # x shape: (batch, channels, time)
  93. mean = torch.mean(x, dim=2)
  94. # Add epsilon for numerical stability to prevent NaN
  95. std = torch.std(x, dim=2) + 1e-8
  96. # Concatenate mean and std
  97. stats = torch.cat([mean, std], dim=1)
  98. if self.projection is not None:
  99. stats = self.projection(stats)
  100. return stats
  101. class AttentiveStatisticalPooling(nn.Module):
  102. """Attentive statistical pooling with learnable attention weights."""
  103. def __init__(self, input_dim: int, attention_dim: int = 128):
  104. super().__init__()
  105. self.input_dim = input_dim
  106. self.attention_dim = attention_dim
  107. self.attention = nn.Sequential(
  108. nn.Linear(input_dim, attention_dim),
  109. nn.Tanh(),
  110. nn.Linear(attention_dim, 1)
  111. )
  112. def forward(self, x):
  113. # x shape: (batch, channels, time)
  114. batch_size, channels, time_steps = x.shape
  115. # Transpose for attention computation
  116. x_t = x.transpose(1, 2) # (batch, time, channels)
  117. # Compute attention weights
  118. attention_weights = self.attention(x_t) # (batch, time, 1)
  119. attention_weights = F.softmax(attention_weights, dim=1)
  120. # Apply attention weights
  121. weighted_x = x_t * attention_weights # (batch, time, channels)
  122. # Compute weighted statistics
  123. mean = torch.sum(weighted_x, dim=1) # (batch, channels)
  124. # Compute weighted standard deviation with numerical stability
  125. diff = x_t - mean.unsqueeze(1)
  126. weighted_var = torch.sum(attention_weights * diff ** 2, dim=1)
  127. # Clamp variance to prevent negative values and add epsilon
  128. weighted_var = torch.clamp(weighted_var, min=1e-8)
  129. std = torch.sqrt(weighted_var)
  130. # Concatenate mean and std
  131. stats = torch.cat([mean, std], dim=1)
  132. return stats
  133. class ECAPA_TDNN(nn.Module):
  134. """
  135. ECAPA-TDNN: Emphasized Channel Attention, Propagation and Aggregation
  136. in TDNN based Speaker Verification.
  137. Reference: "ECAPA-TDNN: Emphasized Channel Attention, Propagation and
  138. Aggregation in TDNN based Speaker Verification" by Desplanques et al.
  139. """
  140. def __init__(self, input_dim: int = 40, channels: int = 512,
  141. embedding_dim: int = 192, num_speakers: int = None,
  142. use_attention_pooling: bool = True):
  143. """
  144. Initialize ECAPA-TDNN model.
  145. Args:
  146. input_dim: Input feature dimension (e.g., 40 for mel spectrograms)
  147. channels: Number of channels in TDNN layers
  148. embedding_dim: Dimension of speaker embeddings
  149. num_speakers: Number of speakers for classification (None for embeddings only)
  150. use_attention_pooling: Whether to use attentive statistical pooling
  151. """
  152. super().__init__()
  153. self.input_dim = input_dim
  154. self.channels = channels
  155. self.embedding_dim = embedding_dim
  156. self.num_speakers = num_speakers
  157. # Frame-level feature extraction
  158. self.frame_conv = nn.Conv1d(input_dim, channels, 5, padding=2)
  159. # ECAPA blocks
  160. self.ecapa1 = self._make_ecapa_block(channels, channels, kernel_size=3, dilation=2)
  161. self.ecapa2 = self._make_ecapa_block(channels, channels, kernel_size=3, dilation=3)
  162. self.ecapa3 = self._make_ecapa_block(channels, channels, kernel_size=3, dilation=4)
  163. self.ecapa4 = self._make_ecapa_block(channels, 3 * channels, kernel_size=1, dilation=1)
  164. # Aggregation
  165. self.aggregation_conv = nn.Conv1d(3 * channels, 1536, 1)
  166. # Pooling
  167. if use_attention_pooling:
  168. self.pooling = AttentiveStatisticalPooling(1536)
  169. pooling_output_dim = 1536 * 2
  170. else:
  171. self.pooling = StatisticalPooling(1536)
  172. pooling_output_dim = 1536 * 2
  173. # Batch normalization after pooling
  174. self.bn_pooling = nn.BatchNorm1d(pooling_output_dim)
  175. # Embedding layer
  176. self.embedding = nn.Linear(pooling_output_dim, embedding_dim)
  177. self.bn_embedding = nn.BatchNorm1d(embedding_dim)
  178. # Classification layer (optional)
  179. if num_speakers is not None:
  180. self.classifier = nn.Linear(embedding_dim, num_speakers)
  181. else:
  182. self.classifier = None
  183. self._initialize_weights()
  184. def _make_ecapa_block(self, in_channels: int, out_channels: int,
  185. kernel_size: int = 3, dilation: int = 1):
  186. """Create an ECAPA block with Res2Conv and SE."""
  187. return nn.Sequential(
  188. Res2Conv1d(in_channels, out_channels, kernel_size, dilation, scale=8),
  189. SEBlock(out_channels, reduction=16)
  190. )
  191. def forward(self, x, return_embedding: bool = False):
  192. """
  193. Forward pass through ECAPA-TDNN.
  194. Args:
  195. x: Input tensor (batch, features, time)
  196. return_embedding: Whether to return embeddings instead of classification
  197. Returns:
  198. If num_speakers is None or return_embedding is True: speaker embeddings
  199. Otherwise: classification logits
  200. """
  201. # Frame-level processing
  202. x = self.frame_conv(x)
  203. # ECAPA blocks with residual connections
  204. x1 = self.ecapa1(x)
  205. x = x + x1
  206. x2 = self.ecapa2(x)
  207. x = x + x2
  208. x3 = self.ecapa3(x)
  209. x = x + x3
  210. # Aggregation (concatenate all frame-level features)
  211. x_agg = torch.cat([x1, x2, x3], dim=1)
  212. x = self.ecapa4(x_agg)
  213. x = self.aggregation_conv(x)
  214. # Statistical pooling
  215. x = self.pooling(x)
  216. x = self.bn_pooling(x)
  217. # Embedding
  218. embedding = self.embedding(x)
  219. embedding = self.bn_embedding(embedding)
  220. if self.classifier is None or return_embedding:
  221. return F.normalize(embedding, p=2, dim=1)
  222. else:
  223. logits = self.classifier(embedding)
  224. return logits
  225. def get_embeddings(self, x):
  226. """Get speaker embeddings."""
  227. return self.forward(x, return_embedding=True)
  228. def _initialize_weights(self):
  229. """Initialize model weights."""
  230. for m in self.modules():
  231. if isinstance(m, nn.Conv1d):
  232. nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
  233. if m.bias is not None:
  234. nn.init.constant_(m.bias, 0)
  235. elif isinstance(m, nn.BatchNorm1d):
  236. nn.init.constant_(m.weight, 1)
  237. nn.init.constant_(m.bias, 0)
  238. elif isinstance(m, nn.Linear):
  239. nn.init.normal_(m.weight, 0, 0.01)
  240. if m.bias is not None:
  241. nn.init.constant_(m.bias, 0)
  242. class TitaNet_S(nn.Module):
  243. """
  244. TitaNet-S: A lightweight variant of TitaNet for speaker recognition.
  245. Based on the TitaNet architecture but with reduced parameters
  246. for efficient speaker recognition.
  247. """
  248. def __init__(self, input_dim: int = 40, channels: List[int] = None,
  249. embedding_dim: int = 192, num_speakers: int = None,
  250. dropout_rate: float = 0.1):
  251. """
  252. Initialize TitaNet-S model.
  253. Args:
  254. input_dim: Input feature dimension
  255. channels: List of channel dimensions for each block
  256. embedding_dim: Dimension of speaker embeddings
  257. num_speakers: Number of speakers for classification
  258. dropout_rate: Dropout rate for regularization
  259. """
  260. super().__init__()
  261. if channels is None:
  262. channels = [128, 256, 512, 1024]
  263. self.input_dim = input_dim
  264. self.channels = channels
  265. self.embedding_dim = embedding_dim
  266. self.num_speakers = num_speakers
  267. # Input projection
  268. self.input_conv = nn.Conv1d(input_dim, channels[0], 1)
  269. self.input_bn = nn.BatchNorm1d(channels[0])
  270. # TitaNet blocks
  271. self.blocks = nn.ModuleList()
  272. for i in range(len(channels) - 1):
  273. block = self._make_titanet_block(
  274. channels[i], channels[i+1], dropout_rate
  275. )
  276. self.blocks.append(block)
  277. # Statistical pooling
  278. self.pooling = AttentiveStatisticalPooling(channels[-1])
  279. # Embedding layers
  280. pooling_dim = channels[-1] * 2
  281. self.embedding_layers = nn.Sequential(
  282. nn.Linear(pooling_dim, channels[-1]),
  283. nn.BatchNorm1d(channels[-1]),
  284. nn.ReLU(inplace=True),
  285. nn.Dropout(dropout_rate),
  286. nn.Linear(channels[-1], embedding_dim),
  287. nn.BatchNorm1d(embedding_dim)
  288. )
  289. # Classification layer (optional)
  290. if num_speakers is not None:
  291. self.classifier = nn.Linear(embedding_dim, num_speakers)
  292. else:
  293. self.classifier = None
  294. self._initialize_weights()
  295. def _make_titanet_block(self, in_channels: int, out_channels: int,
  296. dropout_rate: float):
  297. """Create a TitaNet block."""
  298. return nn.Sequential(
  299. # Depthwise separable convolution
  300. nn.Conv1d(in_channels, in_channels, 3, padding=1, groups=in_channels),
  301. nn.Conv1d(in_channels, out_channels, 1),
  302. nn.BatchNorm1d(out_channels),
  303. nn.ReLU(inplace=True),
  304. nn.Dropout1d(dropout_rate),
  305. # Squeeze-and-Excitation
  306. SEBlock(out_channels, reduction=8),
  307. # Residual connection (if dimensions match)
  308. # Note: Actual residual connection would need dimension matching
  309. )
  310. def forward(self, x, return_embedding: bool = False):
  311. """Forward pass through TitaNet-S."""
  312. # Input projection
  313. x = self.input_conv(x)
  314. x = self.input_bn(x)
  315. x = F.relu(x, inplace=True)
  316. # TitaNet blocks
  317. for block in self.blocks:
  318. x = block(x)
  319. # Statistical pooling
  320. x = self.pooling(x)
  321. # Embedding
  322. embedding = self.embedding_layers(x)
  323. if self.classifier is None or return_embedding:
  324. return F.normalize(embedding, p=2, dim=1)
  325. else:
  326. logits = self.classifier(embedding)
  327. return logits
  328. def get_embeddings(self, x):
  329. """Get speaker embeddings."""
  330. return self.forward(x, return_embedding=True)
  331. def _initialize_weights(self):
  332. """Initialize model weights."""
  333. for m in self.modules():
  334. if isinstance(m, nn.Conv1d):
  335. nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
  336. if m.bias is not None:
  337. nn.init.constant_(m.bias, 0)
  338. elif isinstance(m, nn.BatchNorm1d):
  339. nn.init.constant_(m.weight, 1)
  340. nn.init.constant_(m.bias, 0)
  341. elif isinstance(m, nn.Linear):
  342. nn.init.normal_(m.weight, 0, 0.01)
  343. if m.bias is not None:
  344. nn.init.constant_(m.bias, 0)
  345. class SpeakerNet_M(nn.Module):
  346. """
  347. SpeakerNet-M: Medium-sized speaker recognition network.
  348. A balanced architecture for speaker recognition that provides
  349. good performance with moderate computational requirements.
  350. """
  351. def __init__(self, input_dim: int = 40, hidden_dim: int = 512,
  352. embedding_dim: int = 256, num_speakers: int = None,
  353. num_layers: int = 4, dropout_rate: float = 0.1):
  354. """
  355. Initialize SpeakerNet-M model.
  356. Args:
  357. input_dim: Input feature dimension
  358. hidden_dim: Hidden dimension for TDNN layers
  359. embedding_dim: Dimension of speaker embeddings
  360. num_speakers: Number of speakers for classification
  361. num_layers: Number of TDNN layers
  362. dropout_rate: Dropout rate for regularization
  363. """
  364. super().__init__()
  365. self.input_dim = input_dim
  366. self.hidden_dim = hidden_dim
  367. self.embedding_dim = embedding_dim
  368. self.num_speakers = num_speakers
  369. self.num_layers = num_layers
  370. # Input layer
  371. self.input_layer = TDNN(input_dim, hidden_dim, kernel_size=5)
  372. # TDNN layers with different dilations
  373. self.tdnn_layers = nn.ModuleList()
  374. dilations = [1, 2, 3, 4, 5][:num_layers]
  375. for i, dilation in enumerate(dilations):
  376. layer = TDNN(
  377. hidden_dim, hidden_dim,
  378. kernel_size=3, dilation=dilation
  379. )
  380. self.tdnn_layers.append(layer)
  381. # Context expansion
  382. self.context_conv = nn.Conv1d(hidden_dim, hidden_dim * 2, 1)
  383. self.context_bn = nn.BatchNorm1d(hidden_dim * 2)
  384. # Statistical pooling
  385. self.pooling = StatisticalPooling(hidden_dim * 2)
  386. # Embedding layers
  387. pooling_dim = hidden_dim * 4 # 2x for mean+std
  388. self.embedding_layers = nn.Sequential(
  389. nn.Linear(pooling_dim, hidden_dim),
  390. nn.BatchNorm1d(hidden_dim),
  391. nn.ReLU(inplace=True),
  392. nn.Dropout(dropout_rate),
  393. nn.Linear(hidden_dim, embedding_dim),
  394. nn.BatchNorm1d(embedding_dim)
  395. )
  396. # Classification layer (optional)
  397. if num_speakers is not None:
  398. self.classifier = nn.Sequential(
  399. nn.Dropout(dropout_rate),
  400. nn.Linear(embedding_dim, num_speakers)
  401. )
  402. else:
  403. self.classifier = None
  404. self._initialize_weights()
  405. def forward(self, x, return_embedding: bool = False):
  406. """Forward pass through SpeakerNet-M."""
  407. # Input processing
  408. x = self.input_layer(x)
  409. # TDNN layers with residual connections
  410. for layer in self.tdnn_layers:
  411. residual = x
  412. x = layer(x)
  413. x = x + residual # Residual connection
  414. # Context expansion
  415. x = self.context_conv(x)
  416. x = self.context_bn(x)
  417. x = F.relu(x, inplace=True)
  418. # Statistical pooling
  419. x = self.pooling(x)
  420. # Embedding
  421. embedding = self.embedding_layers(x)
  422. if self.classifier is None or return_embedding:
  423. return F.normalize(embedding, p=2, dim=1)
  424. else:
  425. logits = self.classifier(embedding)
  426. return logits
  427. def get_embeddings(self, x):
  428. """Get speaker embeddings."""
  429. return self.forward(x, return_embedding=True)
  430. def _initialize_weights(self):
  431. """Initialize model weights."""
  432. for m in self.modules():
  433. if isinstance(m, nn.Conv1d):
  434. nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
  435. if m.bias is not None:
  436. nn.init.constant_(m.bias, 0)
  437. elif isinstance(m, nn.BatchNorm1d):
  438. nn.init.constant_(m.weight, 1)
  439. nn.init.constant_(m.bias, 0)
  440. elif isinstance(m, nn.Linear):
  441. nn.init.normal_(m.weight, 0, 0.01)
  442. if m.bias is not None:
  443. nn.init.constant_(m.bias, 0)
  444. def create_voice_recognition_model(model_type: str = "ecapa_tdnn", **kwargs) -> nn.Module:
  445. """
  446. Factory function to create voice recognition models.
  447. Args:
  448. model_type: Type of model ('ecapa_tdnn', 'titanet_s', 'speakernet_m')
  449. **kwargs: Additional arguments for model initialization
  450. Returns:
  451. Voice recognition model instance
  452. """
  453. if model_type.lower() == "ecapa_tdnn":
  454. return ECAPA_TDNN(**kwargs)
  455. elif model_type.lower() == "titanet_s":
  456. return TitaNet_S(**kwargs)
  457. elif model_type.lower() == "speakernet_m":
  458. return SpeakerNet_M(**kwargs)
  459. else:
  460. raise ValueError(f"Unknown model type: {model_type}")
  461. def count_parameters(model: nn.Module) -> Dict[str, int]:
  462. """Count model parameters."""
  463. total_params = sum(p.numel() for p in model.parameters())
  464. trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
  465. return {
  466. "total_parameters": total_params,
  467. "trainable_parameters": trainable_params,
  468. "non_trainable_parameters": total_params - trainable_params
  469. }
  470. class AngularMarginLoss(nn.Module):
  471. """
  472. Angular Margin Loss (ArcFace) for speaker recognition.
  473. Reference: "ArcFace: Additive Angular Margin Loss for Deep Face Recognition"
  474. """
  475. def __init__(self, embedding_dim: int, num_speakers: int,
  476. margin: float = 0.5, scale: float = 64.0):
  477. """
  478. Initialize Angular Margin Loss.
  479. Args:
  480. embedding_dim: Dimension of input embeddings
  481. num_speakers: Number of speaker classes
  482. margin: Angular margin (m)
  483. scale: Scale factor (s)
  484. """
  485. super().__init__()
  486. self.embedding_dim = embedding_dim
  487. self.num_speakers = num_speakers
  488. self.margin = margin
  489. self.scale = scale
  490. # Weight matrix (represents speaker centroids)
  491. self.weight = nn.Parameter(torch.randn(num_speakers, embedding_dim))
  492. nn.init.xavier_uniform_(self.weight)
  493. self.cos_margin = math.cos(margin)
  494. self.sin_margin = math.sin(margin)
  495. self.threshold = math.cos(math.pi - margin)
  496. self.mm = math.sin(math.pi - margin) * margin
  497. def forward(self, embeddings: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
  498. """
  499. Compute Angular Margin Loss with enhanced numerical stability.
  500. Args:
  501. embeddings: Speaker embeddings (batch_size, embedding_dim)
  502. labels: Speaker labels (batch_size,)
  503. Returns:
  504. Loss value
  505. """
  506. # Input validation
  507. if embeddings.size(0) != labels.size(0):
  508. raise ValueError(f"Batch size mismatch: embeddings {embeddings.size(0)} vs labels {labels.size(0)}")
  509. if labels.max() >= self.num_speakers or labels.min() < 0:
  510. raise ValueError(f"Invalid labels: min={labels.min()}, max={labels.max()}, expected 0-{self.num_speakers-1}")
  511. # Check for NaN/Inf in inputs
  512. if torch.isnan(embeddings).any() or torch.isinf(embeddings).any():
  513. raise ValueError("NaN or Inf detected in embeddings")
  514. # Normalize embeddings and weights with enhanced stability
  515. embeddings_norm = F.normalize(embeddings, p=2, dim=1, eps=1e-8)
  516. weight_norm = F.normalize(self.weight, p=2, dim=1, eps=1e-8)
  517. # Compute cosine similarity with numerical stability
  518. cosine = F.linear(embeddings_norm, weight_norm) # (batch_size, num_speakers)
  519. # Enhanced clamping to prevent numerical issues
  520. eps = 1e-7
  521. cosine = torch.clamp(cosine, -1.0 + eps, 1.0 - eps)
  522. # Compute sine with enhanced stability
  523. sine = torch.sqrt(torch.clamp(1.0 - torch.pow(cosine, 2), min=eps))
  524. # Compute phi (cosine with margin) with stability checks
  525. phi = cosine * self.cos_margin - sine * self.sin_margin
  526. # Apply threshold with enhanced stability
  527. phi = torch.where(cosine > self.threshold, phi, cosine - self.mm)
  528. # Create one-hot labels
  529. one_hot = torch.zeros_like(cosine, dtype=cosine.dtype, device=cosine.device)
  530. one_hot.scatter_(1, labels.view(-1, 1).long(), 1)
  531. # Apply margin to target class
  532. output = (one_hot * phi) + ((1.0 - one_hot) * cosine)
  533. # Apply scale with gradient clipping
  534. output = output * self.scale
  535. # Final stability check
  536. if torch.isnan(output).any() or torch.isinf(output).any():
  537. # Fallback to standard cosine similarity
  538. output = cosine * self.scale
  539. # Additional check after fallback
  540. if torch.isnan(output).any() or torch.isinf(output).any():
  541. # Last resort: return a valid loss tensor
  542. return torch.tensor(1.0, device=output.device, requires_grad=True)
  543. # Compute cross entropy with additional stability
  544. try:
  545. loss = F.cross_entropy(output, labels)
  546. if torch.isnan(loss) or torch.isinf(loss):
  547. return torch.tensor(1.0, device=output.device, requires_grad=True)
  548. return loss
  549. except Exception:
  550. return torch.tensor(1.0, device=output.device, requires_grad=True)
  551. class GE2ELoss(nn.Module):
  552. """
  553. Generalized End-to-End Loss for speaker verification.
  554. Reference: "Generalized End-to-End Loss for Speaker Verification"
  555. """
  556. def __init__(self, init_w: float = 10.0, init_b: float = -5.0):
  557. """
  558. Initialize GE2E Loss.
  559. Args:
  560. init_w: Initial value for learnable weight
  561. init_b: Initial value for learnable bias
  562. """
  563. super().__init__()
  564. self.w = nn.Parameter(torch.tensor(init_w))
  565. self.b = nn.Parameter(torch.tensor(init_b))
  566. def forward(self, embeddings: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
  567. """
  568. Compute GE2E Loss with numerical stability.
  569. Args:
  570. embeddings: Speaker embeddings (batch_size, embedding_dim)
  571. labels: Speaker labels (batch_size,)
  572. Returns:
  573. Loss value
  574. """
  575. try:
  576. # Check for NaN/Inf in inputs
  577. if torch.isnan(embeddings).any() or torch.isinf(embeddings).any():
  578. return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
  579. # Normalize embeddings with epsilon for stability
  580. embeddings = F.normalize(embeddings, p=2, dim=1, eps=1e-8)
  581. # Compute centroids for each speaker
  582. unique_labels = torch.unique(labels)
  583. if len(unique_labels) == 0:
  584. return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
  585. centroids = []
  586. for label in unique_labels:
  587. mask = labels == label
  588. if mask.sum() == 0:
  589. continue
  590. centroid = torch.mean(embeddings[mask], dim=0)
  591. # Check centroid validity
  592. if torch.isnan(centroid).any() or torch.isinf(centroid).any():
  593. centroid = torch.zeros_like(centroid)
  594. centroids.append(centroid)
  595. if len(centroids) == 0:
  596. return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
  597. centroids = torch.stack(centroids)
  598. centroids = F.normalize(centroids, p=2, dim=1, eps=1e-8)
  599. # Compute similarities
  600. similarities = torch.mm(embeddings, centroids.t())
  601. # Check similarities for stability
  602. if torch.isnan(similarities).any() or torch.isinf(similarities).any():
  603. return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
  604. # Apply learnable parameters with clamping
  605. w_clamped = torch.clamp(torch.abs(self.w), min=0.1, max=100.0)
  606. b_clamped = torch.clamp(self.b, min=-50.0, max=50.0)
  607. similarities = w_clamped * similarities + b_clamped
  608. # Create target labels (map original labels to centroid indices)
  609. target_labels = torch.zeros_like(labels)
  610. for i, label in enumerate(unique_labels):
  611. mask = labels == label
  612. target_labels[mask] = i
  613. # Final stability check
  614. if torch.isnan(similarities).any() or torch.isinf(similarities).any():
  615. return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
  616. loss = F.cross_entropy(similarities, target_labels)
  617. if torch.isnan(loss) or torch.isinf(loss):
  618. return torch.tensor(1.0, device=embeddings.device, requires_grad=True)
  619. return loss
  620. except Exception:
  621. return torch.tensor(1.0, device=embeddings.device, requires_grad=True)