declare-lab/tango2
92
1# Copyright 2023 The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7# http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14import math15 16import torch17import torch.nn.functional as F18from torch import nn19 20from .resnet import Downsample1D, ResidualTemporalBlock1D, Upsample1D, rearrange_dims21 22 23class DownResnetBlock1D(nn.Module):24 def __init__(25 self,26 in_channels,27 out_channels=None,28 num_layers=1,29 conv_shortcut=False,30 temb_channels=32,31 groups=32,32 groups_out=None,33 non_linearity=None,34 time_embedding_norm="default",35 output_scale_factor=1.0,36 add_downsample=True,37 ):38 super().__init__()39 self.in_channels = in_channels40 out_channels = in_channels if out_channels is None else out_channels41 self.out_channels = out_channels42 self.use_conv_shortcut = conv_shortcut43 self.time_embedding_norm = time_embedding_norm44 self.add_downsample = add_downsample45 self.output_scale_factor = output_scale_factor46 47 if groups_out is None:48 groups_out = groups49 50 # there will always be at least one resnet51 resnets = [ResidualTemporalBlock1D(in_channels, out_channels, embed_dim=temb_channels)]52 53 for _ in range(num_layers):54 resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=temb_channels))55 56 self.resnets = nn.ModuleList(resnets)57 58 if non_linearity == "swish":59 self.nonlinearity = lambda x: F.silu(x)60 elif non_linearity == "mish":61 self.nonlinearity = nn.Mish()62 elif non_linearity == "silu":63 self.nonlinearity = nn.SiLU()64 else:65 self.nonlinearity = None66 67 self.downsample = None68 if add_downsample:69 self.downsample = Downsample1D(out_channels, use_conv=True, padding=1)70 71 def forward(self, hidden_states, temb=None):72 output_states = ()73 74 hidden_states = self.resnets[0](hidden_states, temb)75 for resnet in self.resnets[1:]:76 hidden_states = resnet(hidden_states, temb)77 78 output_states += (hidden_states,)79 80 if self.nonlinearity is not None:81 hidden_states = self.nonlinearity(hidden_states)82 83 if self.downsample is not None:84 hidden_states = self.downsample(hidden_states)85 86 return hidden_states, output_states87 88 89class UpResnetBlock1D(nn.Module):90 def __init__(91 self,92 in_channels,93 out_channels=None,94 num_layers=1,95 temb_channels=32,96 groups=32,97 groups_out=None,98 non_linearity=None,99 time_embedding_norm="default",100 output_scale_factor=1.0,101 add_upsample=True,102 ):103 super().__init__()104 self.in_channels = in_channels105 out_channels = in_channels if out_channels is None else out_channels106 self.out_channels = out_channels107 self.time_embedding_norm = time_embedding_norm108 self.add_upsample = add_upsample109 self.output_scale_factor = output_scale_factor110 111 if groups_out is None:112 groups_out = groups113 114 # there will always be at least one resnet115 resnets = [ResidualTemporalBlock1D(2 * in_channels, out_channels, embed_dim=temb_channels)]116 117 for _ in range(num_layers):118 resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=temb_channels))119 120 self.resnets = nn.ModuleList(resnets)121 122 if non_linearity == "swish":123 self.nonlinearity = lambda x: F.silu(x)124 elif non_linearity == "mish":125 self.nonlinearity = nn.Mish()126 elif non_linearity == "silu":127 self.nonlinearity = nn.SiLU()128 else:129 self.nonlinearity = None130 131 self.upsample = None132 if add_upsample:133 self.upsample = Upsample1D(out_channels, use_conv_transpose=True)134 135 def forward(self, hidden_states, res_hidden_states_tuple=None, temb=None):136 if res_hidden_states_tuple is not None:137 res_hidden_states = res_hidden_states_tuple[-1]138 hidden_states = torch.cat((hidden_states, res_hidden_states), dim=1)139 140 hidden_states = self.resnets[0](hidden_states, temb)141 for resnet in self.resnets[1:]:142 hidden_states = resnet(hidden_states, temb)143 144 if self.nonlinearity is not None:145 hidden_states = self.nonlinearity(hidden_states)146 147 if self.upsample is not None:148 hidden_states = self.upsample(hidden_states)149 150 return hidden_states151 152 153class ValueFunctionMidBlock1D(nn.Module):154 def __init__(self, in_channels, out_channels, embed_dim):155 super().__init__()156 self.in_channels = in_channels157 self.out_channels = out_channels158 self.embed_dim = embed_dim159 160 self.res1 = ResidualTemporalBlock1D(in_channels, in_channels // 2, embed_dim=embed_dim)161 self.down1 = Downsample1D(out_channels // 2, use_conv=True)162 self.res2 = ResidualTemporalBlock1D(in_channels // 2, in_channels // 4, embed_dim=embed_dim)163 self.down2 = Downsample1D(out_channels // 4, use_conv=True)164 165 def forward(self, x, temb=None):166 x = self.res1(x, temb)167 x = self.down1(x)168 x = self.res2(x, temb)169 x = self.down2(x)170 return x171 172 173class MidResTemporalBlock1D(nn.Module):174 def __init__(175 self,176 in_channels,177 out_channels,178 embed_dim,179 num_layers: int = 1,180 add_downsample: bool = False,181 add_upsample: bool = False,182 non_linearity=None,183 ):184 super().__init__()185 self.in_channels = in_channels186 self.out_channels = out_channels187 self.add_downsample = add_downsample188 189 # there will always be at least one resnet190 resnets = [ResidualTemporalBlock1D(in_channels, out_channels, embed_dim=embed_dim)]191 192 for _ in range(num_layers):193 resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=embed_dim))194 195 self.resnets = nn.ModuleList(resnets)196 197 if non_linearity == "swish":198 self.nonlinearity = lambda x: F.silu(x)199 elif non_linearity == "mish":200 self.nonlinearity = nn.Mish()201 elif non_linearity == "silu":202 self.nonlinearity = nn.SiLU()203 else:204 self.nonlinearity = None205 206 self.upsample = None207 if add_upsample:208 self.upsample = Downsample1D(out_channels, use_conv=True)209 210 self.downsample = None211 if add_downsample:212 self.downsample = Downsample1D(out_channels, use_conv=True)213 214 if self.upsample and self.downsample:215 raise ValueError("Block cannot downsample and upsample")216 217 def forward(self, hidden_states, temb):218 hidden_states = self.resnets[0](hidden_states, temb)219 for resnet in self.resnets[1:]:220 hidden_states = resnet(hidden_states, temb)221 222 if self.upsample:223 hidden_states = self.upsample(hidden_states)224 if self.downsample:225 self.downsample = self.downsample(hidden_states)226 227 return hidden_states228 229 230class OutConv1DBlock(nn.Module):231 def __init__(self, num_groups_out, out_channels, embed_dim, act_fn):232 super().__init__()233 self.final_conv1d_1 = nn.Conv1d(embed_dim, embed_dim, 5, padding=2)234 self.final_conv1d_gn = nn.GroupNorm(num_groups_out, embed_dim)235 if act_fn == "silu":236 self.final_conv1d_act = nn.SiLU()237 if act_fn == "mish":238 self.final_conv1d_act = nn.Mish()239 self.final_conv1d_2 = nn.Conv1d(embed_dim, out_channels, 1)240 241 def forward(self, hidden_states, temb=None):242 hidden_states = self.final_conv1d_1(hidden_states)243 hidden_states = rearrange_dims(hidden_states)244 hidden_states = self.final_conv1d_gn(hidden_states)245 hidden_states = rearrange_dims(hidden_states)246 hidden_states = self.final_conv1d_act(hidden_states)247 hidden_states = self.final_conv1d_2(hidden_states)248 return hidden_states249 250 251class OutValueFunctionBlock(nn.Module):252 def __init__(self, fc_dim, embed_dim):253 super().__init__()254 self.final_block = nn.ModuleList(255 [256 nn.Linear(fc_dim + embed_dim, fc_dim // 2),257 nn.Mish(),258 nn.Linear(fc_dim // 2, 1),259 ]260 )261 262 def forward(self, hidden_states, temb):263 hidden_states = hidden_states.view(hidden_states.shape[0], -1)264 hidden_states = torch.cat((hidden_states, temb), dim=-1)265 for layer in self.final_block:266 hidden_states = layer(hidden_states)267 268 return hidden_states269 270 271_kernels = {272 "linear": [1 / 8, 3 / 8, 3 / 8, 1 / 8],273 "cubic": [-0.01171875, -0.03515625, 0.11328125, 0.43359375, 0.43359375, 0.11328125, -0.03515625, -0.01171875],274 "lanczos3": [275 0.003689131001010537,276 0.015056144446134567,277 -0.03399861603975296,278 -0.066637322306633,279 0.13550527393817902,280 0.44638532400131226,281 0.44638532400131226,282 0.13550527393817902,283 -0.066637322306633,284 -0.03399861603975296,285 0.015056144446134567,286 0.003689131001010537,287 ],288}289 290 291class Downsample1d(nn.Module):292 def __init__(self, kernel="linear", pad_mode="reflect"):293 super().__init__()294 self.pad_mode = pad_mode295 kernel_1d = torch.tensor(_kernels[kernel])296 self.pad = kernel_1d.shape[0] // 2 - 1297 self.register_buffer("kernel", kernel_1d)298 299 def forward(self, hidden_states):300 hidden_states = F.pad(hidden_states, (self.pad,) * 2, self.pad_mode)301 weight = hidden_states.new_zeros([hidden_states.shape[1], hidden_states.shape[1], self.kernel.shape[0]])302 indices = torch.arange(hidden_states.shape[1], device=hidden_states.device)303 weight[indices, indices] = self.kernel.to(weight)304 return F.conv1d(hidden_states, weight, stride=2)305 306 307class Upsample1d(nn.Module):308 def __init__(self, kernel="linear", pad_mode="reflect"):309 super().__init__()310 self.pad_mode = pad_mode311 kernel_1d = torch.tensor(_kernels[kernel]) * 2312 self.pad = kernel_1d.shape[0] // 2 - 1313 self.register_buffer("kernel", kernel_1d)314 315 def forward(self, hidden_states, temb=None):316 hidden_states = F.pad(hidden_states, ((self.pad + 1) // 2,) * 2, self.pad_mode)317 weight = hidden_states.new_zeros([hidden_states.shape[1], hidden_states.shape[1], self.kernel.shape[0]])318 indices = torch.arange(hidden_states.shape[1], device=hidden_states.device)319 weight[indices, indices] = self.kernel.to(weight)320 return F.conv_transpose1d(hidden_states, weight, stride=2, padding=self.pad * 2 + 1)321 322 323class SelfAttention1d(nn.Module):324 def __init__(self, in_channels, n_head=1, dropout_rate=0.0):325 super().__init__()326 self.channels = in_channels327 self.group_norm = nn.GroupNorm(1, num_channels=in_channels)328 self.num_heads = n_head329 330 self.query = nn.Linear(self.channels, self.channels)331 self.key = nn.Linear(self.channels, self.channels)332 self.value = nn.Linear(self.channels, self.channels)333 334 self.proj_attn = nn.Linear(self.channels, self.channels, bias=True)335 336 self.dropout = nn.Dropout(dropout_rate, inplace=True)337 338 def transpose_for_scores(self, projection: torch.Tensor) -> torch.Tensor:339 new_projection_shape = projection.size()[:-1] + (self.num_heads, -1)340 # move heads to 2nd position (B, T, H * D) -> (B, T, H, D) -> (B, H, T, D)341 new_projection = projection.view(new_projection_shape).permute(0, 2, 1, 3)342 return new_projection343 344 def forward(self, hidden_states):345 residual = hidden_states346 batch, channel_dim, seq = hidden_states.shape347 348 hidden_states = self.group_norm(hidden_states)349 hidden_states = hidden_states.transpose(1, 2)350 351 query_proj = self.query(hidden_states)352 key_proj = self.key(hidden_states)353 value_proj = self.value(hidden_states)354 355 query_states = self.transpose_for_scores(query_proj)356 key_states = self.transpose_for_scores(key_proj)357 value_states = self.transpose_for_scores(value_proj)358 359 scale = 1 / math.sqrt(math.sqrt(key_states.shape[-1]))360 361 attention_scores = torch.matmul(query_states * scale, key_states.transpose(-1, -2) * scale)362 attention_probs = torch.softmax(attention_scores, dim=-1)363 364 # compute attention output365 hidden_states = torch.matmul(attention_probs, value_states)366 367 hidden_states = hidden_states.permute(0, 2, 1, 3).contiguous()368 new_hidden_states_shape = hidden_states.size()[:-2] + (self.channels,)369 hidden_states = hidden_states.view(new_hidden_states_shape)370 371 # compute next hidden_states372 hidden_states = self.proj_attn(hidden_states)373 hidden_states = hidden_states.transpose(1, 2)374 hidden_states = self.dropout(hidden_states)375 376 output = hidden_states + residual377 378 return output379 380 381class ResConvBlock(nn.Module):382 def __init__(self, in_channels, mid_channels, out_channels, is_last=False):383 super().__init__()384 self.is_last = is_last385 self.has_conv_skip = in_channels != out_channels386 387 if self.has_conv_skip:388 self.conv_skip = nn.Conv1d(in_channels, out_channels, 1, bias=False)389 390 self.conv_1 = nn.Conv1d(in_channels, mid_channels, 5, padding=2)391 self.group_norm_1 = nn.GroupNorm(1, mid_channels)392 self.gelu_1 = nn.GELU()393 self.conv_2 = nn.Conv1d(mid_channels, out_channels, 5, padding=2)394 395 if not self.is_last:396 self.group_norm_2 = nn.GroupNorm(1, out_channels)397 self.gelu_2 = nn.GELU()398 399 def forward(self, hidden_states):400 residual = self.conv_skip(hidden_states) if self.has_conv_skip else hidden_states401 402 hidden_states = self.conv_1(hidden_states)403 hidden_states = self.group_norm_1(hidden_states)404 hidden_states = self.gelu_1(hidden_states)405 hidden_states = self.conv_2(hidden_states)406 407 if not self.is_last:408 hidden_states = self.group_norm_2(hidden_states)409 hidden_states = self.gelu_2(hidden_states)410 411 output = hidden_states + residual412 return output413 414 415class UNetMidBlock1D(nn.Module):416 def __init__(self, mid_channels, in_channels, out_channels=None):417 super().__init__()418 419 out_channels = in_channels if out_channels is None else out_channels420 421 # there is always at least one resnet422 self.down = Downsample1d("cubic")423 resnets = [424 ResConvBlock(in_channels, mid_channels, mid_channels),425 ResConvBlock(mid_channels, mid_channels, mid_channels),426 ResConvBlock(mid_channels, mid_channels, mid_channels),427 ResConvBlock(mid_channels, mid_channels, mid_channels),428 ResConvBlock(mid_channels, mid_channels, mid_channels),429 ResConvBlock(mid_channels, mid_channels, out_channels),430 ]431 attentions = [432 SelfAttention1d(mid_channels, mid_channels // 32),433 SelfAttention1d(mid_channels, mid_channels // 32),434 SelfAttention1d(mid_channels, mid_channels // 32),435 SelfAttention1d(mid_channels, mid_channels // 32),436 SelfAttention1d(mid_channels, mid_channels // 32),437 SelfAttention1d(out_channels, out_channels // 32),438 ]439 self.up = Upsample1d(kernel="cubic")440 441 self.attentions = nn.ModuleList(attentions)442 self.resnets = nn.ModuleList(resnets)443 444 def forward(self, hidden_states, temb=None):445 hidden_states = self.down(hidden_states)446 for attn, resnet in zip(self.attentions, self.resnets):447 hidden_states = resnet(hidden_states)448 hidden_states = attn(hidden_states)449 450 hidden_states = self.up(hidden_states)451 452 return hidden_states453 454 455class AttnDownBlock1D(nn.Module):456 def __init__(self, out_channels, in_channels, mid_channels=None):457 super().__init__()458 mid_channels = out_channels if mid_channels is None else mid_channels459 460 self.down = Downsample1d("cubic")461 resnets = [462 ResConvBlock(in_channels, mid_channels, mid_channels),463 ResConvBlock(mid_channels, mid_channels, mid_channels),464 ResConvBlock(mid_channels, mid_channels, out_channels),465 ]466 attentions = [467 SelfAttention1d(mid_channels, mid_channels // 32),468 SelfAttention1d(mid_channels, mid_channels // 32),469 SelfAttention1d(out_channels, out_channels // 32),470 ]471 472 self.attentions = nn.ModuleList(attentions)473 self.resnets = nn.ModuleList(resnets)474 475 def forward(self, hidden_states, temb=None):476 hidden_states = self.down(hidden_states)477 478 for resnet, attn in zip(self.resnets, self.attentions):479 hidden_states = resnet(hidden_states)480 hidden_states = attn(hidden_states)481 482 return hidden_states, (hidden_states,)483 484 485class DownBlock1D(nn.Module):486 def __init__(self, out_channels, in_channels, mid_channels=None):487 super().__init__()488 mid_channels = out_channels if mid_channels is None else mid_channels489 490 self.down = Downsample1d("cubic")491 resnets = [492 ResConvBlock(in_channels, mid_channels, mid_channels),493 ResConvBlock(mid_channels, mid_channels, mid_channels),494 ResConvBlock(mid_channels, mid_channels, out_channels),495 ]496 497 self.resnets = nn.ModuleList(resnets)498 499 def forward(self, hidden_states, temb=None):500 hidden_states = self.down(hidden_states)501 502 for resnet in self.resnets:503 hidden_states = resnet(hidden_states)504 505 return hidden_states, (hidden_states,)506 507 508class DownBlock1DNoSkip(nn.Module):509 def __init__(self, out_channels, in_channels, mid_channels=None):510 super().__init__()511 mid_channels = out_channels if mid_channels is None else mid_channels512 513 resnets = [514 ResConvBlock(in_channels, mid_channels, mid_channels),515 ResConvBlock(mid_channels, mid_channels, mid_channels),516 ResConvBlock(mid_channels, mid_channels, out_channels),517 ]518 519 self.resnets = nn.ModuleList(resnets)520 521 def forward(self, hidden_states, temb=None):522 hidden_states = torch.cat([hidden_states, temb], dim=1)523 for resnet in self.resnets:524 hidden_states = resnet(hidden_states)525 526 return hidden_states, (hidden_states,)527 528 529class AttnUpBlock1D(nn.Module):530 def __init__(self, in_channels, out_channels, mid_channels=None):531 super().__init__()532 mid_channels = out_channels if mid_channels is None else mid_channels533 534 resnets = [535 ResConvBlock(2 * in_channels, mid_channels, mid_channels),536 ResConvBlock(mid_channels, mid_channels, mid_channels),537 ResConvBlock(mid_channels, mid_channels, out_channels),538 ]539 attentions = [540 SelfAttention1d(mid_channels, mid_channels // 32),541 SelfAttention1d(mid_channels, mid_channels // 32),542 SelfAttention1d(out_channels, out_channels // 32),543 ]544 545 self.attentions = nn.ModuleList(attentions)546 self.resnets = nn.ModuleList(resnets)547 self.up = Upsample1d(kernel="cubic")548 549 def forward(self, hidden_states, res_hidden_states_tuple, temb=None):550 res_hidden_states = res_hidden_states_tuple[-1]551 hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)552 553 for resnet, attn in zip(self.resnets, self.attentions):554 hidden_states = resnet(hidden_states)555 hidden_states = attn(hidden_states)556 557 hidden_states = self.up(hidden_states)558 559 return hidden_states560 561 562class UpBlock1D(nn.Module):563 def __init__(self, in_channels, out_channels, mid_channels=None):564 super().__init__()565 mid_channels = in_channels if mid_channels is None else mid_channels566 567 resnets = [568 ResConvBlock(2 * in_channels, mid_channels, mid_channels),569 ResConvBlock(mid_channels, mid_channels, mid_channels),570 ResConvBlock(mid_channels, mid_channels, out_channels),571 ]572 573 self.resnets = nn.ModuleList(resnets)574 self.up = Upsample1d(kernel="cubic")575 576 def forward(self, hidden_states, res_hidden_states_tuple, temb=None):577 res_hidden_states = res_hidden_states_tuple[-1]578 hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)579 580 for resnet in self.resnets:581 hidden_states = resnet(hidden_states)582 583 hidden_states = self.up(hidden_states)584 585 return hidden_states586 587 588class UpBlock1DNoSkip(nn.Module):589 def __init__(self, in_channels, out_channels, mid_channels=None):590 super().__init__()591 mid_channels = in_channels if mid_channels is None else mid_channels592 593 resnets = [594 ResConvBlock(2 * in_channels, mid_channels, mid_channels),595 ResConvBlock(mid_channels, mid_channels, mid_channels),596 ResConvBlock(mid_channels, mid_channels, out_channels, is_last=True),597 ]598 599 self.resnets = nn.ModuleList(resnets)600 601 def forward(self, hidden_states, res_hidden_states_tuple, temb=None):602 res_hidden_states = res_hidden_states_tuple[-1]603 hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)604 605 for resnet in self.resnets:606 hidden_states = resnet(hidden_states)607 608 return hidden_states609 610 611def get_down_block(down_block_type, num_layers, in_channels, out_channels, temb_channels, add_downsample):612 if down_block_type == "DownResnetBlock1D":613 return DownResnetBlock1D(614 in_channels=in_channels,615 num_layers=num_layers,616 out_channels=out_channels,617 temb_channels=temb_channels,618 add_downsample=add_downsample,619 )620 elif down_block_type == "DownBlock1D":621 return DownBlock1D(out_channels=out_channels, in_channels=in_channels)622 elif down_block_type == "AttnDownBlock1D":623 return AttnDownBlock1D(out_channels=out_channels, in_channels=in_channels)624 elif down_block_type == "DownBlock1DNoSkip":625 return DownBlock1DNoSkip(out_channels=out_channels, in_channels=in_channels)626 raise ValueError(f"{down_block_type} does not exist.")627 628 629def get_up_block(up_block_type, num_layers, in_channels, out_channels, temb_channels, add_upsample):630 if up_block_type == "UpResnetBlock1D":631 return UpResnetBlock1D(632 in_channels=in_channels,633 num_layers=num_layers,634 out_channels=out_channels,635 temb_channels=temb_channels,636 add_upsample=add_upsample,637 )638 elif up_block_type == "UpBlock1D":639 return UpBlock1D(in_channels=in_channels, out_channels=out_channels)640 elif up_block_type == "AttnUpBlock1D":641 return AttnUpBlock1D(in_channels=in_channels, out_channels=out_channels)642 elif up_block_type == "UpBlock1DNoSkip":643 return UpBlock1DNoSkip(in_channels=in_channels, out_channels=out_channels)644 raise ValueError(f"{up_block_type} does not exist.")645 646 647def get_mid_block(mid_block_type, num_layers, in_channels, mid_channels, out_channels, embed_dim, add_downsample):648 if mid_block_type == "MidResTemporalBlock1D":649 return MidResTemporalBlock1D(650 num_layers=num_layers,651 in_channels=in_channels,652 out_channels=out_channels,653 embed_dim=embed_dim,654 add_downsample=add_downsample,655 )656 elif mid_block_type == "ValueFunctionMidBlock1D":657 return ValueFunctionMidBlock1D(in_channels=in_channels, out_channels=out_channels, embed_dim=embed_dim)658 elif mid_block_type == "UNetMidBlock1D":659 return UNetMidBlock1D(in_channels=in_channels, mid_channels=mid_channels, out_channels=out_channels)660 raise ValueError(f"{mid_block_type} does not exist.")661 662 663def get_out_block(*, out_block_type, num_groups_out, embed_dim, out_channels, act_fn, fc_dim):664 if out_block_type == "OutConv1DBlock":665 return OutConv1DBlock(num_groups_out, out_channels, embed_dim, act_fn)666 elif out_block_type == "ValueFunction":667 return OutValueFunctionBlock(fc_dim, embed_dim)668 return None669 