File size: 6,803 Bytes
d22cc5f
 
 
ae5507d
 
8a16cc4
ae5507d
d22cc5f
ae5507d
d22cc5f
 
92d86ab
ae5507d
8a16cc4
d22cc5f
 
8a16cc4
d22cc5f
8a16cc4
d22cc5f
8a16cc4
 
 
 
 
 
e25a431
8a16cc4
 
 
 
 
d22cc5f
92d86ab
8a16cc4
d22cc5f
 
8a16cc4
 
 
 
 
d22cc5f
8a16cc4
 
 
 
d22cc5f
8a16cc4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a613a36
d22cc5f
8a16cc4
 
d22cc5f
8a16cc4
 
d22cc5f
92d86ab
d22cc5f
ae5507d
8a16cc4
a6a017b
d22cc5f
8a16cc4
 
 
 
 
 
993e6ad
8a16cc4
 
92d86ab
 
8a16cc4
 
 
 
 
 
 
 
 
 
 
 
993e6ad
52119a0
8a16cc4
d22cc5f
8a16cc4
 
 
993e6ad
 
 
8a16cc4
52119a0
8a16cc4
 
 
 
 
d22cc5f
993e6ad
8a16cc4
92d86ab
8a16cc4
92d86ab
8a16cc4
 
 
 
 
 
 
 
 
 
 
 
 
d22cc5f
8a16cc4
a613a36
 
52119a0
 
d22cc5f
ae5507d
8a16cc4
52119a0
8a16cc4
 
92d86ab
 
52119a0
 
 
993e6ad
 
92d86ab
52119a0
 
92d86ab
993e6ad
 
8a16cc4
92d86ab
ae5507d
52119a0
d22cc5f
 
52119a0
 
ae5507d
 
92d86ab
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
import torch
import torch.nn as nn
import librosa
import numpy as np
import gradio as gr
from transformers import XLMRobertaModel, XLMRobertaTokenizer, WavLMModel
from huggingface_hub import hf_hub_download
import warnings

warnings.filterwarnings('ignore')

# ── CONFIGURATION ──
REPO_ID = "anggars/neural-mathrock"
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
TARGET_COLS = ['mbti', 'emotion', 'vibe', 'intensity', 'tempo']

print("Downloading model from Hugging Face Hub...")
model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
ckpt = torch.load(model_path, map_location=DEVICE, weights_only=False)

# ── LOAD LABEL ENCODERS ──
le_mbti = ckpt['le_mbti']
le_emotion = ckpt['le_emotion']
le_vibe = ckpt['le_vibe']
le_intensity = ckpt['le_intensity']
le_tempo = ckpt['le_tempo']

NUM_MBTI = len(le_mbti.classes_)
NUM_EMOTION = len(le_emotion.classes_)
NUM_VIBE = len(le_vibe.classes_)
NUM_INTENSITY = len(le_intensity.classes_)
NUM_TEMPO = len(le_tempo.classes_)

# ── ARCHITECTURE ──
class HybridMultimodalModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.text_model = XLMRobertaModel.from_pretrained('anggars/xlm-mbti')
        self.audio_model = WavLMModel.from_pretrained('microsoft/wavlm-base')
        
        self.audio_proj = nn.Linear(768, 256)
        
        self.fusion = nn.Sequential(
            nn.Linear(1024, 512),
            nn.BatchNorm1d(512),
            nn.ReLU(),
            nn.Dropout(0.4),
        )
        
        self.head_mbti = nn.Linear(512, NUM_MBTI)
        self.head_emotion = nn.Linear(512, NUM_EMOTION)
        self.head_vibe = nn.Linear(512, NUM_VIBE)
        self.head_intensity = nn.Linear(512, NUM_INTENSITY)
        self.head_tempo = nn.Linear(512, NUM_TEMPO)

    def forward(self, input_ids, attention_mask, audio_values):
        text_out = self.text_model(input_ids=input_ids, attention_mask=attention_mask)
        text_feat = text_out.pooler_output
        
        audio_out = self.audio_model(audio_values).last_hidden_state
        audio_feat = self.audio_proj(audio_out.mean(dim=1))
        
        fused = self.fusion(torch.cat([text_feat, audio_feat], dim=-1))
        
        return {
            'mbti': self.head_mbti(fused),
            'emotion': self.head_emotion(fused),
            'vibe': self.head_vibe(fused),
            'intensity': self.head_intensity(fused),
            'tempo': self.head_tempo(fused),
        }

