Download model/skysense.py from OneScience-Group/SkySense: direct link, hf CLI and curl.
- Browser
- Download file 6.61 kB
-
https://huggingface.co/OneScience-Group/SkySense/resolve/main/model/skysense.py
- Command line
-
hf download hf://OneScience-Group/SkySense/model/skysense.py
-
curl -L -o skysense.py https://huggingface.co/OneScience-Group/SkySense/resolve/main/model/skysense.py
6.61 kB
| """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} | |
| 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() | |