Spaces:
Running on Zero
Running on Zero
Download consistency_decoder.py from ddeshler/Neural_Music_Generator: direct link, hf CLI and curl.
- Browser
- Download file 34.1 kB
-
https://huggingface.co/spaces/ddeshler/Neural_Music_Generator/resolve/main/consistency_decoder.py
- Command line
-
hf download hf://spaces/ddeshler/Neural_Music_Generator/consistency_decoder.py
-
curl -L -o consistency_decoder.py https://huggingface.co/spaces/ddeshler/Neural_Music_Generator/resolve/main/consistency_decoder.py
34.1 kB
| # https://gist.github.com/mrsteyk/74ad3ec2f6f823111ae4c90e168505ac | |
| import torch | |
| import torch.nn.functional as F | |
| import torch.nn as nn | |
| class TimestepEmbedding(nn.Module): | |
| def __init__(self, n_time=1024, n_emb=320, n_out=1280) -> None: | |
| super().__init__() | |
| self.emb = nn.Embedding(n_time, n_emb) | |
| self.f_1 = nn.Linear(n_emb, n_out) | |
| self.f_2 = nn.Linear(n_out, n_out) | |
| def forward(self, x) -> torch.Tensor: | |
| x = self.emb(x) | |
| x = self.f_1(x) | |
| x = F.silu(x) | |
| return self.f_2(x) | |
| class PositionalEmbedding(nn.Module): | |
| def __init__(self, pe_dim=320, out_dim=1280, max_positions=10000, endpoint=True): | |
| super().__init__() | |
| self.num_channels = pe_dim | |
| self.max_positions = max_positions | |
| self.endpoint = endpoint | |
| self.f_1 = nn.Linear(pe_dim, out_dim) | |
| self.f_2 = nn.Linear(out_dim, out_dim) | |
| def forward(self, x): | |
| freqs = torch.arange(start=0, end=self.num_channels//2, dtype=torch.float32, device=x.device) | |
| freqs = freqs / (self.num_channels // 2 - (1 if self.endpoint else 0)) | |
| freqs = (1 / self.max_positions) ** freqs | |
| x = x.ger(freqs.to(x.dtype)) | |
| x = torch.cat([x.cos(), x.sin()], dim=1) | |
| x = self.f_1(x) | |
| x = F.silu(x) | |
| return self.f_2(x) | |
| class ImageEmbedding(nn.Module): | |
| def __init__(self, in_channels, out_channels=320) -> None: | |
| super().__init__() | |
| self.f = nn.Conv1d(in_channels, out_channels, kernel_size=3, padding=1) | |
| def forward(self, x) -> torch.Tensor: | |
| return self.f(x) | |
| class ImageUnembedding(nn.Module): | |
| def __init__(self, in_channels=320, out_channels=3) -> None: | |
| super().__init__() | |
| self.gn = nn.GroupNorm(32, in_channels) | |
| self.f = nn.Conv1d(in_channels, out_channels, kernel_size=3, padding=1) | |
| def forward(self, x) -> torch.Tensor: | |
| return self.f(F.silu(self.gn(x))) | |
| class ConvResblock(nn.Module): | |
| def __init__(self, in_features, out_features, t_dim) -> None: | |
| super().__init__() | |
| self.f_t = nn.Linear(t_dim, out_features * 2) | |
| self.gn_1 = nn.GroupNorm(32, in_features) | |
| self.f_1 = nn.Conv1d(in_features, out_features, kernel_size=3, padding=1) | |
| self.gn_2 = nn.GroupNorm(32, out_features) | |
| self.f_2 = nn.Conv1d(out_features, out_features, kernel_size=3, padding=1) | |
| skip_conv = in_features != out_features | |
| self.f_s = ( | |
| nn.Conv1d(in_features, out_features, kernel_size=1, padding=0) | |
| if skip_conv | |
| else nn.Identity() | |
| ) | |
| def forward(self, x, t): | |
| x_skip = x | |
| t = self.f_t(F.silu(t)) | |
| t = t.chunk(2, dim=1) | |
| t_1 = t[0].unsqueeze(dim=2) + 1 | |
| t_2 = t[1].unsqueeze(dim=2) | |
| gn_1 = F.silu(self.gn_1(x)) | |
| f_1 = self.f_1(gn_1) | |
| gn_2 = self.gn_2(f_1) | |
| return self.f_s(x_skip) + self.f_2(F.silu(gn_2 * t_1 + t_2)) | |
| # Also ConvResblock | |
| class Downsample(nn.Module): | |
| def __init__(self, in_channels, t_dim, ratio=2) -> None: | |
| super().__init__() | |
| self.ratio = ratio | |
| self.f_t = nn.Linear(t_dim, in_channels * 2) | |
| self.gn_1 = nn.GroupNorm(32, in_channels) | |
| self.f_1 = nn.Conv1d(in_channels, in_channels, kernel_size=3, padding=1) | |
| self.gn_2 = nn.GroupNorm(32, in_channels) | |
| self.f_2 = nn.Conv1d(in_channels, in_channels, kernel_size=3, padding=1) | |
| def forward(self, x, t) -> torch.Tensor: | |
| x_skip = x | |
| t = self.f_t(F.silu(t)) | |
| t_1, t_2 = t.chunk(2, dim=1) | |
| t_1 = t_1.unsqueeze(2) + 1 | |
| t_2 = t_2.unsqueeze(2) | |
| gn_1 = F.silu(self.gn_1(x)) | |
| avg_pool1d = F.avg_pool1d(gn_1, kernel_size=self.ratio, stride=None) | |
| f_1 = self.f_1(avg_pool1d) | |
| gn_2 = self.gn_2(f_1) | |
| f_2 = self.f_2(F.silu(t_2 + (t_1 * gn_2))) | |
| return f_2 + F.avg_pool1d(x_skip, kernel_size=self.ratio, stride=None) | |
| # Also ConvResblock | |
| class Upsample(nn.Module): | |
| def __init__(self, in_channels, t_dim, ratio=2) -> None: | |
| super().__init__() | |
| self.ratio = ratio | |
| self.f_t = nn.Linear(t_dim, in_channels * 2) | |
| self.gn_1 = nn.GroupNorm(32, in_channels) | |
| self.f_1 = nn.Conv1d(in_channels, in_channels, kernel_size=3, padding=1) | |
| self.gn_2 = nn.GroupNorm(32, in_channels) | |
| self.f_2 = nn.Conv1d(in_channels, in_channels, kernel_size=3, padding=1) | |
| def forward(self, x, t) -> torch.Tensor: | |
| x_skip = x | |
| t = self.f_t(F.silu(t)) | |
| t_1, t_2 = t.chunk(2, dim=1) | |
| t_1 = t_1.unsqueeze(2) + 1 | |
| t_2 = t_2.unsqueeze(2) | |
| gn_1 = F.silu(self.gn_1(x)) | |
| upsample = F.interpolate(gn_1, scale_factor=self.ratio, mode='nearest') | |
| f_1 = self.f_1(upsample) | |
| gn_2 = self.gn_2(f_1) | |
| f_2 = self.f_2(F.silu(t_2 + (t_1 * gn_2))) | |
| return f_2 + F.interpolate(x_skip.float(), scale_factor=self.ratio, mode='nearest').to(x_skip.dtype) | |
| class ConsistencyDecoderUNet(nn.Module): | |
| def __init__(self, in_channels=3, z_dec_channels=None, c0=320, c1=640, c2=1024, pe_dim=320, t_dim=1280, ratios=[8, 5, 4]) -> None: | |
| super().__init__() | |
| if z_dec_channels is not None: | |
| in_channels += z_dec_channels | |
| self.embed_image = ImageEmbedding(in_channels=in_channels, out_channels=c0) | |
| self.embed_time = PositionalEmbedding(pe_dim=pe_dim, out_dim=t_dim) | |
| down_0 = nn.ModuleList([ | |
| ConvResblock(c0, c0, t_dim), | |
| ConvResblock(c0, c0, t_dim), | |
| ConvResblock(c0, c0, t_dim), | |
| Downsample(c0, t_dim, ratios[0]), | |
| ]) | |
| down_1 = nn.ModuleList([ | |
| ConvResblock(c0, c1, t_dim), | |
| ConvResblock(c1, c1, t_dim), | |
| ConvResblock(c1, c1, t_dim), | |
| Downsample(c1, t_dim, ratios[1]), | |
| ]) | |
| down_2 = nn.ModuleList([ | |
| ConvResblock(c1, c2, t_dim), | |
| ConvResblock(c2, c2, t_dim), | |
| ConvResblock(c2, c2, t_dim), | |
| Downsample(c2, t_dim, ratios[2]), | |
| ]) | |
| down_3 = nn.ModuleList([ | |
| ConvResblock(c2, c2, t_dim), | |
| ConvResblock(c2, c2, t_dim), | |
| ConvResblock(c2, c2, t_dim), | |
| ]) | |
| self.down = nn.ModuleList([ | |
| down_0, | |
| down_1, | |
| down_2, | |
| down_3, | |
| ]) | |
| self.mid = nn.ModuleList([ | |
| ConvResblock(c2, c2, t_dim), | |
| ConvResblock(c2, c2, t_dim), | |
| ]) | |
| up_3 = nn.ModuleList([ | |
| ConvResblock(c2 * 2, c2, t_dim), | |
| ConvResblock(c2 * 2, c2, t_dim), | |
| ConvResblock(c2 * 2, c2, t_dim), | |
| ConvResblock(c2 * 2, c2, t_dim), | |
| Upsample(c2, t_dim, ratios[2]), | |
| ]) | |
| up_2 = nn.ModuleList([ | |
| ConvResblock(c2 * 2, c2, t_dim), | |
| ConvResblock(c2 * 2, c2, t_dim), | |
| ConvResblock(c2 * 2, c2, t_dim), | |
| ConvResblock(c2 + c1, c2, t_dim), | |
| Upsample(c2, t_dim, ratios[1]), | |
| ]) | |
| up_1 = nn.ModuleList([ | |
| ConvResblock(c2 + c1, c1, t_dim), | |
| ConvResblock(c1 * 2, c1, t_dim), | |
| ConvResblock(c1 * 2, c1, t_dim), | |
| ConvResblock(c0 + c1, c1, t_dim), | |
| Upsample(c1, t_dim, ratios[0]), | |
| ]) | |
| up_0 = nn.ModuleList([ | |
| ConvResblock(c0 + c1, c0, t_dim), | |
| ConvResblock(c0 * 2, c0, t_dim), | |
| ConvResblock(c0 * 2, c0, t_dim), | |
| ConvResblock(c0 * 2, c0, t_dim), | |
| ]) | |
| self.up = nn.ModuleList([ | |
| up_0, | |
| up_1, | |
| up_2, | |
| up_3, | |
| ]) | |
| self.output = ImageUnembedding(in_channels=c0, out_channels=1) | |
| def get_last_layer_weight(self): | |
| return self.output.f.weight | |
| def initialize_weights(self): | |
| # Initialize transformer layers: | |
| def _basic_init(module): | |
| if isinstance(module, nn.Linear): | |
| torch.nn.init.xavier_uniform_(module.weight) | |
| if module.bias is not None: | |
| nn.init.constant_(module.bias, 0) | |
| self.apply(_basic_init) | |
| # Zero-out adaLN modulation layers in DiT blocks: | |
| # for block in self.blocks: | |
| # nn.init.constant_(block.adaLN_modulation[-1].weight, 0) | |
| # nn.init.constant_(block.adaLN_modulation[-1].bias, 0) | |
| # Zero-out output layers: | |
| nn.init.constant_(self.output.gn.weight, 0) | |
| nn.init.constant_(self.output.gn.bias, 0) | |
| nn.init.constant_(self.output.f.weight, 0) | |
| nn.init.constant_(self.output.f.bias, 0) | |
| def forward(self, x, t=None, z_dec=None) -> torch.Tensor: | |
| if z_dec is not None: | |
| # if z_dec.shape[-2] != x.shape[-2] or z_dec.shape[-1] != x.shape[-1]: | |
| if z_dec.shape[-1] != x.shape[-1]: | |
| # assert x.shape[-2] // z_dec.shape[-2] == x.shape[-1] // z_dec.shape[-1] | |
| # z_dec = F.upsample_nearest(z_dec, scale_factor=x.shape[-2] // z_dec.shape[-2]) | |
| z_dec = F.interpolate(z_dec, scale_factor=x.shape[-1] // z_dec.shape[-1], mode='nearest') | |
| x = torch.cat([x, z_dec], dim=1) | |
| x = self.embed_image(x) | |
| if t is None: | |
| t = torch.zeros(x.shape[0], device=x.device) | |
| t = self.embed_time(t) | |
| skips = [x] | |
| for down in self.down: | |
| for block in down: | |
| x = block(x, t) | |
| skips.append(x) | |
| for mid in self.mid: | |
| x = mid(x, t) | |
| for up in self.up[::-1]: | |
| for block in up: | |
| if isinstance(block, ConvResblock): | |
| x = torch.concat([x, skips.pop()], dim=1) | |
| x = block(x, t) | |
| return self.output(x) | |
| class PixelShuffle1D(torch.nn.Module): | |
| """ | |
| 1D pixel shuffler. https://arxiv.org/pdf/1609.05158.pdf | |
| Upscales sample length, downscales channel length | |
| "short" is input, "long" is output | |
| """ | |
| def __init__(self, upscale_factor): | |
| super(PixelShuffle1D, self).__init__() | |
| self.upscale_factor = upscale_factor | |
| def forward(self, x): | |
| batch_size = x.shape[0] | |
| short_channel_len = x.shape[1] | |
| short_width = x.shape[2] | |
| long_channel_len = short_channel_len // self.upscale_factor | |
| long_width = self.upscale_factor * short_width | |
| x = x.contiguous().view([batch_size, self.upscale_factor, long_channel_len, short_width]) | |
| x = x.permute(0, 2, 3, 1).contiguous() | |
| x = x.view(batch_size, long_channel_len, long_width) | |
| return x | |
| class PixelUnshuffle1D(torch.nn.Module): | |
| """ | |
| Inverse of 1D pixel shuffler | |
| Upscales channel length, downscales sample length | |
| "long" is input, "short" is output | |
| """ | |
| def __init__(self, downscale_factor): | |
| super(PixelUnshuffle1D, self).__init__() | |
| self.downscale_factor = downscale_factor | |
| def forward(self, x): | |
| batch_size = x.shape[0] | |
| long_channel_len = x.shape[1] | |
| long_width = x.shape[2] | |
| short_channel_len = long_channel_len * self.downscale_factor | |
| short_width = long_width // self.downscale_factor | |
| x = x.contiguous().view([batch_size, long_channel_len, short_width, self.downscale_factor]) | |
| x = x.permute(0, 3, 1, 2).contiguous() | |
| x = x.view([batch_size, short_channel_len, short_width]) | |
| return x | |
| class ConvPixelUnshuffleDownSampleLayer(nn.Module): | |
| def __init__( | |
| self, | |
| in_channels: int, | |
| out_channels: int, | |
| factor: int, | |
| kernel_size: int = 3 | |
| ): | |
| super().__init__() | |
| assert out_channels % factor == 0, f'{out_channels}, {factor}' | |
| # self.conv = ConvLayer( | |
| # in_channels=in_channels, | |
| # out_channels=out_channels // out_ratio, | |
| # kernel_size=kernel_size, | |
| # use_bias=True, | |
| # norm=None, | |
| # act_func=None, | |
| # ) | |
| self.norm = nn.GroupNorm(32, in_channels) | |
| self.conv = nn.Conv1d(in_channels, out_channels // factor, kernel_size=kernel_size, padding=kernel_size // 2) | |
| self.pixel_unshuffle = PixelUnshuffle1D(factor) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.norm(x) | |
| x = self.conv(x) | |
| x = self.pixel_unshuffle(x) | |
| return x | |
| class PixelUnshuffleChannelAveragingDownSampleLayer(nn.Module): | |
| def __init__( | |
| self, | |
| in_channels: int, | |
| out_channels: int, | |
| factor: int, | |
| ): | |
| super().__init__() | |
| self.in_channels = in_channels | |
| self.out_channels = out_channels | |
| assert in_channels * factor % out_channels == 0, f'{in_channels} {factor} {out_channels}' | |
| self.group_size = in_channels * factor // out_channels | |
| self.pixel_unshuffle = PixelUnshuffle1D(factor) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.pixel_unshuffle(x) | |
| B, C, L = x.shape | |
| x = x.view(B, self.out_channels, self.group_size, L) | |
| x = x.mean(dim=2) | |
| return x | |
| class DownsampleV3(nn.Module): | |
| def __init__(self, in_channels, out_channels, ratio): | |
| super().__init__() | |
| self.conv = ConvPixelUnshuffleDownSampleLayer(in_channels, out_channels, ratio, ratio * 2 + 1) | |
| self.shortcut = PixelUnshuffleChannelAveragingDownSampleLayer(in_channels, out_channels, ratio) | |
| def forward(self, x, t=None): | |
| x = self.conv(x) + self.shortcut(x) | |
| return x | |
| class DownsampleV2(nn.Module): | |
| def __init__(self, in_channels, out_channels, t_dim, ratio=2) -> None: | |
| super().__init__() | |
| self.ratio = ratio | |
| self.conv = ConvPixelUnshuffleDownSampleLayer(in_channels, in_channels, ratio) | |
| self.shortcut = PixelUnshuffleChannelAveragingDownSampleLayer(in_channels, in_channels, ratio) | |
| self.f_t = nn.Linear(t_dim, in_channels * 2) | |
| self.gn_1 = nn.GroupNorm(32, in_channels) | |
| self.f_1 = nn.Conv1d(in_channels, in_channels, kernel_size=3, padding=1) | |
| self.gn_2 = nn.GroupNorm(32, in_channels) | |
| self.f_2 = nn.Conv1d(in_channels, in_channels, kernel_size=3, padding=1) | |
| def forward(self, x, t) -> torch.Tensor: | |
| x_skip = x | |
| t = self.f_t(F.silu(t)) | |
| t_1, t_2 = t.chunk(2, dim=1) | |
| t_1 = t_1.unsqueeze(2) + 1 | |
| t_2 = t_2.unsqueeze(2) | |
| gn_1 = F.silu(self.gn_1(x)) | |
| avg_pool1d = self.conv(gn_1);print(gn_1.shape, avg_pool1d.shape) | |
| f_1 = self.f_1(avg_pool1d) | |
| gn_2 = self.gn_2(f_1) | |
| f_2 = self.f_2(F.silu(t_2 + (t_1 * gn_2))) | |
| return f_2 + self.shortcut(x_skip) | |
| class ConvPixelShuffleUpSampleLayer(nn.Module): | |
| def __init__( | |
| self, | |
| in_channels: int, | |
| out_channels: int, | |
| factor: int, | |
| kernel_size: int = 3 | |
| ): | |
| super().__init__() | |
| # self.conv = ConvLayer( | |
| # in_channels=in_channels, | |
| # out_channels=out_channels * out_ratio, | |
| # kernel_size=kernel_size, | |
| # use_bias=True, | |
| # norm=None, | |
| # act_func=None, | |
| # ) | |
| self.norm = nn.GroupNorm(32, in_channels) | |
| self.conv = nn.Conv1d(in_channels, out_channels * factor, kernel_size=kernel_size, padding=kernel_size // 2) | |
| self.pixel_shuffle = PixelShuffle1D(factor) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.norm(x) | |
| x = self.conv(x) | |
| x = self.pixel_shuffle(x) | |
| return x | |
| class ChannelDuplicatingPixelUnshuffleUpSampleLayer(nn.Module): | |
| def __init__( | |
| self, | |
| in_channels: int, | |
| out_channels: int, | |
| factor: int, | |
| ): | |
| super().__init__() | |
| self.in_channels = in_channels | |
| self.out_channels = out_channels | |
| assert out_channels * factor % in_channels == 0, f'{out_channels} {factor} {in_channels}' | |
| self.repeats = out_channels * factor // in_channels | |
| self.pixel_shuffle = PixelShuffle1D(factor) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = x.repeat_interleave(self.repeats, dim=1) | |
| x = self.pixel_shuffle(x) | |
| return x | |
| class UpsampleV3(nn.Module): | |
| def __init__(self, in_channels, out_channels, ratio): | |
| super().__init__() | |
| self.conv = ConvPixelShuffleUpSampleLayer(in_channels, out_channels, ratio, ratio * 2 + 1) | |
| self.shortcut = ChannelDuplicatingPixelUnshuffleUpSampleLayer(in_channels, out_channels, ratio) | |
| def forward(self, x, t=None): | |
| x = self.conv(x) + self.shortcut(x) | |
| return x | |
| class UpsampleV2(nn.Module): | |
| def __init__(self, in_channels, t_dim, ratio=2) -> None: | |
| super().__init__() | |
| self.ratio = ratio | |
| self.f_t = nn.Linear(t_dim, in_channels * 2) | |
| self.gn_1 = nn.GroupNorm(32, in_channels) | |
| self.f_1 = nn.Conv1d(in_channels, in_channels, kernel_size=3, padding=1) | |
| self.gn_2 = nn.GroupNorm(32, in_channels) | |
| self.f_2 = nn.Conv1d(in_channels, in_channels, kernel_size=3, padding=1) | |
| self.conv = ConvPixelShuffleUpSampleLayer(in_channels, in_channels, ratio) | |
| self.shortcut = ChannelDuplicatingPixelUnshuffleUpSampleLayer(in_channels, in_channels, ratio) | |
| def forward(self, x, t) -> torch.Tensor: | |
| x_skip = x | |
| t = self.f_t(F.silu(t)) | |
| t_1, t_2 = t.chunk(2, dim=1) | |
| t_1 = t_1.unsqueeze(2) + 1 | |
| t_2 = t_2.unsqueeze(2) | |
| gn_1 = F.silu(self.gn_1(x)) | |
| upsample = self.conv(gn_1) | |
| f_1 = self.f_1(upsample) | |
| gn_2 = self.gn_2(f_1) | |
| f_2 = self.f_2(F.silu(t_2 + (t_1 * gn_2))) | |
| return f_2 + self.shortcut(x_skip) | |
| def modulate(x, shift, scale): | |
| if scale.ndim == 3: | |
| return x * (1 + scale) + shift | |
| else: | |
| return x * (1 + scale.unsqueeze(-1)) + shift.unsqueeze(-1) | |
| class AdaLNConvBlock(nn.Module): | |
| def __init__(self, hidden_features, t_dim, dilation=1, type='linear') -> None: | |
| super().__init__() | |
| self.norm = nn.GroupNorm(32, hidden_features) | |
| self.conv1 = nn.Conv1d(hidden_features, hidden_features, kernel_size=3, dilation=dilation, padding=1 * dilation) | |
| self.conv2 = nn.Conv1d(hidden_features, hidden_features, kernel_size=3, dilation=dilation, padding=1 * dilation) | |
| if type == 'linear': | |
| self.adaLN_modulation = nn.Sequential( | |
| nn.SiLU(), | |
| nn.Linear(t_dim, 3 * hidden_features, bias=True) | |
| ) | |
| elif type == 'conv': | |
| self.adaLN_modulation = nn.Sequential( | |
| nn.SiLU(), | |
| nn.Conv1d(t_dim, 3 * hidden_features, kernel_size=3, dilation=dilation, padding=1 * dilation) | |
| ) | |
| else: | |
| raise NotImplementedError() | |
| self.type = type | |
| def forward(self, x, t): | |
| x_skip = x | |
| if t.ndim == 3: | |
| gate, shift, scale = self.adaLN_modulation(t).permute(0, 2, 1).chunk(3, dim=1) | |
| else: | |
| gate, shift, scale = self.adaLN_modulation(t).chunk(3, dim=1) | |
| if gate.ndim == 2: | |
| gate = gate.unsqueeze(-1) | |
| x = modulate(self.norm(x), shift, scale) | |
| x = F.silu(self.conv1(x)) | |
| x = gate * self.conv2(x) | |
| return x_skip + x | |
| class DylanDecoderUNet(nn.Module): | |
| def __init__(self, in_channels=3, z_dec_channels=None, channels=[320, 640, 1024], pe_dim=320, t_dim=1280, ratios=[8, 5, 4]) -> None: | |
| super().__init__() | |
| if z_dec_channels is not None: | |
| in_channels += z_dec_channels | |
| self.embed_image = ImageEmbedding(in_channels=in_channels, out_channels=channels[0]) | |
| self.embed_time = PositionalEmbedding(pe_dim=pe_dim, out_dim=t_dim) | |
| assert len(channels) == len(ratios), f'{len(channels)} != {len(ratios)}' | |
| depths = [3] * len(channels) | |
| self.down = nn.ModuleList([]) | |
| for i, (channel, depth, ratio) in enumerate(zip(channels, depths, ratios)): | |
| blocks = nn.ModuleList([]) | |
| for _ in range(depth): | |
| blocks.append(AdaLNConvBlock(channel, t_dim, dilation=2 ** _)) | |
| if ratio > 1: | |
| if i == len(channels) - 1: | |
| blocks.append(DownsampleV3(channel, channel, ratio)) | |
| else: | |
| blocks.append(DownsampleV3(channel, channels[i+1], ratio)) | |
| self.down.append(blocks) | |
| self.mid = nn.ModuleList([ | |
| AdaLNConvBlock(channels[-1], t_dim) for _ in range(depths[-1]) | |
| ]) | |
| depths = [3] * len(channels) | |
| self.skip_projs = nn.ModuleList([]) | |
| self.up = nn.ModuleList([]) | |
| self.up.append(nn.ModuleList([AdaLNConvBlock(channels[-1], t_dim)])) | |
| for i, (channel, depth, ratio) in reversed(list(enumerate(zip(channels, depths, ratios)))): | |
| blocks = nn.ModuleList([]) | |
| if ratio > 1: | |
| if i == len(channels) - 1: | |
| blocks.append(UpsampleV3(channel, channel, ratio)) | |
| else: | |
| blocks.append(UpsampleV3(channels[i+1], channel, ratio)) | |
| self.skip_projs.insert(0, nn.Conv1d(channel * 2, channel, kernel_size=1)) | |
| for _ in range(depth): | |
| blocks.append(AdaLNConvBlock(channel, t_dim, dilation=2 ** _)) | |
| self.up.insert(0, blocks) | |
| self.output = ImageUnembedding(in_channels=channels[0], out_channels=1) | |
| self.initialize_weights() | |
| def initialize_weights(self): | |
| # Initialize transformer layers: | |
| def _basic_init(module): | |
| if isinstance(module, nn.Linear): | |
| torch.nn.init.xavier_uniform_(module.weight) | |
| if module.bias is not None: | |
| nn.init.constant_(module.bias, 0) | |
| # Initialize like nn.Linear (instead of nn.Conv1d): | |
| # if isinstance(module, nn.Conv1d): | |
| # w = module.weight.data | |
| # nn.init.xavier_uniform_(w.view([w.shape[0], -1])) | |
| # if module.bias is not None: | |
| # nn.init.constant_(module.bias, 0) | |
| self.apply(_basic_init) | |
| # Initialize timestep embedding MLP: | |
| nn.init.normal_(self.embed_time.f_1.weight, std=0.02) | |
| nn.init.normal_(self.embed_time.f_2.weight, std=0.02) | |
| # Zero-out adaLN modulation layers: | |
| for blocks in self.up: | |
| for block in blocks: | |
| try: | |
| nn.init.constant_(block.adaLN_modulation[-1].weight, 0) | |
| nn.init.constant_(block.adaLN_modulation[-1].bias, 0) | |
| except: | |
| continue | |
| # Zero-out output layers: | |
| # nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0) | |
| # nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0) | |
| nn.init.constant_(self.output.f.weight, 0) | |
| nn.init.constant_(self.output.f.bias, 0) | |
| def forward(self, x, t=None, z_dec=None) -> torch.Tensor: | |
| if z_dec is not None: | |
| if z_dec.shape[-1] != x.shape[-1]: | |
| z_dec = F.interpolate(z_dec, scale_factor=x.shape[-1] // z_dec.shape[-1], mode='nearest') | |
| x = torch.cat([x, z_dec], dim=1) | |
| x = self.embed_image(x) | |
| if t is None: | |
| t = torch.zeros(x.shape[0], device=x.device) | |
| t = self.embed_time(t) | |
| skips = [x] | |
| for down in self.down: | |
| for block in down: | |
| x = block(x, t) | |
| skips.append(x) | |
| for mid in self.mid: | |
| x = mid(x, t) | |
| skips.pop() | |
| for i, up in enumerate(reversed(self.up)): | |
| for block in up: | |
| if isinstance(block, UpsampleV3): | |
| x = block(x, t) | |
| x = torch.cat([x, skips.pop()], dim=1) | |
| x = self.skip_projs[-i](x) | |
| else: | |
| x = block(x, t) | |
| return self.output(x) | |
| class DylanDecoderUNet2(nn.Module): | |
| def __init__(self, in_channels=3, z_dec_channels=None, channels=[320, 640, 1024], pe_dim=320, t_dim=1280, ratios=[8, 5, 4], type='linear') -> None: | |
| super().__init__() | |
| self.embed_image = ImageEmbedding(in_channels=in_channels, out_channels=channels[0]) | |
| self.embed_time = PositionalEmbedding(pe_dim=pe_dim, out_dim=z_dec_channels) | |
| self.embed_z = ImageEmbedding(in_channels=z_dec_channels, out_channels=z_dec_channels) | |
| self.type = type | |
| assert len(channels) == len(ratios), f'{len(channels)} != {len(ratios)}' | |
| depths = [3] * len(channels) | |
| self.c_projs = nn.ModuleList([ | |
| nn.Linear(z_dec_channels, channel) for channel in channels | |
| ]) | |
| self.down = nn.ModuleList([]) | |
| for i, (channel, depth, ratio) in enumerate(zip(channels, depths, ratios)): | |
| blocks = nn.ModuleList([]) | |
| for _ in range(depth): | |
| blocks.append(AdaLNConvBlock(channel, channel, dilation=2 ** _, type=type)) | |
| if ratio > 1: | |
| if i == len(channels) - 1: | |
| blocks.append(DownsampleV3(channel, channel, ratio)) | |
| else: | |
| blocks.append(DownsampleV3(channel, channels[i+1], ratio)) | |
| self.down.append(blocks) | |
| self.mid = nn.ModuleList([ | |
| AdaLNConvBlock(channels[-1], channels[-1], type=type) for _ in range(depths[-1]) | |
| ]) | |
| depths = [3] * len(channels) | |
| self.skip_projs = nn.ModuleList([]) | |
| self.up = nn.ModuleList([]) | |
| self.up.append(nn.ModuleList([AdaLNConvBlock(channels[-1], channels[-1], type=type)])) | |
| for i, (channel, depth, ratio) in reversed(list(enumerate(zip(channels, depths, ratios)))): | |
| blocks = nn.ModuleList([]) | |
| if ratio > 1: | |
| if i == len(channels) - 1: | |
| blocks.append(UpsampleV3(channel, channel, ratio)) | |
| else: | |
| blocks.append(UpsampleV3(channels[i+1], channel, ratio)) | |
| self.skip_projs.insert(0, nn.Conv1d(channel * 2, channel, kernel_size=1)) | |
| for _ in range(depth): | |
| blocks.append(AdaLNConvBlock(channel, channel, dilation=2 ** _, type=type)) | |
| self.up.insert(0, blocks) | |
| self.output = ImageUnembedding(in_channels=channels[0], out_channels=1) | |
| self.initialize_weights() | |
| def initialize_weights(self): | |
| # Initialize transformer layers: | |
| def _basic_init(module): | |
| if isinstance(module, nn.Linear): | |
| torch.nn.init.xavier_uniform_(module.weight) | |
| if module.bias is not None: | |
| nn.init.constant_(module.bias, 0) | |
| # Initialize like nn.Linear (instead of nn.Conv1d): | |
| # if isinstance(module, nn.Conv1d): | |
| # w = module.weight.data | |
| # nn.init.xavier_uniform_(w.view([w.shape[0], -1])) | |
| # if module.bias is not None: | |
| # nn.init.constant_(module.bias, 0) | |
| self.apply(_basic_init) | |
| # Initialize timestep embedding MLP: | |
| nn.init.normal_(self.embed_time.f_1.weight, std=0.02) | |
| nn.init.normal_(self.embed_time.f_2.weight, std=0.02) | |
| # Zero-out adaLN modulation layers: | |
| for blocks in self.up: | |
| for block in blocks: | |
| try: | |
| nn.init.constant_(block.adaLN_modulation[-1].weight, 0) | |
| nn.init.constant_(block.adaLN_modulation[-1].bias, 0) | |
| except: | |
| continue | |
| # Zero-out output layers: | |
| # nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0) | |
| # nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0) | |
| nn.init.constant_(self.output.f.weight, 0) | |
| nn.init.constant_(self.output.f.bias, 0) | |
| def interpolate(self, x, z_dec): | |
| if z_dec is not None: | |
| if z_dec.shape[-1] != x.shape[-1]: | |
| z_dec = F.interpolate(z_dec, scale_factor=x.shape[-1] // z_dec.shape[-1], mode='nearest') | |
| if self.type == 'linear': | |
| z_dec = z_dec.permute(0, 2, 1) | |
| return z_dec | |
| return None | |
| def forward(self, x, t=None, z_dec=None) -> torch.Tensor: | |
| if z_dec is not None: | |
| z_dec = self.embed_z(z_dec) | |
| x = self.embed_image(x) | |
| if t is None: | |
| t = torch.zeros(x.shape[0], device=x.device) | |
| t = self.embed_time(t) | |
| if self.type == 'linear': | |
| t = t.unsqueeze(1) | |
| elif self.type == 'conv': | |
| t = t.unsqueeze(-1) | |
| skips = [x] | |
| for i, down in enumerate(self.down): | |
| c = self.c_projs[i](t) + self.c_projs[i](self.interpolate(x, z_dec)) | |
| for block in down: | |
| x = block(x, c) | |
| skips.append(x) | |
| c = self.c_projs[-1](t) + self.c_projs[-1](self.interpolate(x, z_dec)) | |
| for mid in self.mid: | |
| x = mid(x, c) | |
| skips.pop() | |
| for i, up in enumerate(reversed(self.up)): | |
| for block in up: | |
| if isinstance(block, UpsampleV3): | |
| x = block(x, c) | |
| x = torch.cat([x, skips.pop()], dim=1) | |
| x = self.skip_projs[-i](x) | |
| c = self.c_projs[-i](t) + self.c_projs[-i](self.interpolate(x, z_dec)) | |
| else: | |
| x = block(x, c) | |
| return self.output(x) | |
| class ConsistencyDecoderUNetV2(nn.Module): | |
| def __init__(self, in_channels=3, z_dec_channels=None, channels=[320, 640, 1024], pe_dim=320, t_dim=1280, ratios=[8, 5, 4]) -> None: | |
| super().__init__() | |
| if z_dec_channels is not None: | |
| in_channels += z_dec_channels | |
| self.embed_image = ImageEmbedding(in_channels=in_channels, out_channels=channels[0]) | |
| self.embed_time = PositionalEmbedding(pe_dim=pe_dim, out_dim=t_dim) | |
| assert len(channels) == len(ratios), f'{len(channels)} != {len(ratios)}' | |
| self.down = nn.ModuleList([]) | |
| self.down.append(nn.ModuleList([ | |
| ConvResblock(channels[0], channels[0], t_dim), | |
| ConvResblock(channels[0], channels[0], t_dim), | |
| ConvResblock(channels[0], channels[0], t_dim), | |
| DownsampleV2(channels[0], t_dim, ratios[0]), | |
| ])) | |
| # Levels 1..N-1 | |
| for i in range(1, len(channels)): | |
| c_prev = channels[i - 1] | |
| c_cur = channels[i] | |
| self.down.append(nn.ModuleList([ | |
| ConvResblock(c_prev, c_cur, t_dim), | |
| ConvResblock(c_cur, c_cur, t_dim), | |
| ConvResblock(c_cur, c_cur, t_dim), | |
| DownsampleV2(c_cur, t_dim, ratios[i]), | |
| ])) | |
| # Bottom (no downsample), uses last channel again | |
| c_bot = channels[-1] | |
| self.down.append(nn.ModuleList([ | |
| ConvResblock(c_bot, c_bot, t_dim), | |
| ConvResblock(c_bot, c_bot, t_dim), | |
| ConvResblock(c_bot, c_bot, t_dim), | |
| ])) | |
| self.mid = nn.ModuleList([ | |
| ConvResblock(channels[-1], channels[-1], t_dim), | |
| ConvResblock(channels[-1], channels[-1], t_dim), | |
| ]) | |
| self.up = nn.ModuleList([]) | |
| self.up.append(nn.ModuleList([ | |
| ConvResblock(channels[-1] * 2, channels[-1], t_dim), | |
| ConvResblock(channels[-1] * 2, channels[-1], t_dim), | |
| ConvResblock(channels[-1] * 2, channels[-1], t_dim), | |
| ConvResblock(channels[-1] * 2, channels[-1], t_dim), | |
| UpsampleV2(channels[-1], t_dim, ratios[-1]), | |
| ])) | |
| self.up.append(nn.ModuleList([ | |
| ConvResblock(channels[-1] * 2, channels[-1], t_dim), | |
| ConvResblock(channels[-1] * 2, channels[-1], t_dim), | |
| ConvResblock(channels[-1] * 2, channels[-1], t_dim), | |
| ConvResblock(channels[-1] + channels[-2], channels[-1], t_dim), | |
| UpsampleV2(channels[-1], t_dim, ratios[-2]), | |
| ])) | |
| for i in range(1, len(channels) - 1): | |
| c_prev = channels[-i] | |
| c_cur = channels[-i-1] | |
| c_next = channels[-i-2] | |
| self.up.append(nn.ModuleList([ | |
| ConvResblock(c_prev + c_cur, c_cur, t_dim), | |
| ConvResblock(c_cur * 2, c_cur, t_dim), | |
| ConvResblock(c_cur * 2, c_cur, t_dim), | |
| ConvResblock(c_next + c_cur, c_cur, t_dim), | |
| UpsampleV2(c_cur, t_dim, ratios[-i-2]), | |
| ])) | |
| self.up.append(nn.ModuleList([ | |
| ConvResblock(channels[0] + channels[1], channels[0], t_dim), | |
| ConvResblock(channels[0] * 2, channels[0], t_dim), | |
| ConvResblock(channels[0] * 2, channels[0], t_dim), | |
| ConvResblock(channels[0] * 2, channels[0], t_dim), | |
| ])) | |
| self.output = ImageUnembedding(in_channels=channels[0], out_channels=1) | |
| def get_last_layer_weight(self): | |
| return self.output.f.weight | |
| def forward(self, x, t=None, z_dec=None) -> torch.Tensor: | |
| if z_dec is not None: | |
| # if z_dec.shape[-2] != x.shape[-2] or z_dec.shape[-1] != x.shape[-1]: | |
| if z_dec.shape[-1] != x.shape[-1]: | |
| # assert x.shape[-2] // z_dec.shape[-2] == x.shape[-1] // z_dec.shape[-1] | |
| # z_dec = F.upsample_nearest(z_dec, scale_factor=x.shape[-2] // z_dec.shape[-2]) | |
| z_dec = F.interpolate(z_dec, scale_factor=x.shape[-1] // z_dec.shape[-1], mode='nearest') | |
| x = torch.cat([x, z_dec], dim=1) | |
| x = self.embed_image(x) | |
| if t is None: | |
| t = torch.zeros(x.shape[0], device=x.device) | |
| t = self.embed_time(t) | |
| skips = [x] | |
| for down in self.down: | |
| for block in down: | |
| x = block(x, t) | |
| skips.append(x) | |
| for mid in self.mid: | |
| x = mid(x, t) | |
| for up in self.up: | |
| for block in up: | |
| if isinstance(block, ConvResblock): | |
| x = torch.concat([x, skips.pop()], dim=1) | |
| x = block(x, t) | |
| return self.output(x) |