import torch as th import torch.nn as nn import torch.nn.functional as F import numpy as np from geometry import relaxed_distortion_measure from pathlib import Path import json, os class AutoEncoderClass(nn.Module): def __init__( self, x_dim, z_dim=2, h_dims=[64,32], loss_type='l2', # 'l2', 'l1', 'huber' actv=nn.ReLU(), # activation function iso = True, iso_reg = 1e-7, ): super(AutoEncoderClass, self).__init__() self.name = 'autoencoder' self.z_dim = z_dim self.h_dims = h_dims self.loss_type = loss_type.lower() self.actv = actv self.iso = iso self.iso_reg = iso_reg if isinstance(x_dim, int): self.input_shape = (x_dim,) self.flattened_dim = x_dim elif isinstance(x_dim, tuple): self.input_shape = x_dim self.flattened_dim = int(np.prod(x_dim)) else: raise ValueError("x_dim must be int or tuple") enc_layers = [] in_dim = self.flattened_dim for h_dim in self.h_dims: enc_layers.append(nn.Linear(in_dim, h_dim)) enc_layers.append(self.actv) in_dim = h_dim enc_layers.append(nn.Linear(in_dim, self.z_dim)) self.encoder = nn.Sequential(*enc_layers) dec_layers = [] in_dim = self.z_dim for h_dim in reversed(self.h_dims): dec_layers.append(nn.Linear(in_dim, h_dim)) dec_layers.append(self.actv) in_dim = h_dim dec_layers.append(nn.Linear(in_dim, self.flattened_dim)) self.decoder = nn.Sequential(*dec_layers) def encode(self, x): return self.encoder(x.view(x.size(0), -1)) def decode(self, z): return self.decoder(z).view(z.size(0), *self.input_shape) def forward(self, x): z = self.encode(x); return self.decode(z), z def compute_loss(self, x, x_hat): if self.loss_type == 'l2': loss = nn.functional.mse_loss(x_hat, x) elif self.loss_type == 'l1': loss = nn.functional.l1_loss(x_hat, x) elif self.loss_type == 'huber': loss = nn.functional.huber_loss(x_hat, x) elif self.loss_type == 'mseloss': loss = nn.functional.mse_loss(x_hat,x) else: raise ValueError(f"Unknown loss_type: {self.loss_type}") if self.iso: iso_loss = relaxed_distortion_measure(self.decode, self.encode(x), eta=0.2, metric = 'identity') total_loss = loss + self.iso_reg * iso_loss else: iso_loss = th.tensor(0.0, device=x.device, dtype=loss.dtype) total_loss =loss return total_loss, loss, iso_loss def load_saved_model(device='cpu'): """ load AutoEncoderClass instance using saved config.json & model.pt Args: model_name device (str): 'cpu' / 'cuda:0' Returns: ae (AutoEncoderClass): model instance config (dict): configuration """ current_dir = Path(__file__).parent config_path = current_dir / "config.json" model_path = current_dir / "model.pt" if not config_path.exists(): raise FileNotFoundError(f"Cannot find path. Current path: {os.getcwd()}") with open(config_path, 'r') as f: config = json.load(f) actv_map = { 'Tanh': th.nn.Tanh(), 'ReLU': th.nn.ReLU(), 'SiLU': th.nn.SiLU(), 'GELU': th.nn.GELU(), } actv_fn = actv_map[config['actv']] ae = AutoEncoderClass( x_dim=config['x_dim'][0], z_dim=config['z_dim'], h_dims=config['h_dims'], loss_type=config['loss_type'], actv=actv_fn, iso=config['iso'], iso_reg=config['iso_reg'], ).to(device) ae.load_state_dict(th.load(model_path, map_location=device)) ae.eval() return ae, config