add vqvae
This commit is contained in:
49
train/vqvae.py
Normal file
49
train/vqvae.py
Normal file
@@ -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)
|
||||
Reference in New Issue
Block a user