Views
No views yet
| Training Loss | Epoch | Step | Validation Loss | Accuracy | Balanced Accuracy | Precision | Recall | F1 |
|---|---|---|---|---|---|---|---|---|
| 0.7498 | 1.0 | 4609 | 0.6200 | 0.7547 | 0.7835 | 0.7821 | 0.7835 | 0.7819 |
| 0.5963 | 2.0 | 9218 | 0.5831 | 0.78 | 0.8069 | 0.8216 | 0.8069 | 0.8118 |
| 0.5422 | 3.0 | 13827 | 0.5668 | 0.79 | 0.8154 | 0.8255 | 0.8154 | 0.8196 |
| 0.7421 | 4.0 | 18436 | 0.5599 | 0.7889 | 0.8165 | 0.8235 | 0.8165 | 0.8187 |
1from transformers.modeling_outputs import SequenceClassifierOutput
2from transformers import AutoProcessor, WhisperForAudioClassification, AutoConfig, PreTrainedModel, WhisperModel
3import torch.nn as nn
4
5class WhisperClassifier(nn.Module):
6 def __init__(self, hidden_size, num_labels=5, dropout=0.2):
7 super().__init__()
8 self.pool_norm = nn.LayerNorm(hidden_size)
9 self.pre_dropout = nn.Dropout(dropout)
10
11 mid1 = max(hidden_size // 2, num_labels * 4)
12 mid2 = max(hidden_size // 4, num_labels * 2)
13
14 self.classifier = nn.Sequential(
15 nn.Linear(hidden_size, mid1),
16 nn.GELU(),
17 nn.Dropout(dropout),
18 nn.LayerNorm(mid1),
19 nn.Linear(mid1, mid2),
20 nn.GELU(),
21 nn.Dropout(dropout),
22 nn.LayerNorm(mid2),
23 nn.Linear(mid2, num_labels),
24 )
25
26 def forward(self, hidden_states, attention_mask=None):
27 if attention_mask is not None:
28 lengths = attention_mask.sum(dim=1, keepdim=True)
29 masked = hidden_states * attention_mask.unsqueeze(-1)
30 pooled = masked.sum(dim=1) / lengths
31 else:
32 pooled = hidden_states.mean(dim=1)
33 x = self.pool_norm(pooled)
34 x = self.pre_dropout(x)
35 logits = self.classifier(x)
36 return logits
37
38class WhisperForEmotionClassification(PreTrainedModel):
39 config_class = AutoConfig
40
41 def __init__(
42 self, config, model_name="openai/whisper-small", num_labels=5, dropout=0.2
43 ):
44 super().__init__(config)
45 self.encoder = WhisperModel.from_pretrained(model_name).encoder
46 hidden_size = config.hidden_size
47 self.classifier = WhisperClassifier(
48 hidden_size, num_labels=num_labels, dropout=dropout
49 )
50 self.post_init()
51
52 def forward(self, input_features, attention_mask=None, labels=None):
53 encoder_output = self.encoder(
54 input_features=input_features,
55 attention_mask=attention_mask,
56 return_dict=True,
57 )
58 hidden_states = encoder_output.last_hidden_state
59 logits = self.classifier(hidden_states, attention_mask=attention_mask)
60 loss = None
61 if labels is not None:
62 loss = nn.CrossEntropyLoss()(
63 logits.view(-1, logits.size(-1)), labels.view(-1)
64 )
65 return SequenceClassifierOutput(
66 loss=loss,
67 logits=logits,
68 )
69
70EMOTION_LABELS = ['neutral', 'angry', 'positive', 'sad', 'other']
71
72model_name = "nixiieee/whisper-small-emotion-classifier-dusha"
73processor = WhisperProcessor.from_pretrained("openai/whisper-small", return_attention_mask=True)
74config = AutoConfig.from_pretrained(model_name)
75model = WhisperForEmotionClassification.from_pretrained(model_name, num_labels=5, dropout=0.1)
76model.eval()
77
78# load audio
79wav, sr = torchaudio.load("audio.wav")
80# resample if necessary
81wav = torchaudio.functional.resample(wav, sr, 16000)
82input_features = processor(wav[0], sampling_rate=16000, return_tensors="pt")
83
84with torch.no_grad():
85 pred_ids = model.generate(**input_features)
86
87pred = pred_ids.logits.argmax(dim=-1).item()
88print("Predicted emotion:", EMOTION_LABELS[pred])