| """Wavelength-conditioned, visible-token DOFA masked autoencoder.""" |
|
|
| import math |
|
|
| import torch |
| from torch import nn |
| from torch.nn import functional as F |
|
|
|
|
| class WavelengthEncoder(nn.Module): |
| def __init__(self, dimension): |
| super().__init__() |
| self.dimension = dimension |
| self.mlp = nn.Sequential(nn.Linear(dimension, dimension), nn.GELU(), |
| nn.Linear(dimension, dimension)) |
|
|
| def forward(self, wavelengths): |
| half = self.dimension // 2 |
| scale = torch.exp(torch.arange(half, device=wavelengths.device, |
| dtype=wavelengths.dtype) * |
| (-math.log(10000.0) / max(half - 1, 1))) |
| phase = wavelengths[:, None] * 1000.0 * scale[None] |
| encoded = torch.cat((phase.sin(), phase.cos()), -1) |
| encoded = F.pad(encoded, (0, self.dimension - encoded.shape[-1])) |
| return encoded + self.mlp(encoded) |
|
|
|
|
| class DynamicPatchWeights(nn.Module): |
| def __init__(self, wavelength_dim, embed_dim, decoder_dim, patch_size, heads): |
| super().__init__() |
| layer = nn.TransformerEncoderLayer(wavelength_dim, heads, wavelength_dim * 2, |
| batch_first=True, norm_first=True, dropout=0.0, |
| activation="gelu") |
| self.context = nn.TransformerEncoder(layer, 1) |
| self.patch_size = patch_size |
| self.embed_dim = embed_dim |
| self.decoder_dim = decoder_dim |
| self.encoder_weight = nn.Linear(wavelength_dim, embed_dim * patch_size**2) |
| self.encoder_bias = nn.Linear(wavelength_dim, embed_dim) |
| self.decoder_weight = nn.Linear(wavelength_dim, decoder_dim * patch_size**2) |
| self.decoder_bias = nn.Linear(wavelength_dim, patch_size**2) |
|
|
| def forward(self, wavelength_features): |
| context = self.context(wavelength_features.unsqueeze(0)).squeeze(0) |
| channels = len(context) |
| encoder_weight = self.encoder_weight(context).reshape( |
| channels, self.embed_dim, self.patch_size, self.patch_size).permute(1, 0, 2, 3) |
| encoder_bias = self.encoder_bias(context).mean(0) |
| decoder_weight = self.decoder_weight(context).reshape( |
| channels, self.patch_size**2, self.decoder_dim) |
| decoder_bias = self.decoder_bias(context) |
| return encoder_weight, encoder_bias, decoder_weight, decoder_bias |
|
|
|
|
| class DOFA(nn.Module): |
| def __init__(self, image_size=224, patch_size=16, embed_dim=64, encoder_depth=2, |
| encoder_heads=4, decoder_dim=32, decoder_depth=1, decoder_heads=4, |
| wavelength_embed_dim=32, hypernetwork_heads=4, mask_ratio=0.75): |
| super().__init__() |
| if image_size % patch_size: |
| raise ValueError("image_size must be divisible by patch_size") |
| self.image_size, self.patch_size = image_size, patch_size |
| self.mask_ratio = mask_ratio |
| self.num_patches = (image_size // patch_size) ** 2 |
| self.wavelength_encoder = WavelengthEncoder(wavelength_embed_dim) |
| self.dynamic_weights = DynamicPatchWeights(wavelength_embed_dim, embed_dim, decoder_dim, |
| patch_size, hypernetwork_heads) |
| self.encoder_position = nn.Parameter(torch.zeros(1, self.num_patches, embed_dim)) |
| encoder_layer = nn.TransformerEncoderLayer(embed_dim, encoder_heads, embed_dim * 4, |
| batch_first=True, norm_first=True, |
| dropout=0.0, activation="gelu") |
| self.encoder = nn.TransformerEncoder(encoder_layer, encoder_depth) |
| self.decoder_input = nn.Linear(embed_dim, decoder_dim) |
| self.decoder_position = nn.Parameter(torch.zeros(1, self.num_patches, decoder_dim)) |
| self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim)) |
| decoder_layer = nn.TransformerEncoderLayer(decoder_dim, decoder_heads, decoder_dim * 4, |
| batch_first=True, norm_first=True, |
| dropout=0.0, activation="gelu") |
| self.decoder = nn.TransformerEncoder(decoder_layer, decoder_depth) |
| nn.init.normal_(self.encoder_position, std=0.02) |
| nn.init.normal_(self.decoder_position, std=0.02) |
| nn.init.normal_(self.mask_token, std=0.02) |
|
|
| def patchify(self, images): |
| p = self.patch_size |
| batch, channels, height, width = images.shape |
| patches = images.reshape(batch, channels, height // p, p, width // p, p) |
| return patches.permute(0, 2, 4, 1, 3, 5).reshape(batch, self.num_patches, |
| channels, p * p) |
|
|
| def unpatchify(self, patches): |
| p, side = self.patch_size, self.image_size // self.patch_size |
| batch, _, channels, _ = patches.shape |
| return patches.reshape(batch, side, side, channels, p, p).permute( |
| 0, 3, 1, 4, 2, 5).reshape(batch, channels, self.image_size, self.image_size) |
|
|
| def random_mask(self, batch, ratio, device): |
| visible = max(1, int(self.num_patches * (1 - ratio))) |
| order = torch.rand(batch, self.num_patches, device=device).argsort(1) |
| visible_indices = order[:, :visible] |
| mask = torch.ones(batch, self.num_patches, dtype=torch.bool, device=device) |
| mask.scatter_(1, visible_indices, False) |
| return visible_indices, mask |
|
|
| def forward(self, images, wavelengths, mask_ratio=None): |
| if images.shape[-2:] != (self.image_size, self.image_size): |
| raise ValueError(f"Expected {self.image_size}x{self.image_size} images") |
| if images.shape[1] != wavelengths.numel(): |
| raise ValueError("Image channels and wavelengths must have equal lengths") |
| wavelength_features = self.wavelength_encoder(wavelengths) |
| encoder_weight, encoder_bias, decoder_weight, decoder_bias = self.dynamic_weights( |
| wavelength_features) |
| tokens = F.conv2d(images, encoder_weight, encoder_bias, |
| stride=self.patch_size).flatten(2).transpose(1, 2) |
| visible_indices, mask = self.random_mask(len(images), self.mask_ratio if mask_ratio is None |
| else mask_ratio, images.device) |
| gather = visible_indices.unsqueeze(-1).expand(-1, -1, tokens.shape[-1]) |
| visible_tokens = torch.gather(tokens + self.encoder_position, 1, gather) |
| encoded = self.encoder(visible_tokens) |
| decoded_tokens = self.mask_token.expand(len(images), self.num_patches, -1).clone() |
| visible_decoded = self.decoder_input(encoded).to(decoded_tokens.dtype) |
| decoded_tokens.scatter_(1, visible_indices.unsqueeze(-1).expand(-1, -1, |
| decoded_tokens.shape[-1]), visible_decoded) |
| decoded = self.decoder(decoded_tokens + self.decoder_position) |
| predictions = torch.einsum("bnd,cpd->bncp", decoded, decoder_weight) |
| predictions = predictions + decoder_bias[None, None] |
| targets = self.patchify(images) |
| patch_error = (predictions - targets).square().mean((2, 3)) |
| loss = (patch_error * mask).sum() / mask.sum().clamp_min(1) |
| return {"loss": loss, "reconstruction": self.unpatchify(predictions), "mask": mask, |
| "visible_indices": visible_indices, "features": encoded} |
|
|
|
|
| __all__ = ["DOFA", "WavelengthEncoder", "DynamicPatchWeights"] |
|
|