Spaces:
Sleeping
Sleeping
| 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") |