Views
No views yet

MultiTaskBackbone — a frozen/finetuned hubert_base encoder
(facebook/hubert-base-ls960) with per-task learnable softmax weights over hidden
layers, mean-pooled over time, feeding three independent MLP classification heads.facebook/hubert-base-ls960| task | classes | labels |
|---|---|---|
| emotion | 7 | angry, disgusted, fearful, happy, neutral, sad, surprised |
| gender | 2 | F, M |
| age | 4 | adult, child, senior, young |
pip install torch transformers huggingface_hub librosa numpy1import json
2
3import librosa
4import numpy as np
5import torch
6import torch.nn as nn
7import torch.nn.functional as F
8from huggingface_hub import hf_hub_download
9from transformers import HubertModel, Wav2Vec2FeatureExtractor
10
11REPO_ID = "kazega0/KazEGA-Hubert"
12PRETRAINED = "facebook/hubert-base-ls960"
13SAMPLE_RATE = 16_000
14MAX_LENGTH = 160_000
15AUDIO_PATH = "sample.wav"
16DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
17
18class MultiTaskHubert(nn.Module):
19 def __init__(self, num_emotions, num_genders, num_ages, pretrained=PRETRAINED):
20 super().__init__()
21 self.hubert = HubertModel.from_pretrained(pretrained, output_hidden_states=True)
22 hidden_size = self.hubert.config.hidden_size
23 num_layers = self.hubert.config.num_hidden_layers + 1
24
25 self.emotion_weights = nn.Parameter(torch.ones(num_layers))
26 self.gender_weights = nn.Parameter(torch.ones(num_layers))
27 self.age_weights = nn.Parameter(torch.ones(num_layers))
28
29 self.emotion_head = nn.Sequential(
30 nn.Linear(hidden_size, 256), nn.ReLU(), nn.Dropout(0.2),
31 nn.Linear(256, num_emotions))
32 self.gender_head = nn.Sequential(
33 nn.Linear(hidden_size, 256), nn.ReLU(), nn.Dropout(0.1),
34 nn.Linear(256, num_genders))
35 self.age_head = nn.Sequential(
36 nn.Linear(hidden_size, 256), nn.ReLU(), nn.Dropout(0.1),
37 nn.Linear(256, num_ages))
38
39 def forward(self, input_values, input_length):
40 hidden = torch.stack(self.hubert(input_values).hidden_states, dim=0)
41 feat_len = int(self.hubert._get_feat_extract_output_lengths(torch.tensor(input_length)))
42 hidden = hidden[:, :, :feat_len, :]
43
44 def pool(layer_weights):
45 w = torch.softmax(layer_weights, dim=0)
46 return (w.view(-1, 1, 1, 1) * hidden).sum(dim=0).mean(dim=1)
47
48 return (self.emotion_head(pool(self.emotion_weights)),
49 self.gender_head(pool(self.gender_weights)),
50 self.age_head(pool(self.age_weights)))
51
52ckpt = torch.load(hf_hub_download(REPO_ID, "model.pt"), map_location="cpu", weights_only=False)
53encoders = json.load(open(hf_hub_download(REPO_ID, "label_encoders.json")))
54id2label = {task: {idx: name for name, idx in m.items()} for task, m in encoders.items()}
55
56model = MultiTaskHubert(ckpt["num_emotions"], ckpt["num_genders"], ckpt["num_ages"])
57model.load_state_dict(ckpt["model_state_dict"])
58model.to(DEVICE).eval()
59
60processor = Wav2Vec2FeatureExtractor.from_pretrained(PRETRAINED)
61
62audio, _ = librosa.load(AUDIO_PATH, sr=SAMPLE_RATE, mono=True)
63audio = audio[:MAX_LENGTH]
64raw_length = len(audio)
65audio = np.pad(audio, (0, MAX_LENGTH - raw_length))
66input_values = processor(audio, sampling_rate=SAMPLE_RATE,
67 return_tensors="pt").input_values.to(DEVICE)
68
69with torch.no_grad():
70 logits = model(input_values, raw_length)
71
72for task, task_logits in zip(("emotion", "gender", "age"), logits):
73 probs = F.softmax(task_logits[0], dim=0)
74 idx = int(probs.argmax())
75 print(f"{task:8s} {id2label[task][idx]:10s} ({probs[idx]:.1%})")