Spaces:
Running on Zero
Running on Zero
Download cell_fm/models/vae/vae_model.py from BoHuangLab/CELL-FM: direct link, hf CLI and curl.
- Browser
- Download file 5.82 kB
-
https://huggingface.co/spaces/BoHuangLab/CELL-FM/resolve/main/cell_fm/models/vae/vae_model.py
- Command line
-
hf download hf://spaces/BoHuangLab/CELL-FM/cell_fm/models/vae/vae_model.py
-
curl -L -o vae_model.py https://huggingface.co/spaces/BoHuangLab/CELL-FM/resolve/main/cell_fm/models/vae/vae_model.py
5.82 kB
| 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 |