LEGO / src /ae.py
EndeavoringYoon's picture
Update src/ae.py
ff9048e verified
Raw History Blame Contribute Delete
3.93 kB
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