Files
modelscope-scepter/scepter/modules/model/network/autoencoder/ae_kl.py
T
2024-04-09 14:02:28 +08:00

269 lines
9.4 KiB
Python

# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from collections import OrderedDict
import numpy as np
import torch
from scepter.modules.model.network.train_module import TrainModule
from scepter.modules.model.registry import BACKBONES, LOSSES, MODELS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
class DiagonalGaussianDistribution(object):
def __init__(self, mean, logvar, deterministic=False):
self.mean = mean
self.logvar = torch.clamp(logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(
self.mean).to(device=self.mean.device)
def sample(self):
x = self.mean + self.std * torch.randn(
self.mean.shape).to(device=self.mean.device)
return x
def kl(self, other=None):
if self.deterministic:
return torch.Tensor([0.])
else:
if other is None:
return 0.5 * torch.sum(
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
dim=[1, 2, 3])
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var +
self.var / other.var - 1.0 - self.logvar + other.logvar,
dim=[1, 2, 3])
def nll(self, sample, dims=[1, 2, 3]):
if self.deterministic:
return torch.Tensor([0.])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(logtwopi + self.logvar +
torch.pow(sample - self.mean, 2) / self.var,
dim=dims)
def mode(self):
# print('*** use DiagonalGaussianDistribution.mode() ***')
return self.mean
@MODELS.register_class()
class AutoencoderKL(TrainModule):
para_dict = {
'ENCODER': {},
'DECODER': {},
'LOSS': {},
'EMBED_DIM': {
'value': 4,
'description': ''
},
'PRETRAINED_MODEL': {
'value': None,
'description': ''
},
'IGNORE_KEYS': {
'value': [],
'description': ''
},
'BATCH_SIZE': {
'value': 16,
'description': ''
},
}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.encoder_cfg = self.cfg.ENCODER
self.decoder_cfg = self.cfg.DECODER
self.loss_cfg = self.cfg.get('LOSS', None)
self.embed_dim = self.cfg.get('EMBED_DIM', 4)
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
self.batch_size = self.cfg.get('BATCH_SIZE', 16)
self.construct_network()
self.init_network()
def construct_network(self):
z_channels = self.encoder_cfg.Z_CHANNELS
self.encoder = BACKBONES.build(self.encoder_cfg, logger=self.logger)
self.decoder = BACKBONES.build(self.decoder_cfg, logger=self.logger)
self.conv1 = torch.nn.Conv2d(2 * z_channels, 2 * self.embed_dim, 1)
self.conv2 = torch.nn.Conv2d(self.embed_dim, z_channels, 1)
if self.loss_cfg is not None:
self.loss = LOSSES.build(self.loss_cfg, logger=self.logger)
def init_network(self):
if self.pretrained_model is not None:
with FS.get_from(self.pretrained_model,
wait_finish=True) as local_model:
self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
if path.find('.safetensors') > -1:
from safetensors import safe_open
sd = OrderedDict()
with safe_open(path, framework='pt', device='cpu') as f:
for k in f.keys():
sd[k] = f.get_tensor(k)
else:
sd = torch.load(path, map_location='cpu')
if path.find('.pt') > -1 and 'state_dict' in sd:
sd = sd['state_dict']
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
sd = sd['state_dict']
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
for ik in ignore_keys:
if ik in k:
if we.rank == 0:
self.logger.info(
'ignore key {} from state_dict.'.format(k))
ignored = True
break
k = k.replace('post_quant_conv',
'conv2') if 'post_quant_conv' in k else k
k = k.replace('quant_conv', 'conv1') if 'quant_conv' in k else k
if not ignored:
new_sd[k] = v
missing, unexpected = self.load_state_dict(new_sd, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def encode(self, x, return_mom=False):
return torch.cat([
self._encode(batch, return_mom=return_mom)
for batch in x.split(self.batch_size, dim=0)
],
dim=0)
def decode(self, z):
return torch.cat(
[self._decode(batch) for batch in z.split(self.batch_size, dim=0)],
dim=0)
def sample(self, moments):
mean, logvar = torch.chunk(moments, 2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
z = posterior.sample()
return z
def _encode(self, x, return_mom=False):
h = self.encoder(x)
moments = self.conv1(h)
if return_mom:
return moments
mean, logvar = torch.chunk(moments, 2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
z = posterior.sample()
return z
def _decode(self, z):
z = self.conv2(z)
dec = self.decoder(z)
return dec
def share_forward(self, image, sample_posterior=True):
posterior = self.encode(image)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
dec = self.decode(z)
return dec, posterior
def forward(self, **kwargs):
if self.training:
ret = self.forward_train(**kwargs)
else:
ret = self.forward_test(**kwargs)
return ret
def forward_train(self,
image=None,
sample_posterior=True,
optimizer_idx=0,
**kwargs):
reconstructions, posterior = self.share_forward(
image, sample_posterior)
ret = {}
if optimizer_idx == 0:
# train encoder+decoder+logvar
aeloss, log_dict_ae = self.loss(image,
reconstructions,
posterior,
optimizer_idx,
self.global_step,
last_layer=self.get_last_layer(),
split='train')
ret['loss'] = aeloss
ret.update(log_dict_ae)
if optimizer_idx == 1:
# train the discriminator
discloss, log_dict_disc = self.loss(
image,
reconstructions,
posterior,
optimizer_idx,
self.global_step,
last_layer=self.get_last_layer(),
split='train')
ret['loss'] = discloss
ret.update(log_dict_disc)
return ret
def forward_test(self, image=None, sample_posterior=True, **kwargs):
reconstructions, posterior = self.share_forward(
image, sample_posterior)
ret = {}
aeloss, log_dict_ae = self.loss(image,
reconstructions,
posterior,
0,
self.global_step,
last_layer=self.get_last_layer(),
split='val')
discloss, log_dict_disc = self.loss(image,
reconstructions,
posterior,
1,
self.global_step,
last_layer=self.get_last_layer(),
split='val')
ret.update(log_dict_ae)
ret.update(log_dict_disc)
return ret
def get_last_layer(self):
return self.decoder.conv_out.weight
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
AutoencoderKL.para_dict,
set_name=True)