WideCodec / neucodec /tconv /base.py
huseinzolkepliscicom's picture
Add neucodec package for self-contained inference
259eeac verified
Raw
History Blame Contribute Delete
2.89 kB
import einops
import torch
from torch import nn
import torch.nn.functional as F
from ..distill_layers import Conv1d
def trend_pool(x, kernel_size):
if kernel_size > 1:
pool_args = dict(kernel_size=kernel_size, stride=1, padding=kernel_size // 2)
return F.avg_pool1d(F.max_pool1d(x.abs(), **pool_args), **pool_args)
# return F.avg_pool1d(F.max_pool1d(x, **pool_args), **pool_args) # woabs
else:
return x
class TrendPool(nn.Module):
def __init__(self, kernel_size=5):
super().__init__()
self.kernel_size = kernel_size
def forward(self, x):
return trend_pool(x, self.kernel_size)
class FirstBlock(nn.Module):
def __init__(
self,
target_dim,
conv_kernels=(7, 7, 7, 7),
pool_kernels=(1, 3, 5, 9),
dilation_rate=2,
):
super().__init__()
assert target_dim % len(pool_kernels) == 0
each_dim = target_dim // len(pool_kernels)
blocks = []
for conv_kernel, pool_kernel in zip(conv_kernels, pool_kernels):
conv_dilation = pool_kernel // dilation_rate + 1
conv_padding = (conv_kernel - 1) * conv_dilation // 2
blocks.append(
nn.Sequential(
TrendPool(pool_kernel),
Conv1d(
1,
each_dim,
kernel_size=conv_kernel,
dilation=conv_dilation,
padding=conv_padding,
),
)
)
self.blocks = nn.ModuleList(blocks)
def forward(self, x):
return torch.cat([block(x) for block in self.blocks], dim=1)
class EnhanceBlock(FirstBlock):
def __init__(self, dim):
super().__init__(4, conv_kernels=(7, 7, 7, 7), pool_kernels=(1, 3, 5, 9))
self.dim = dim
self.merge_layer = nn.Sequential(
# nn.LeakyReLU(), # ! if active or use InstanceNorm1d
nn.InstanceNorm1d(4, affine=True),
nn.Conv1d(4, 1, kernel_size=1),
)
def forward(self, x):
x = einops.rearrange(x, "b c t -> (b c) 1 t", c=self.dim)
y = super().forward(x)
y = self.merge_layer(y)
y = einops.rearrange(y, "(b c) 1 t -> b c t", c=self.dim)
return y # ! x + y or x + y * x
class SimpleEnhanceBlock(FirstBlock):
def __init__(self, dim):
super().__init__(4, conv_kernels=(7, 7, 7, 7), pool_kernels=(1, 3, 5, 9))
self.dim = dim
self.merge_layer = nn.Sequential(
# nn.LeakyReLU(), # ! if active or use InstanceNorm1d
nn.InstanceNorm1d(4, affine=True),
nn.Conv1d(4, self.dim, kernel_size=1),
)
def forward(self, x):
xi = x[:, :1, :]
yi = super().forward(xi)
y = self.merge_layer(yi)
return x + y * x