| import torch |
| import torch.nn as nn |
| from transformers import Wav2Vec2Model |
|
|
|
|
| |
|
|
| class FrozenExtractor(nn.Module): |
| """ |
| Loads facebook/wav2vec2-base, keeps only the CNN feature extractor |
| + the first n_shallow transformer layers, freezes everything. |
| Input : (B, T_wave) raw float32 waveform, normalised to [-1, 1] |
| Output: (B, T', 768) contextual features |
| """ |
|
|
| def __init__(self, model_name: str, n_shallow: int): |
| super().__init__() |
| _w2v = Wav2Vec2Model.from_pretrained(model_name) |
| self.feat_extractor = _w2v.feature_extractor |
| self.feat_projection = _w2v.feature_projection |
| self.pos_conv_embed = _w2v.encoder.pos_conv_embed |
| self.layer_norm = _w2v.encoder.layer_norm |
| self.enc_dropout = _w2v.encoder.dropout |
| self.shallow_layers = nn.ModuleList(_w2v.encoder.layers[:n_shallow]) |
| del _w2v |
| torch.cuda.empty_cache() |
| for p in self.parameters(): |
| p.requires_grad = False |
|
|
| @torch.no_grad() |
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| h = self.feat_extractor(x) |
| h = h.transpose(1, 2) |
| h, _ = self.feat_projection(h) |
| pe = self.pos_conv_embed(h) |
| h = self.layer_norm(h + pe) |
| h = self.enc_dropout(h) |
| for layer in self.shallow_layers: |
| h = layer(h)[0] |
| return h |
|
|
|
|
| |
|
|
| class LinearProjection(nn.Module): |
| def __init__(self, in_dim: int, out_dim: int, dropout: float): |
| super().__init__() |
| self.proj = nn.Sequential( |
| nn.Linear(in_dim, out_dim), |
| nn.LayerNorm(out_dim), |
| nn.GELU(), |
| nn.Dropout(dropout), |
| ) |
|
|
| def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: |
| return self.proj(x) * mask.unsqueeze(-1).float() |
|
|
|
|
| class AGRU(nn.Module): |
| def __init__(self, dim: int, num_heads: int): |
| super().__init__() |
| self.gru = nn.GRU(dim, dim // 2, batch_first=True, bidirectional=True) |
| self.attn = nn.MultiheadAttention(embed_dim=dim, num_heads=num_heads, batch_first=True) |
| self.norm = nn.LayerNorm(dim) |
|
|
| def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: |
| gru_out, _ = self.gru(x) |
| attn_out, _ = self.attn(gru_out, gru_out, gru_out, key_padding_mask=~mask) |
| out = self.norm(gru_out + attn_out) * mask.unsqueeze(-1).float() |
| return torch.cat([gru_out, out], dim=-1) |
|
|
|
|
| class AttentionPooling(nn.Module): |
| def __init__(self, dim: int): |
| super().__init__() |
| self.attn = nn.Linear(dim, 1) |
|
|
| def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: |
| scores = self.attn(x).squeeze(-1).masked_fill(~mask, float("-inf")) |
| weights = torch.softmax(scores, dim=1) |
| return (x * weights.unsqueeze(-1)).sum(dim=1) |
|
|
|
|
| class PersonalityRegressor(nn.Module): |
| def __init__( |
| self, |
| in_dim: int, |
| ffn_dim: int, |
| mlp_hidden: int, |
| num_traits: int, |
| num_heads: int, |
| dropout: float, |
| ): |
| super().__init__() |
| enc_layer = nn.TransformerEncoderLayer( |
| d_model=in_dim, |
| nhead=num_heads, |
| dim_feedforward=ffn_dim, |
| dropout=dropout, |
| batch_first=True, |
| ) |
| self.transformer = nn.TransformerEncoder(enc_layer, num_layers=1) |
| self.pool = AttentionPooling(in_dim) |
| self.head = nn.Sequential( |
| nn.LayerNorm(in_dim), |
| nn.Linear(in_dim, mlp_hidden), |
| nn.GELU(), |
| nn.Dropout(0.3), |
| nn.Linear(mlp_hidden, num_traits), |
| nn.Sigmoid(), |
| ) |
|
|
| def forward(self, x: torch.Tensor, mask: torch.Tensor): |
| x = self.transformer(x, src_key_padding_mask=~mask) |
| emb = self.pool(x, mask) |
| return self.head(emb), emb |
|
|