print("Initializing architecture and loading weights...")
model = HybridMultimodalModel().to(DEVICE)
model.load_state_dict(ckpt['model_state'], strict=False)
model.eval()

tokenizer = XLMRobertaTokenizer.from_pretrained('anggars/xlm-mbti')

# ── HYBRID INFERENCE PIPELINE ──
def predict_multimodal(audio_path, lyrics):
    if audio_path is None:
        return [{"Error": "Audio file cannot be empty."}] * 5
        
    try:
        # 1. Text Processing
        text = str(lyrics).strip() if lyrics else "[INSTRUMENTAL]"
        enc = tokenizer(text, truncation=True, padding='max_length', max_length=128, return_tensors='pt').to(DEVICE)

        # 2. Audio Processing & DSP
        wav, sr = librosa.load(audio_path, sr=16000)
        
        onset_env = librosa.onset.onset_strength(y=wav, sr=sr)
        tempo_bpm, _ = librosa.beat.beat_track(onset_envelope=onset_env, sr=sr)
        print(f"Calculated BPM: {tempo_bpm}")

        # 3. Chunking (20-second sequential blocks)
        chunk_size = 16000 * 20
        chunks = []
        for i in range(0, len(wav), chunk_size):
            chunk = wav[i:i + chunk_size]
            if len(chunk) < chunk_size:
                chunk = np.pad(chunk, (0, chunk_size - len(chunk)))
            chunks.append(chunk)

        if not chunks:
            chunks = [np.zeros(chunk_size)]

        all_logits = {col: [] for col in TARGET_COLS}
        
        # 4. Block-by-block Prediction
        with torch.no_grad():
            for chunk in chunks:
                audio_tensor = torch.tensor(chunk, dtype=torch.float32).unsqueeze(0).to(DEVICE)
                out = model(enc['input_ids'], enc['attention_mask'], audio_tensor)
                for col in TARGET_COLS:
                    all_logits[col].append(out[col][0])
        
        # 5. Averaging and Formatting Results
        final_results = []
        encoders_map = {
            'mbti': le_mbti, 'emotion': le_emotion, 'vibe': le_vibe, 
            'intensity': le_intensity, 'tempo': le_tempo
        }
        
        for col in TARGET_COLS:
            avg_logits = torch.stack(all_logits[col]).mean(dim=0)
            le = encoders_map[col]
            
            # TEMPO LOGIC OVERRIDE
            if col == 'tempo':
                tempo_classes = list(le.classes_)
                try:
                    fast_idx = tempo_classes.index('Fast')
                    mod_idx = tempo_classes.index('Moderate')
                    if tempo_bpm > 125:
                        avg_logits[fast_idx] += 5.0 
                    elif tempo_bpm > 90:
                        avg_logits[mod_idx] += 3.0
                except ValueError:
                    pass
            
            probs = torch.nn.functional.softmax(avg_logits, dim=0).cpu().numpy()
            classes = le.classes_
            
            pred_dict = {str(classes[i]): float(probs[i]) for i in range(len(classes))}
            sorted_preds = dict(sorted(pred_dict.items(), key=lambda item: item[1], reverse=True)[:3])
            final_results.append(sorted_preds)
            
        return final_results
        
    except Exception as e:
        return [{"Error": f"Processing failed: {str(e)}"}] * 5

# ── GRADIO UI ──
with gr.Blocks(theme=gr.themes.Monochrome()) as demo:
    gr.Markdown("# Neural Math Rock Multimodal Analysis")
    gr.Markdown("Analyzing MBTI, Emotion, and Vibes using Hybrid Deep Learning & DSP.")
    
    with gr.Row():
        with gr.Column():
            audio_input = gr.Audio(type="filepath", label="Upload Song (Full Analysis)")
            lyrics_input = gr.Textbox(lines=5, label="Lyrics Content", placeholder="Paste lyrics here...")
            btn = gr.Button("START HYBRID ANALYSIS", variant="primary")
        
        with gr.Column():
            out_mbti = gr.Label(label="Personality (MBTI)")
            out_emotion = gr.Label(label="Emotional State")
            out_vibe = gr.Label(label="Acoustic Vibe")
            out_intensity = gr.Label(label="Intensity Level")
            out_tempo = gr.Label(label="Tempo Classification (BPM Adjusted)")

    btn.click(
        fn=predict_multimodal,
        inputs=[audio_input, lyrics_input],
        outputs=[out_mbti, out_emotion, out_vibe, out_intensity, out_tempo]
    )

if __name__ == "__main__":
    demo.launch()