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