2d-motion-interface / adapters.py
KanameYOkoYAMA's picture
Deploy 2D Motion Interface demo
0cdc216 verified
Raw
History Blame Contribute Delete
2.14 kB
"""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}")