"""Adapter definitions, copied from evaluate_with_realworld.py so the Space does not depend on the evaluation script.""" import torch.nn as nn class AdapterLinear(nn.Module): def __init__(self, dim=81): super().__init__() self.norm = nn.LayerNorm(dim) self.linear = nn.Linear(dim, dim) nn.init.zeros_(self.linear.weight) nn.init.zeros_(self.linear.bias) def forward(self, x): return x + self.linear(self.norm(x)) class AdapterResidual(nn.Module): def __init__(self, dim=81, hidden=512, dropout=0.1): super().__init__() self.net = nn.Sequential( nn.LayerNorm(dim), nn.Linear(dim, hidden), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden, dim), ) nn.init.zeros_(self.net[-1].weight) nn.init.zeros_(self.net[-1].bias) def forward(self, x): return x + self.net(x) class AdapterConv1d(nn.Module): def __init__(self, dim=81, hidden=512, kernel_size=3, dropout=0.1): super().__init__() self.norm = nn.LayerNorm(dim) self.conv1 = nn.Conv1d(dim, hidden, kernel_size=kernel_size, padding=kernel_size // 2) self.act = nn.GELU() self.drop = nn.Dropout(dropout) self.conv2 = nn.Conv1d(hidden, dim, kernel_size=kernel_size, padding=kernel_size // 2) nn.init.zeros_(self.conv2.weight) nn.init.zeros_(self.conv2.bias) def forward(self, x): # x: (B, T, dim) r = self.norm(x) r = r.permute(0, 2, 1) r = self.conv1(r) r = self.act(r) r = self.drop(r) r = self.conv2(r) r = r.permute(0, 2, 1) return x + r def build_adapter(adapter_type, dim=81, hidden=512, kernel_size=3): if adapter_type == "linear": return AdapterLinear(dim=dim) elif adapter_type == "residual": return AdapterResidual(dim=dim, hidden=hidden) elif adapter_type == "conv1d": return AdapterConv1d(dim=dim, hidden=hidden, kernel_size=kernel_size) else: raise ValueError(f"Unknown adapter_type: {adapter_type}")