File size: 7,489 Bytes
1d4cac8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | """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"]
|