SpectralGPT / model /spectralgpt.py
zhangrenchao's picture
Update SpectralGPT model package
7a2d30b verified
Raw
History Blame Contribute Delete
6.71 kB
import torch
from torch import nn
class SpectralGPT(nn.Module):
"""Compact SpectralGPT masked autoencoder for 12-band spectral images."""
def __init__(self, image_size=24, in_channels=12, patch_size=8,
spectral_patch_size=3, embed_dim=48, encoder_depth=2,
encoder_heads=4, decoder_dim=32, decoder_depth=1,
decoder_heads=4, mask_ratio=0.9, spectral_angle_weight=0.1,
spectral_gradient_weight=0.1):
super().__init__()
if image_size % patch_size or in_channels % spectral_patch_size:
raise ValueError("Image and spectral dimensions must be divisible by token sizes")
self.image_size = image_size
self.in_channels = in_channels
self.patch_size = patch_size
self.spectral_patch_size = spectral_patch_size
self.spatial_tokens = (image_size // patch_size) ** 2
self.spectral_tokens = in_channels // spectral_patch_size
self.num_tokens = self.spatial_tokens * self.spectral_tokens
self.token_pixels = patch_size * patch_size * spectral_patch_size
self.mask_ratio = mask_ratio
self.spectral_angle_weight = spectral_angle_weight
self.spectral_gradient_weight = spectral_gradient_weight
self.patch_embed = nn.Conv3d(
1, embed_dim,
kernel_size=(spectral_patch_size, patch_size, patch_size),
stride=(spectral_patch_size, patch_size, patch_size),
)
self.spatial_pos = nn.Parameter(torch.zeros(1, self.spatial_tokens, embed_dim))
self.spectral_pos = nn.Parameter(torch.zeros(1, self.spectral_tokens, embed_dim))
encoder_layer = nn.TransformerEncoderLayer(
embed_dim, encoder_heads, embed_dim * 4, batch_first=True, norm_first=True
)
self.encoder = nn.TransformerEncoder(encoder_layer, encoder_depth)
self.encoder_norm = nn.LayerNorm(embed_dim)
self.decoder_embed = nn.Linear(embed_dim, decoder_dim)
self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))
self.decoder_pos = nn.Linear(embed_dim, decoder_dim, bias=False)
decoder_layer = nn.TransformerEncoderLayer(
decoder_dim, decoder_heads, decoder_dim * 4, batch_first=True, norm_first=True
)
self.decoder = nn.TransformerEncoder(decoder_layer, decoder_depth)
self.decoder_norm = nn.LayerNorm(decoder_dim)
self.decoder_pred = nn.Linear(decoder_dim, self.token_pixels)
nn.init.normal_(self.spatial_pos, std=0.02)
nn.init.normal_(self.spectral_pos, std=0.02)
nn.init.normal_(self.mask_token, std=0.02)
def _positions(self):
return (self.spatial_pos[:, None] + self.spectral_pos[:, :, None]).reshape(
1, self.num_tokens, -1
)
def patchify(self, images):
p, k = self.patch_size, self.spectral_patch_size
n, c, h, w = images.shape
if (c, h, w) != (self.in_channels, self.image_size, self.image_size):
raise ValueError(f"Expected [N,{self.in_channels},{self.image_size},{self.image_size}]")
x = images.reshape(n, c // k, k, h // p, p, w // p, p)
x = x.permute(0, 1, 3, 5, 2, 4, 6)
return x.reshape(n, self.num_tokens, self.token_pixels)
def unpatchify(self, tokens):
p, k = self.patch_size, self.spectral_patch_size
n = tokens.shape[0]
s = self.image_size // p
x = tokens.reshape(n, self.spectral_tokens, s, s, k, p, p)
x = x.permute(0, 1, 4, 2, 5, 3, 6)
return x.reshape(n, self.in_channels, self.image_size, self.image_size)
@staticmethod
def random_masking(tokens, mask_ratio):
n, length, dim = tokens.shape
keep = max(1, int(length * (1.0 - mask_ratio)))
order = torch.argsort(torch.rand(n, length, device=tokens.device), dim=1)
restore = torch.argsort(order, dim=1)
keep_ids = order[:, :keep]
visible = torch.gather(tokens, 1, keep_ids.unsqueeze(-1).expand(-1, -1, dim))
mask = torch.ones(n, length, device=tokens.device)
mask[:, :keep] = 0
mask = torch.gather(mask, 1, restore)
return visible, mask, restore
def forward(self, images, mask_ratio=None):
ratio = self.mask_ratio if mask_ratio is None else mask_ratio
embedded = self.patch_embed(images.unsqueeze(1)).flatten(2).transpose(1, 2)
positions = self._positions()
visible, mask, restore = self.random_masking(embedded + positions, ratio)
latent = self.encoder_norm(self.encoder(visible))
decoded_visible = self.decoder_embed(latent)
missing = self.num_tokens - decoded_visible.shape[1]
full = torch.cat([decoded_visible, self.mask_token.expand(images.shape[0], missing, -1)], 1)
full = torch.gather(full, 1, restore.unsqueeze(-1).expand(-1, -1, full.shape[-1]))
prediction = self.decoder_pred(self.decoder_norm(self.decoder(full + self.decoder_pos(positions))))
target = self.patchify(images)
token_error = (prediction - target).pow(2).mean(-1)
masked_mse = (token_error * mask).sum() / mask.sum().clamp_min(1)
mask_image = self.unpatchify(mask.unsqueeze(-1).expand(-1, -1, self.token_pixels))
predicted_image = self.unpatchify(prediction)
completed = images * (1.0 - mask_image) + predicted_image * mask_image
spectral_mask = mask_image.any(dim=1)
completed_norm = completed.norm(dim=1)
target_norm = images.norm(dim=1)
valid_sam = spectral_mask & (completed_norm > 1e-6) & (target_norm > 1e-6)
cosine = (completed * images).sum(dim=1) / (completed_norm * target_norm).clamp_min(1e-6)
angles = torch.acos(cosine.clamp(-1.0, 1.0))
spectral_angle = (angles * valid_sam).sum() / valid_sam.sum().clamp_min(1)
completed_gradient = completed[:, 1:] - completed[:, :-1]
target_gradient = images[:, 1:] - images[:, :-1]
gradient_mask = torch.maximum(mask_image[:, 1:], mask_image[:, :-1]).bool()
spectral_gradient = ((completed_gradient - target_gradient).abs() * gradient_mask).sum() / gradient_mask.sum().clamp_min(1)
loss = (masked_mse + self.spectral_angle_weight * spectral_angle
+ self.spectral_gradient_weight * spectral_gradient)
return {
"loss": loss,
"masked_mse": masked_mse,
"spectral_angle": spectral_angle,
"spectral_gradient": spectral_gradient,
"prediction": prediction,
"mask": mask,
"mask_image": mask_image,
"prediction_image": predicted_image,
"reconstruction": completed,
}