anggars commited on
Commit
dd98f8b
·
verified ·
1 Parent(s): dd5c062

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +5 -5
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.KeepConv = nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1),
49
  nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2, 2),
50
- nn.KeepConv2d = nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),
51
  nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2, 2),
52
- nn.KeepConv3d = 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.KeepLinear = nn.Linear(128, 512)
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(TEMPO_CLASSES))
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)