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