Version 3 / rewrite

This commit is contained in:
City
2023-10-11 06:11:27 +02:00
parent 451cb196b7
commit 31c3eb5a82
6 changed files with 188 additions and 210 deletions
+47 -22
View File
@@ -2,30 +2,55 @@ import torch
import torch.nn as nn
import numpy as np
class Block(nn.Module):
def __init__(self, size):
super().__init__()
self.join = nn.ReLU()
self.long = nn.Sequential(
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
nn.Dropout(0.2)
)
def forward(self, x):
y = self.long(x)
z = self.join(y + x)
return z
class Interposer(nn.Module):
def __init__(self):
super().__init__()
# it looks like a spaceship if you squint :D
module_list = [
#############)
#############)
#||#
#||#
nn.Conv2d(4, 32, kernel_size=5, padding=2),
nn.ReLU(),
nn.Conv2d(32, 128, kernel_size=7, padding=3),
nn.ReLU(),
nn.Conv2d(128, 32, kernel_size=7, padding=3),
nn.ReLU(),
nn.Conv2d(32, 4, kernel_size=5, padding=2),
#||#
#||#
#############)
#############)
]
self.chan = 4 # in/out channels
self.hid = 128
self.sequential = nn.Sequential(*module_list)
# expand channels
self.head_join = nn.ReLU()
self.head_short = nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1)
self.head_long = nn.Sequential(
nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
nn.LeakyReLU(0.1),
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
)
# not sure if this is how residuals work
self.core = nn.Sequential(
Block(self.hid),
Block(self.hid),
Block(self.hid),
)
# reduce channels
self.tail = nn.Sequential(
nn.ReLU(),
nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1)
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.sequential(x)
def forward(self, x):
y = self.head_join(
self.head_long(x)+
self.head_short(x)
)
z = self.core(y)
return self.tail(z)