| """Compact, trainable SkySense reproduction for multi-modal remote sensing data.""" |
|
|
| import torch |
| from torch import nn |
| from torch.nn import functional as F |
|
|
|
|
| class SpatialEncoder(nn.Module): |
| def __init__(self, in_channels, embed_dim, patch_size): |
| super().__init__() |
| self.projection = nn.Sequential( |
| nn.Conv2d(in_channels, embed_dim, patch_size, patch_size), |
| nn.GELU(), |
| nn.Conv2d(embed_dim, embed_dim, 3, padding=1), |
| nn.GELU(), |
| ) |
|
|
| def forward(self, images): |
| batch, time, channels, height, width = images.shape |
| features = self.projection(images.reshape(batch * time, channels, height, width)) |
| _, dim, out_height, out_width = features.shape |
| return features.reshape(batch, time, dim, out_height, out_width) |
|
|
|
|
| class SkySense(nn.Module): |
| """Factorized spatial-temporal encoder with geo-context prototypes.""" |
|
|
| def __init__( |
| self, |
| hr_channels=3, |
| s2_channels=10, |
| s1_channels=2, |
| embed_dim=32, |
| hr_patch_size=16, |
| s2_patch_size=8, |
| s1_patch_size=8, |
| temporal_depth=2, |
| temporal_heads=4, |
| num_regions=16, |
| prototypes_per_region=4, |
| num_classes=6, |
| ): |
| super().__init__() |
| self.num_regions = num_regions |
| self.hr_encoder = SpatialEncoder(hr_channels, embed_dim, hr_patch_size) |
| self.s2_encoder = SpatialEncoder(s2_channels, embed_dim, s2_patch_size) |
| self.s1_encoder = SpatialEncoder(s1_channels, embed_dim, s1_patch_size) |
| self.date_embedding = nn.Embedding(366, embed_dim) |
| self.modality_embedding = nn.Parameter(torch.zeros(3, embed_dim)) |
| self.fusion_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) |
| layer = nn.TransformerEncoderLayer( |
| d_model=embed_dim, |
| nhead=temporal_heads, |
| dim_feedforward=embed_dim * 4, |
| dropout=0.0, |
| activation="gelu", |
| batch_first=True, |
| norm_first=True, |
| ) |
| self.temporal_fusion = nn.TransformerEncoder(layer, temporal_depth) |
| modality_layer = nn.TransformerEncoderLayer( |
| d_model=embed_dim, nhead=temporal_heads, dim_feedforward=embed_dim * 4, |
| dropout=0.0, activation="gelu", batch_first=True, norm_first=True, |
| ) |
| self.modality_fusion = nn.TransformerEncoder(modality_layer, 1) |
| self.modality_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) |
| self.prototypes = nn.Parameter( |
| torch.randn(num_regions, prototypes_per_region, embed_dim) * 0.02 |
| ) |
| self.decoder = nn.Sequential( |
| nn.Conv2d(embed_dim * 2, embed_dim, 3, padding=1), |
| nn.GELU(), |
| nn.Conv2d(embed_dim, num_classes, 1), |
| ) |
| nn.init.normal_(self.date_embedding.weight, std=0.02) |
| nn.init.normal_(self.modality_embedding, std=0.02) |
| nn.init.normal_(self.fusion_token, std=0.02) |
| nn.init.normal_(self.modality_token, std=0.02) |
|
|
| def _add_context(self, features, dates, modality_index): |
| if dates.dtype != torch.long: |
| raise TypeError(f"dates must use torch.int64, got {dates.dtype}") |
| if dates.shape != features.shape[:2]: |
| raise ValueError(f"dates shape {tuple(dates.shape)} does not match image batch/time {tuple(features.shape[:2])}") |
| if torch.any((dates < 0) | (dates > 364)): |
| raise ValueError("dates must contain day-of-year values in [0, 364]") |
| date_context = self.date_embedding(dates).unsqueeze(-1).unsqueeze(-1) |
| modality = self.modality_embedding[modality_index].view(1, 1, -1, 1, 1) |
| return features + date_context + modality |
|
|
| def encode_modalities(self, hr, s2, s1, dates_hr, dates_s2, dates_s1): |
| return ( |
| self._add_context(self.hr_encoder(hr), dates_hr, 0), |
| self._add_context(self.s2_encoder(s2), dates_s2, 1), |
| self._add_context(self.s1_encoder(s1), dates_s1, 2), |
| ) |
|
|
| def _aggregate_time(self, features): |
| batch, time, dim, height, width = features.shape |
| sequence = features.permute(0, 3, 4, 1, 2).reshape(-1, time, dim) |
| token = self.fusion_token.expand(sequence.shape[0], -1, -1) |
| fused = self.temporal_fusion(torch.cat([token, sequence], dim=1))[:, 0] |
| return fused.reshape(batch, height, width, dim).permute(0, 3, 1, 2) |
|
|
| def forward(self, hr, s2, s1, dates_hr, dates_s2, dates_s1, region): |
| if region.dtype != torch.long: |
| raise TypeError(f"region must use torch.int64, got {region.dtype}") |
| if region.shape != (hr.shape[0],): |
| raise ValueError(f"region must have shape [{hr.shape[0]}], got {tuple(region.shape)}") |
| if torch.any((region < 0) | (region >= self.num_regions)): |
| raise ValueError(f"region IDs must be in [0, {self.num_regions - 1}]") |
| modality_features = self.encode_modalities(hr, s2, s1, dates_hr, dates_s2, dates_s1) |
| aggregated = [self._aggregate_time(feature) for feature in modality_features] |
| target_size = aggregated[0].shape[-2:] |
| aligned = [aggregated[0]] + [ |
| F.interpolate(feature, size=target_size, mode="bilinear", align_corners=False) |
| for feature in aggregated[1:] |
| ] |
| batch, dim, out_height, out_width = aligned[0].shape |
| modalities = torch.stack(aligned, dim=1).permute(0, 3, 4, 1, 2).reshape(-1, 3, dim) |
| token = self.modality_token.expand(modalities.shape[0], -1, -1) |
| fused = self.modality_fusion(torch.cat([token, modalities], dim=1))[:, 0] |
| fused = fused.reshape(batch, out_height, out_width, dim) |
|
|
| regional_prototypes = self.prototypes[region] |
| query = F.normalize(fused, dim=-1) |
| keys = F.normalize(regional_prototypes, dim=-1) |
| attention = torch.einsum("bhwd,bpd->bhwp", query, keys).softmax(dim=-1) |
| geo_context = torch.einsum("bhwp,bpd->bhwd", attention, regional_prototypes) |
| output = torch.cat([fused, geo_context], dim=-1).permute(0, 3, 1, 2) |
| logits = self.decoder(output) |
| logits = F.interpolate(logits, size=hr.shape[-2:], mode="bilinear", align_corners=False) |
| return {"logits": logits, "features": modality_features, "fused": fused} |
|
|
| @staticmethod |
| def cross_modal_alignment_loss(features): |
| pooled = [F.normalize(feature.mean(dim=(1, 3, 4)), dim=-1) for feature in features] |
| losses = [1.0 - (pooled[i] * pooled[j]).sum(dim=-1).mean() for i in range(3) for j in range(i + 1, 3)] |
| return torch.stack(losses).mean() |
|
|