"""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"]