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")