Rasha-Abd-El-Khalik
initial deploy
396c01a
Raw
History Blame Contribute Delete
4.18 kB
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)