Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -45,14 +45,14 @@ class AudioMathRockModel(nn.Module):
|
|
| 45 |
self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB(stype='power', top_db=80.0)
|
| 46 |
|
| 47 |
self.cnn_extractor = nn.Sequential(
|
| 48 |
-
nn.
|
| 49 |
nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2, 2),
|
| 50 |
-
nn.
|
| 51 |
nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2, 2),
|
| 52 |
-
nn.
|
| 53 |
nn.BatchNorm2d(128), nn.ReLU(), nn.AdaptiveAvgPool2d((1, 1)),
|
| 54 |
nn.Flatten(),
|
| 55 |
-
nn.
|
| 56 |
)
|
| 57 |
|
| 58 |
self.wlm_proj = nn.Linear(768, 512)
|
|
@@ -72,7 +72,7 @@ class AudioMathRockModel(nn.Module):
|
|
| 72 |
)
|
| 73 |
self.vibe_head = nn.Linear(512, len(VIBE_LABELS))
|
| 74 |
self.int_head = nn.Linear(512, len(INTENSITY_LABELS))
|
| 75 |
-
self.tmp_head = nn.Linear(512, len(
|
| 76 |
|
| 77 |
def forward(self, wavlm_values, clap_values):
|
| 78 |
wavlm_feats = self.wavlm(wavlm_values).last_hidden_state.mean(dim=1)
|
|
|
|
| 45 |
self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB(stype='power', top_db=80.0)
|
| 46 |
|
| 47 |
self.cnn_extractor = nn.Sequential(
|
| 48 |
+
nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1),
|
| 49 |
nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2, 2),
|
| 50 |
+
nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),
|
| 51 |
nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2, 2),
|
| 52 |
+
nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),
|
| 53 |
nn.BatchNorm2d(128), nn.ReLU(), nn.AdaptiveAvgPool2d((1, 1)),
|
| 54 |
nn.Flatten(),
|
| 55 |
+
nn.Linear(128, 512)
|
| 56 |
)
|
| 57 |
|
| 58 |
self.wlm_proj = nn.Linear(768, 512)
|
|
|
|
| 72 |
)
|
| 73 |
self.vibe_head = nn.Linear(512, len(VIBE_LABELS))
|
| 74 |
self.int_head = nn.Linear(512, len(INTENSITY_LABELS))
|
| 75 |
+
self.tmp_head = nn.Linear(512, len(TEMPO_LABELS))
|
| 76 |
|
| 77 |
def forward(self, wavlm_values, clap_values):
|
| 78 |
wavlm_feats = self.wavlm(wavlm_values).last_hidden_state.mean(dim=1)
|