Spaces:
Running on Zero
Running on Zero
| """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}") | |