DOFA / model /dofa.py
zhangrenchao's picture
Upload DOFA model package
1d4cac8 verified
Raw
History Blame Contribute Delete
7.49 kB
"""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"]