ChristophSchuhmann's picture
Release best Gemini-tuned Whisper Base and Small with code, normalization and evaluation
cd9b2d8 verified
Raw History Blame Contribute Delete
2.88 kB
"""Matched small heads on cached frozen pooled and temporal features."""
import torch
from torch import nn
from torch.nn import functional as F
class ScalarHeads(nn.Module):
def __init__(self, linear=False):
super().__init__()
self.linear = linear
if linear:
self.output = nn.Linear(256, 193)
else:
self.first = nn.Parameter(torch.empty(193, 256, 64))
self.bias = nn.Parameter(torch.zeros(193, 64))
self.second = nn.Parameter(torch.empty(193, 64))
self.final_bias = nn.Parameter(torch.zeros(193))
nn.init.normal_(self.first, std=.02)
nn.init.normal_(self.second, std=.02)
self.dropout = nn.Dropout(.1)
def forward(self, features):
if self.linear:
return self.output(features)
hidden = self.dropout(F.gelu(torch.einsum('bi,sih->bsh', features, self.first) + self.bias))
return (hidden * self.second[None]).sum(-1) + self.final_bias
class Probe(nn.Module):
def __init__(self, n_classes, linear=False):
super().__init__()
self.scalars = ScalarHeads(linear)
def head(d, out):
if linear:
return nn.Linear(d, out) if (d + 1) * out <= 50000 else nn.Sequential(nn.Linear(d, 64), nn.Linear(64, out))
return nn.Sequential(nn.Linear(d, 64), nn.GELU(), nn.Dropout(.1), nn.Linear(64, out))
self.timbre_head, self.identity_head = head(256, 128), head(256, 250)
self.frame_head = head(64, 3)
self.event_head = head(64, n_classes)
for name, module in [('timbre', self.timbre_head), ('identity', self.identity_head), ('frame', self.frame_head), ('event', self.event_head)]:
if sum(p.numel() for p in module.parameters()) > 50000:
raise ValueError('Head parameter cap exceeded: ' + name)
def forward(self, features, frame_features, starts, ends):
scalar = self.scalars(features)
temporal = self.frame_head(frame_features)
# Prefix sums pool only each event's frames without per-event CUDA sync.
width = frame_features.shape[1]
a = starts.clamp(0, width - 1)
b = torch.maximum(ends, a + 1).clamp(max=width)
prefix = F.pad(frame_features.cumsum(1), (0, 0, 1, 0))
ai = a[..., None].expand(-1, -1, frame_features.shape[-1])
bi = b[..., None].expand_as(ai)
event = (prefix.gather(1, bi) - prefix.gather(1, ai)) / (b - a)[..., None]
return {'scores': scalar[:, :192], 'cps': scalar[:, 192],
'timbre': F.normalize(self.timbre_head(features), dim=-1),
'identity': F.normalize(self.identity_head(features), dim=-1),
'frame': temporal[:, :, 0], 'onset': temporal[:, :, 1], 'log_duration': temporal[:, :, 2],
'event_class': self.event_head(event)}