Revert the config dicts

This commit is contained in:
henryruhs
2025-03-11 14:43:10 +01:00
parent e5f983b2bf
commit 1dfd230fc5
7 changed files with 60 additions and 96 deletions
+11 -14
View File
@@ -7,31 +7,28 @@ from torch import Tensor, nn
class NLD(nn.Module):
def __init__(self, config_parser : ConfigParser) -> None:
super().__init__()
self.config =\
{
'input_channels': config_parser.getint('training.model.discriminator', 'input_channels'),
'num_filters': config_parser.getint('training.model.discriminator', 'num_filters'),
'kernel_size': config_parser.getint('training.model.discriminator', 'kernel_size'),
'num_layers': config_parser.getint('training.model.discriminator', 'num_layers')
}
self.config_input_channels = config_parser.getint('training.model.discriminator', 'input_channels')
self.config_num_filters = config_parser.getint('training.model.discriminator', 'num_filters')
self.config_kernel_size = config_parser.getint('training.model.discriminator', 'kernel_size')
self.config_num_layers = config_parser.getint('training.model.discriminator', 'num_layers')
self.layers = self.create_layers()
self.sequences = nn.Sequential(*self.layers)
def create_layers(self) -> nn.ModuleList:
padding = math.ceil((self.config.get('kernel_size') - 1) / 2)
current_filters = self.config.get('num_filters')
padding = math.ceil((self.config_kernel_size - 1) / 2)
current_filters = self.config_num_filters
layers = nn.ModuleList(
[
nn.Conv2d(self.config.get('input_channels'), current_filters, kernel_size = self.config.get('kernel_size'), stride = 2, padding = padding),
nn.Conv2d(self.config_input_channels, current_filters, kernel_size = self.config_kernel_size, stride = 2, padding = padding),
nn.LeakyReLU(0.2, True)
])
for _ in range(1, self.config.get('num_layers')):
for _ in range(1, self.config_num_layers):
previous_filters = current_filters
current_filters = min(current_filters * 2, 512)
layers +=\
[
nn.Conv2d(previous_filters, current_filters, kernel_size = self.config.get('kernel_size'), stride = 2, padding = padding),
nn.Conv2d(previous_filters, current_filters, kernel_size = self.config_kernel_size, stride = 2, padding = padding),
nn.InstanceNorm2d(current_filters),
nn.LeakyReLU(0.2, True)
]
@@ -40,10 +37,10 @@ class NLD(nn.Module):
current_filters = min(current_filters * 2, 512)
layers +=\
[
nn.Conv2d(previous_filters, current_filters, kernel_size = self.config.get('kernel_size'), padding = padding),
nn.Conv2d(previous_filters, current_filters, kernel_size = self.config_kernel_size, padding = padding),
nn.InstanceNorm2d(current_filters),
nn.LeakyReLU(0.2, True),
nn.Conv2d(current_filters, 1, kernel_size = self.config.get('kernel_size'), padding = padding)
nn.Conv2d(current_filters, 1, kernel_size = self.config_kernel_size, padding = padding)
]
return layers
+7 -10
View File
@@ -8,10 +8,7 @@ from torch import Tensor, nn
class UNet(nn.Module):
def __init__(self, config_parser : ConfigParser) -> None:
super().__init__()
self.config =\
{
'output_size': config_parser.getint('training.model.generator', 'output_size')
}
self.config_output_size = config_parser.getint('training.model.generator', 'output_size')
self.down_samples = self.create_down_samples()
self.up_samples = self.create_up_samples()
@@ -25,20 +22,20 @@ class UNet(nn.Module):
DownSample(256, 512)
])
if self.config.get('output_size') == 128:
if self.config_output_size == 128:
down_samples.extend(
[
DownSample(512, 512)
])
if self.config.get('output_size') == 256:
if self.config_output_size == 256:
down_samples.extend(
[
DownSample(512, 1024),
DownSample(1024, 1024)
])
if self.config.get('output_size') == 512:
if self.config_output_size == 512:
down_samples.extend(
[
DownSample(512, 1024),
@@ -51,20 +48,20 @@ class UNet(nn.Module):
def create_up_samples(self) -> nn.ModuleList:
up_samples = nn.ModuleList()
if self.config.get('output_size') == 128:
if self.config_output_size == 128:
up_samples.extend(
[
UpSample(512, 512)
])
if self.config.get('output_size') == 256:
if self.config_output_size == 256:
up_samples.extend(
[
UpSample(1024, 1024),
UpSample(2048, 512)
])
if self.config.get('output_size') == 512:
if self.config_output_size == 512:
up_samples.extend(
[
UpSample(2048, 2048),