Janaxx's picture
fix: restore original model architecture matching checkpoint
2024f1c
Raw
History Blame Contribute Delete
6.1 kB
import torch
import torch.nn as nn
from transformers import AutoModel
from .config import (
PRETRAINED_MODEL, # "facebook/wav2vec2-base"
HIDDEN_DIM, # 768 — Wav2Vec2-base transformer output size
CLASSIFIER_DIM, # 256 — middle layer of the classification head
DROPOUT, # 0.3
NUM_CLASSES, # 4 — angry, happy, sad, neutral
SAMPLE_RATE, # 16000
MAX_DURATION_SEC, # 5
)
from .config import MELD_DROPOUT # NEW: 0.4 — slightly higher dropout for MELD
#number of audio samples in one clip
MAX_SAMPLES = SAMPLE_RATE * MAX_DURATION_SEC # 80 000
class Wav2Vec2EmotionClassifier(nn.Module):
def __init__(self, pretrained_model=PRETRAINED_MODEL, dropout=DROPOUT):
super().__init__()
print(f"Loading pretrained backbone: {pretrained_model} ...")
self.wav2vec2 = AutoModel.from_pretrained(pretrained_model)
print(f"Backbone loaded")
self.classifier = nn.Sequential(
nn.Linear(HIDDEN_DIM, CLASSIFIER_DIM),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(CLASSIFIER_DIM, NUM_CLASSES),
)
self._pretrained_model = pretrained_model
self._dropout = dropout
self._init_classifier_weights()
self.freeze_feature_extractor()
def _init_classifier_weights(self):
for module in self.classifier.modules():
if isinstance(module, nn.Linear):
nn.init.xavier_uniform_(module.weight)
nn.init.zeros_(module.bias)
def freeze_feature_extractor(self):
for param in self.wav2vec2.feature_extractor.parameters():
param.requires_grad = False
print("CNN feature extractor: FROZEN "
f"(trainable params: {self._count_trainable():,})")
def unfreeze_feature_extractor(self):
for param in self.wav2vec2.feature_extractor.parameters():
param.requires_grad = True
print("CNN feature extractor: UNFROZEN "
f"(trainable params: {self._count_trainable():,}) "
" use LR 1e-6 for extractor layers")
def forward(self, input_values):
outputs = self.wav2vec2(
input_values=input_values,
attention_mask=None,
)
hidden_states = outputs.last_hidden_state
pooled = hidden_states.mean(dim=1)
logits = self.classifier(pooled)
return logits
def _count_trainable(self):
return sum(p.numel() for p in self.parameters() if p.requires_grad)
def count_parameters(self):
total = sum(p.numel() for p in self.parameters())
trainable = sum(p.numel() for p in self.parameters() if p.requires_grad)
return trainable, total
def print_summary(self):
trainable, total = self.count_parameters()
frozen = total - trainable
print("\n" + "=" * 60)
print(" Wav2Vec2 Emotion Classifier — Model Summary")
print("=" * 60)
print(f" Backbone : {self._pretrained_model}")
print(f" Input shape : (batch, {MAX_SAMPLES}) "
f"[{MAX_DURATION_SEC}s @ {SAMPLE_RATE}Hz]")
print(f" Transformer output: (batch, ~249, {HIDDEN_DIM})")
print(f" After pooling : (batch, {HIDDEN_DIM})")
print(f" Classifier : {HIDDEN_DIM} {CLASSIFIER_DIM} "
f" {NUM_CLASSES}")
print(f" Dropout : {self._dropout}")
print(f" Num classes : {NUM_CLASSES} "
"(angry, happy, sad, neutral)")
print("-" * 60)
print(f" Total params : {total:>12,}")
print(f" Trainable : {trainable:>12,}")
print(f" Frozen : {frozen:>12,}")
print("=" * 60 + "\n")
def get_layer_groups(self):
return {
'head': list(self.classifier.parameters()),
'transformer': list(self.wav2vec2.encoder.parameters()),
'extractor': list(self.wav2vec2.feature_extractor.parameters()),
}
def build_model(device, pretrained_model=PRETRAINED_MODEL, dropout=DROPOUT):
"""
pretrained_model: defaults to wav2vec2-base (IEMOCAP).
Pass MELD_PRETRAINED_MODEL for MELD (hubert-base).
dropout: defaults to 0.3 (IEMOCAP).
Pass MELD_DROPOUT (0.4) for MELD.
"""
print(f"\nBuilding model on device: {device}")
model = Wav2Vec2EmotionClassifier(pretrained_model=pretrained_model, dropout=dropout)
model = model.to(device)
model.print_summary()
return model
if __name__ == "__main__":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f" Device: {device}")
model = build_model(device)
batch_size = 4
dummy_input = torch.randn(batch_size, MAX_SAMPLES).to(device)
model.eval()
with torch.no_grad():
output = model(dummy_input)
print(f" Input shape : {dummy_input.shape}")
print(f" Output shape : {output.shape}")
print(f" Output sample (logits): {output[0].cpu().numpy()}")
assert output.shape == (batch_size, NUM_CLASSES), \
f"Shape mismatch Expected ({batch_size}, {NUM_CLASSES}), got {output.shape}"
print("\n--- Freeze / Unfreeze Test ---")
model.freeze_feature_extractor()
trainable_frozen, _ = model.count_parameters()
model.unfreeze_feature_extractor()
trainable_unfrozen, total = model.count_parameters()
assert trainable_unfrozen > trainable_frozen, \
"unfreeze_feature_extractor() did not increase trainable params!"
print(f" Frozen trainable params : {trainable_frozen:,}")
print(f" Unfrozen trainable params: {trainable_unfrozen:,}")
print(f" Difference (extractor) : {trainable_unfrozen - trainable_frozen:,}")
print("\nLayer Groups Test")
groups = model.get_layer_groups()
for name, params in groups.items():
n = sum(p.numel() for p in params)
print(f" '{name}' group: {n:,} params")
print("\nAll tests passed Model is ready")