Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -40,6 +40,9 @@ class HybridMultimodalModel(nn.Module):
|
|
| 40 |
|
| 41 |
self.audio_proj = nn.Linear(768, 256)
|
| 42 |
|
|
|
|
|
|
|
|
|
|
| 43 |
self.fusion = nn.Sequential(
|
| 44 |
nn.Linear(1024, 512),
|
| 45 |
nn.BatchNorm1d(512),
|
|
@@ -60,7 +63,10 @@ class HybridMultimodalModel(nn.Module):
|
|
| 60 |
audio_out = self.audio_model(audio_values).last_hidden_state
|
| 61 |
audio_feat = self.audio_proj(audio_out.mean(dim=1))
|
| 62 |
|
| 63 |
-
|
|
|
|
|
|
|
|
|
|
| 64 |
|
| 65 |
return {
|
| 66 |
'mbti': self.head_mbti(fused),
|
|
|
|
| 40 |
|
| 41 |
self.audio_proj = nn.Linear(768, 256)
|
| 42 |
|
| 43 |
+
self.text_gate = nn.Sequential(nn.Linear(768, 768), nn.Sigmoid())
|
| 44 |
+
self.audio_gate = nn.Sequential(nn.Linear(256, 256), nn.Sigmoid())
|
| 45 |
+
|
| 46 |
self.fusion = nn.Sequential(
|
| 47 |
nn.Linear(1024, 512),
|
| 48 |
nn.BatchNorm1d(512),
|
|
|
|
| 63 |
audio_out = self.audio_model(audio_values).last_hidden_state
|
| 64 |
audio_feat = self.audio_proj(audio_out.mean(dim=1))
|
| 65 |
|
| 66 |
+
gated_text = text_feat * self.text_gate(text_feat)
|
| 67 |
+
gated_audio = audio_feat * self.audio_gate(audio_feat)
|
| 68 |
+
|
| 69 |
+
fused = self.fusion(torch.cat([gated_text, gated_audio], dim=-1))
|
| 70 |
|
| 71 |
return {
|
| 72 |
'mbti': self.head_mbti(fused),
|