Views
No views yet
nixiieee/dusha_balanced dataset (cropped version of original dusha dataset).torch==2.5.1 torchaudio==2.5.1 transformers==4.49.0 accelerate==1.5.21from transformers import (
2 PreTrainedModel,
3 AutoConfig,
4 AutoModel,
5 AutoProcessor,
6)
7from transformers.modeling_outputs import SequenceClassifierOutput
8import torch
9import torch.nn as nn
10import torchaudio
11
12class EmotionClassifier(nn.Module):
13 def __init__(self, hidden_size, num_labels=5, dropout=0.2):
14 super().__init__()
15 self.pool_norm = nn.LayerNorm(hidden_size)
16 self.pre_dropout = nn.Dropout(dropout)
17
18 mid1 = max(hidden_size // 2, num_labels * 4)
19 mid2 = max(hidden_size // 4, num_labels * 2)
20
21 self.classifier = nn.Sequential(
22 nn.Linear(hidden_size, mid1),
23 nn.GELU(),
24 nn.Dropout(dropout),
25 nn.LayerNorm(mid1),
26 nn.Linear(mid1, mid2),
27 nn.GELU(),
28 nn.Dropout(dropout),
29 nn.LayerNorm(mid2),
30 nn.Linear(mid2, num_labels),
31 )
32
33 def forward(self, hidden_states, attention_mask=None):
34 if attention_mask is not None:
35 lengths = attention_mask.sum(dim=1, keepdim=True)
36 masked = hidden_states * attention_mask.unsqueeze(-1)
37 pooled = masked.sum(dim=1) / lengths
38 else:
39 pooled = hidden_states.mean(dim=1)
40 x = self.pool_norm(pooled)
41 x = self.pre_dropout(x)
42 logits = self.classifier(x)
43 return logits
44
45class ModelForEmotionClassification(PreTrainedModel):
46 config_class = AutoConfig
47
48 def __init__(
49 self, config, model_name, num_labels=5, dropout=0.2
50 ):
51 super().__init__(config)
52 self.encoder = AutoModel.from_pretrained(model_name, trust_remote_code=True).model.encoder
53 hidden_size = config.encoder['d_model']
54 self.classifier = EmotionClassifier(
55 hidden_size, num_labels=num_labels, dropout=dropout
56 )
57 self.post_init()
58
59 def forward(
60 self,
61 input_features: torch.Tensor,
62 input_lengths: torch.Tensor,
63 attention_mask: torch.Tensor = None,
64 labels: torch.Tensor = None
65 ) -> SequenceClassifierOutput:
66 encoded, out_lens = self.encoder(input_features, input_lengths)
67 hidden_states = encoded.transpose(1, 2)
68
69 if attention_mask is None:
70 max_t = hidden_states.size(1)
71 attention_mask = (
72 torch.arange(max_t, device=out_lens.device)
73 .unsqueeze(0)
74 .lt(out_lens.unsqueeze(1))
75 .long()
76 )
77
78 logits = self.classifier(hidden_states, attention_mask=attention_mask)
79
80 loss = None
81 if labels is not None:
82 loss_fct = nn.CrossEntropyLoss()
83 loss = loss_fct(logits, labels)
84
85 return SequenceClassifierOutput(loss=loss, logits=logits)
86
87model_name = "nixiieee/gigaam-rnnt-emotion-classifier-dusha"
88processor = AutoProcessor.from_pretrained(model_name, trust_remote_code=True)
89config = AutoConfig.from_pretrained(model_name, trust_remote_code=True)
90model = ModelForEmotionClassification.from_pretrained(model_name, config=config, model_name=model_name)
91model.eval()
92
93# load audio
94wav, sr = torchaudio.load("audio.wav")
95# resample if necessary
96wav = torchaudio.functional.resample(wav, sr, 16000)
97input_features = processor(wav[0], sampling_rate=16000, return_tensors="pt")
98
99with torch.no_grad():
100 pred_ids = model.generate(**input_features)
101
102pred = pred_ids.logits.argmax(dim=-1).item()
103print("Predicted emotion:", config.id2label[pred])