This commit is contained in:
Andras Schmelczer 2024-06-29 10:14:12 +01:00
commit 137ba1c475
No known key found for this signature in database
GPG key ID: FC8F2C3D3D1A718C
2 changed files with 12 additions and 18 deletions

View file

@ -28,7 +28,6 @@ class HistogramNet(nn.Module):
self._use_elu = use_elu self._use_elu = use_elu
self._leaky_relu_alpha = leaky_relu_alpha self._leaky_relu_alpha = leaky_relu_alpha
self._use_residual = use_residual self._use_residual = use_residual
self.print_og_result = False
self._convolutions = nn.ModuleList( self._convolutions = nn.ModuleList(
self._make_conv_layer(in_channels=in_channels, out_channels=out_channels) self._make_conv_layer(in_channels=in_channels, out_channels=out_channels)
@ -55,7 +54,7 @@ class HistogramNet(nn.Module):
in_channels=in_channels, in_channels=in_channels,
out_channels=out_channels, out_channels=out_channels,
kernel_size=self._kernel_size, kernel_size=self._kernel_size,
padding=1, padding=self._kernel_size // 2,
bias=False, bias=False,
), ),
( (
@ -75,21 +74,22 @@ class HistogramNet(nn.Module):
in_channels=channels, in_channels=channels,
out_channels=channels, out_channels=channels,
kernel_size=self._kernel_size, kernel_size=self._kernel_size,
padding=1, padding=self._kernel_size // 2,
bias=False, bias=False,
), ),
( (
nn.ELU(self._elu_alpha) nn.ELU(self._elu_alpha)
if self._use_elu if self._use_elu
else nn.LeakyReLU(self._leaky_relu_alpha)( else nn.LeakyReLU(self._leaky_relu_alpha)
nn.InstanceNorm3d if self._use_instance_norm else nn.BatchNorm3d ),
)(channels) (nn.InstanceNorm3d if self._use_instance_norm else nn.BatchNorm3d)(
channels
), ),
nn.Conv3d( nn.Conv3d(
in_channels=channels, in_channels=channels,
out_channels=channels, out_channels=channels,
kernel_size=self._kernel_size, kernel_size=self._kernel_size,
padding=1, padding=self._kernel_size // 2,
bias=False, bias=False,
), ),
( (
@ -108,7 +108,7 @@ class HistogramNet(nn.Module):
in_channels=in_channels, in_channels=in_channels,
out_channels=out_channels, out_channels=out_channels,
kernel_size=self._kernel_size, kernel_size=self._kernel_size,
padding=1, padding=self._kernel_size // 2,
), ),
( (
nn.ELU(self._elu_alpha) nn.ELU(self._elu_alpha)
@ -129,10 +129,6 @@ class HistogramNet(nn.Module):
for deconv in self._deconvolutions: for deconv in self._deconvolutions:
x = deconv(x) x = deconv(x)
if self.print_og_result:
logging.info(f"Original result {torch.sum(x)}")
self.print_og_result = False
return self._normalize(x) return self._normalize(x)
@staticmethod @staticmethod
@ -144,7 +140,6 @@ class HistogramNet(nn.Module):
def _initialize_weights(self): def _initialize_weights(self):
for m in self.modules(): for m in self.modules():
if isinstance(m, (nn.Conv3d, nn.ConvTranspose3d)): if isinstance(m, (nn.Conv3d, nn.ConvTranspose3d)):
# Applying He normal initialization
nn.init.kaiming_normal_(m.weight, mode="fan_in", nonlinearity="relu") nn.init.kaiming_normal_(m.weight, mode="fan_in", nonlinearity="relu")
if m.bias is not None: if m.bias is not None:
nn.init.constant_(m.bias, 0) nn.init.constant_(m.bias, 0)

View file

@ -22,6 +22,7 @@ def random_hparam_search(
device: torch.device, device: torch.device,
) -> None: ) -> None:
for _ in count(): for _ in count():
run_id = get_next_run_name(tensorboard_path)
current_hyperparameters = { current_hyperparameters = {
k: v.rvs() if hasattr(v, "rvs") else choice(v) k: v.rvs() if hasattr(v, "rvs") else choice(v)
for k, v in choice(hyperparameters).items() for k, v in choice(hyperparameters).items()
@ -29,11 +30,9 @@ def random_hparam_search(
serialized_hparams = json.dumps( serialized_hparams = json.dumps(
current_hyperparameters, indent=2, sort_keys=True current_hyperparameters, indent=2, sort_keys=True
) )
logging.info( logging.info(f"Starting {run_id} with hparams {serialized_hparams}")
f"Starting {get_next_run_name(tensorboard_path)} with hparams {serialized_hparams}"
)
log_dir = tensorboard_path / get_next_run_name(tensorboard_path) log_dir = tensorboard_path / run_id
try: try:
model = train( model = train(
@ -46,7 +45,7 @@ def random_hparam_search(
device=device, device=device,
**current_hyperparameters, **current_hyperparameters,
) )
model_path = models_path / get_next_run_name(models_path) model_path = models_path / run_id
save_model(model, current_hyperparameters, model_path) save_model(model, current_hyperparameters, model_path)
del model del model
except KeyboardInterrupt as e: except KeyboardInterrupt as e: