diff --git a/train/vqvae.py b/train/vqvae.py new file mode 100644 index 0000000..a753b42 --- /dev/null +++ b/train/vqvae.py @@ -0,0 +1,49 @@ +import torch +import torch.nn as nn +import torch.optim as optim + +class Encoder(nn.Module): + def __init__(self, input_channels, latent_dim): + super(Encoder, self).__init__() + self.conv1 = nn.Conv2d(input_channels, 64, kernel_size=4, stride=2, padding=1) + self.conv2 = nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=1) + self.conv3 = nn.Conv2d(128, latent_dim, kernel_size=4, stride=2, padding=1) + + def forward(self, x): + x = torch.relu(self.conv1(x)) + x = torch.relu(self.conv2(x)) + x = torch.relu(self.conv3(x)) + return x + +class VectorQuantizer(nn.Module): + def __init__(self, num_embeddings, embedding_dim): + super(VectorQuantizer, self).__init__() + self.embedding_dim = embedding_dim + self.embedding = nn.Embedding(num_embeddings, embedding_dim) + + def forward(self, x): + x_flat = x.view(-1, self.embedding_dim) + + indices = torch.argmin(torch.cdist(x_flat.unsqueeze(0), self.embedding.weight), dim=1) + + quantized = self.embedding(indices).view(x.size()) + + return quantized, indices + +class VQVAE(nn.Module): + def __init__(self, input_channels, latent_dim, num_embeddings, embedding_dim): + super(VQVAE, self).__init__() + self.encoder = Encoder(input_channels, latent_dim) + self.vector_quantizer = VectorQuantizer(num_embeddings, embedding_dim) + + def forward(self, x): + x = self.encoder(x) + quantized, indices = self.vector_quantizer(x) + return quantized, indices + +input_channels = 3 +latent_dim = 256 +num_embeddings = 512 +embedding_dim = 64 + +model = VQVAE(input_channels, latent_dim, num_embeddings, embedding_dim)