Spaces:
Running
Running
Download src/ae.py from EndeavoringYoon/LEGO: direct link, hf CLI and curl.
- Browser
- Download file 3.93 kB
-
https://huggingface.co/spaces/EndeavoringYoon/LEGO/resolve/main/src/ae.py
- Command line
-
hf download hf://spaces/EndeavoringYoon/LEGO/src/ae.py
-
curl -L -o ae.py https://huggingface.co/spaces/EndeavoringYoon/LEGO/resolve/main/src/ae.py
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 |