import torch import torch.nn as nn from transformers import Wav2Vec2Model # ── Shallow frozen wav2vec2 extractor (from train_val.ipynb) ───────────────── 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 # (B, T', 768) # ── Deep audio inference modules (from test__1_.ipynb) ─────────────────────── 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) # (B, T, dim*2) 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 # (B, num_traits), (B, in_dim)