Spaces:
Running on Zero
Running on Zero
File size: 5,817 Bytes
e794567 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | import os
import torch
import torch.nn as nn
from diffusers.models import AutoencoderKL
from .vae_config import VAEConfig
from transformers import PreTrainedModel
from cell_fm.pipeline.utils import VAEOutput
from cell_fm.logging import logger
class VAEModel(PreTrainedModel):
config_class = VAEConfig
def __init__(self, config):
super().__init__(config)
self.config = config
# Calculate number of downsampling and upsampling layers required
self.num_down_blocks = config.num_down_blocks # Input size / latent space size
self.num_up_blocks = self.num_down_blocks
# Generate down_block_types and up_block_types based on the required layers
down_block_types = tuple(["DownEncoderBlock2D"] * self.num_down_blocks)
up_block_types = tuple(["UpDecoderBlock2D"] * self.num_up_blocks)
# Initialize AutoencoderKL with custom parameters from config
self.vae = AutoencoderKL(
in_channels=config.in_channels,
out_channels=config.out_channels,
down_block_types=down_block_types,
up_block_types=up_block_types,
block_out_channels=config.vae_block_out_channels,
latent_channels=config.latent_channels, # Latent space dimensions
)
self.load_pretrained_weights(config, checkpoint_path=config.vae_loadcheck_path)
def load_pretrained_weights(self, config, checkpoint_path):
"""
Load pretrained weights from a given state_dict.
"""
if config.ft or config.infer:
if config.ft:
logger.info(f"Finetune from checkpoint: {checkpoint_path}")
else:
logger.info(f"Infer from checkpoint: {checkpoint_path}")
if os.path.splitext(checkpoint_path)[1] == '.safetensors':
from safetensors.torch import load_file
checkpoints_state = load_file(checkpoint_path)
else:
checkpoints_state = torch.load(checkpoint_path, map_location="cpu")
if "model" in checkpoints_state:
checkpoints_state = checkpoints_state["model"]
elif "module" in checkpoints_state:
checkpoints_state = checkpoints_state["module"]
model_state_dict = self.state_dict()
filtered_state_dict = {k: v for k, v in checkpoints_state.items() if k in model_state_dict and v.size() == model_state_dict[k].size()}
IncompatibleKeys = self.load_state_dict(filtered_state_dict, strict=False)
# IncompatibleKeys = self.load_state_dict(checkpoints_state, strict=False)
IncompatibleKeys = IncompatibleKeys._asdict()
missing_keys = []
for keys in IncompatibleKeys["missing_keys"]:
if keys.find("dummy") == -1:
missing_keys.append(keys)
unexpected_keys = []
for keys in IncompatibleKeys["unexpected_keys"]:
if keys.find("dummy") == -1:
unexpected_keys.append(keys)
if len(missing_keys) > 0:
logger.info(
"Missing keys in {}: {}".format(
checkpoint_path,
missing_keys,
)
)
if len(unexpected_keys) > 0:
logger.info(
"Unexpected keys {}: {}".format(
checkpoint_path,
unexpected_keys,
)
)
def encode(self, x):
"""Encodes input into latent space."""
return self.vae.encode(x).latent_dist
def decode(self, latents):
"""Decodes latent space into reconstructed input."""
return self.vae.decode(latents)
def forward(self, batched_data):
x = batched_data['protein_img']
"""Forward pass through the VAE."""
latent_dist = self.encode(x)
latents = latent_dist.sample() # Reparameterization trick
recon_x = self.decode(latents).sample
total_loss, recon_loss, kl_loss = self.compute_loss(x, recon_x, latent_dist)
# Log the losses
log_loss = {
"total_loss": total_loss.item(),
"recon_loss": recon_loss.item(),
"kl_loss": kl_loss.item(),
}
return VAEOutput(total_loss, log_loss)
def compute_loss(self, x, recon_x, latent_dist):
"""Compute reconstruction and KL divergence loss."""
recon_loss = nn.MSELoss()(recon_x, x)
kl_loss = -0.5 * torch.mean(1 + latent_dist.logvar - latent_dist.mean.pow(2) - latent_dist.logvar.exp())
total_loss = self.config.recon_loss_coeff * recon_loss + self.config.kl_loss_coeff * kl_loss
return total_loss, recon_loss, kl_loss
def sample(self, num_samples=1, latent_size=64, device="cpu"):
"""
Generate samples from the latent space.
Args:
num_samples (int): Number of samples to generate.
latent_size (int): Size of the latent space.
device (str): Device to perform sampling on.
Returns:
torch.Tensor: Generated images.
"""
# Sample from a standard normal distribution in latent space
latents = torch.randn((num_samples, self.config.latent_channels, latent_size, latent_size), device=device) # Shape matches latent dimensions
# Decode latents to generate images
with torch.no_grad():
generated_images = self.decode(latents).sample
return generated_images
def reconstruct(self, x):
latent_dist = self.encode(x)
latents = latent_dist.sample() # Reparameterization trick
recon_x = self.decode(latents).sample
return recon_x |