| import torch.nn as nn |
| from transformers import ViTModel |
|
|
| class SEMViTAutoencoder(nn.Module): |
| def __init__(self): |
| super().__init__() |
| |
| |
| self.encoder = ViTModel.from_pretrained( |
| 'google/vit-base-patch16-224', |
| add_pooling_layer=False |
| ) |
| |
| |
| |
| |
| |
| self.decoder = nn.Sequential( |
| nn.ConvTranspose2d(768, 256, kernel_size=4, stride=2, padding=1), |
| nn.BatchNorm2d(256), |
| nn.ReLU(), |
| nn.ConvTranspose2d(256, 128, kernel_size=4, stride=2, padding=1), |
| nn.BatchNorm2d(128), |
| nn.ReLU(), |
| nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1), |
| nn.BatchNorm2d(64), |
| nn.ReLU(), |
| nn.ConvTranspose2d(64, 3, kernel_size=4, stride=2, padding=1), |
| nn.Sigmoid() |
| ) |
|
|
| def forward(self, x): |
| |
| outputs = self.encoder(x, interpolate_pos_encoding=True) |
| |
| latent = outputs.last_hidden_state[:, 1:, :].transpose(1, 2).reshape(-1, 768, 32, 32) |
| |
| return self.decoder(latent) |