Framework de segmentação semântica para imagens de satélite multiespectrais, comparando duas arquiteturas:
- Prithvi: Foundation model geoespacial IBM/NASA (ViT pré-treinado em dados HLS — Landsat/Sentinel-2 harmonizados). Variantes
tiny,100m,300m,600m - AttentionResUNet: Encoder ResNet34/50 + decoder com attention gates espaciais e de canal
O framework é agnóstico a sensor e número de classes: a quantidade de bandas, a ordem delas e o número de classes são definidos via YAML. O mapeamento das bandas de entrada para os pesos pré-treinados do HLS é explícito (model_bands), o que evita inicializar com peso errado quando a ordem das bandas difere do HLS.
- Arquitetura do Projeto
- Setup — Windows (venv)
- Setup — Docker
- Pipeline de Dados
- Mapeamento de Bandas
- Treinamento
- Inferência
- Exportação de Modelos
- Estrutura de Diretórios
- Configurações
- Métricas e Monitoramento
- Benchmark
SatSegmentation/
├── config/
│ ├── prithvi.yaml # Hiperparâmetros do Prithvi (incluindo model_bands)
│ └── unet.yaml # Hiperparâmetros do U-Net
├── scripts/
│ ├── prithvi/
│ │ ├── train_prithvi.py # Loop de treino Prithvi
│ │ ├── predict_prithvi.py# Inferência em batch (TorchScript FP16)
│ │ └── prithvi_fp16.py # Conversão para FP16 + TorchScript
│ └── unet/
│ ├── train_unet.py # Loop de treino U-Net
│ ├── predict_unet.py # Inferência ONNX com tiling e overlap
│ └── export_onnx_unet.py
├── src/
│ ├── model.py # AttentionResUNet + Prithvi11BandsModel (com _HLS_BAND_INDEX)
│ ├── dataset.py # SegDatasetMemmap, PrithviDataset, augmentação GPU
│ ├── metrics.py # FocalDiceLoss, mIoU, Dice, Kappa, plots
│ ├── utils.py # Pesos de classe, avaliação, file pairing
│ ├── checkpoint_model.py # Save/load de estado completo de treino
│ ├── eval.py # Loop de validação com CSV + confusion matrix
│ └── fix_terratorch.py # Patch de instalação do TerraTorch (one-time)
├── data/
│ ├── img/ # GeoTIFFs de entrada
│ └── mask/ # Máscaras de segmentação (*_mask.tif)
└── memmap_output/ # Dados pré-processados (gerado automaticamente)
git clone https://github.com/Ga0512/SatSegmentation.git
cd SatSegmentation
python -m venv venv
venv\Scripts\activate
python -m pip install --upgrade pip
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121
pip install -r requirements.txtPatch obrigatório (executar uma única vez após instalar):
python -m src.fix_terratorchVariável de ambiente para rasterio (necessária em toda sessão):
$env:PROJ_LIB = "$pwd\venv\Lib\site-packages\rasterio\proj_data"- Docker
- NVIDIA Container Toolkit (para acesso à GPU dentro do container)
sudo apt-get update && sudo apt-get install -y nvidia-container-toolkit
sudo systemctl restart docker
# Verificar acesso à GPU
docker run --rm --gpus all nvidia/cuda:12.1.1-base-ubuntu22.04 nvidia-smidocker build -t satseg .
docker run -it --gpus all -v ${PWD}:/app satsegDentro do container:
python -m src.fix_terratorchO pipeline transforma GeoTIFFs brutos em arquivos memory-mapped binários (.dat) para carregamento eficiente durante o treino.
GeoTIFFs (./data/) → build_memmap_and_stats() → ./memmap_output/*.dat + stats.npz
↓
SegDatasetMemmap / PrithviDataset
↓
DataLoader (workers CPU, pin_memory)
↓
Normalização na GPU (clamp + z-score)
↓
Augmentação GPU (Kornia)
- Imagens armazenadas como
uint16(2 bytes/pixel), máscaras comoint16 - Crops não-sobrepostos do tamanho
crop_sizeextraídos das imagens originais - Estatísticas por banda (percentis p2/p98, média, desvio padrão) calculadas uma única vez e salvas em
stats.npz - Normalização executada na GPU (clamp percentílico + z-score por banda) — os workers de CPU fazem apenas leitura e cast de tipo
- A construção do memmap é automática na primeira execução; execuções subsequentes reutilizam os arquivos existentes
A máscara passa por duas defesas durante a construção do memmap:
- NoData: pixels com o valor
NoDataValuedeclarado nos metadados do GeoTIFF (da imagem ou da máscara) são marcados comoIGNORE_INDEX = 255. Esses pixels são ignorados pela loss (ignore_index=255) e nas métricas. - Índices fora de
[0, num_classes): quandonum_classesé passado ao builder, qualquer valor de classe residual (ex: anotações antigas com classe 4 em um setup atual de 3 classes) é convertido emIGNORE_INDEXe logado. Evita crash doF.one_hotnoFocalDiceLoss(CUDA device-side assert) e o silencioso enviesamento de métricas.
class 0 é classe válida (background), não NoData — é tratada como qualquer outra classe pela loss e pelas métricas.
O Prithvi foi pré-treinado em 6 bandas HLS, nesta ordem:
| Posição HLS | Banda |
|---|---|
| 0 | BLUE |
| 1 | GREEN |
| 2 | RED |
| 3 | NIR_NARROW |
| 4 | SWIR_1 |
| 5 | SWIR_2 |
No YAML, o campo dataset.model_bands declara qual banda HLS cada posição da sua entrada representa. O Prithvi11BandsModel lê isso e copia o peso pré-treinado correto para cada canal de patch_embed.proj, em vez de assumir cegamente a ordem [BLUE, GREEN, RED, ...].
Landsat 8/9 — B4, B5, B6, B7 (canônico para exploração mineral, clay/silica via SWIR):
dataset:
num_bands: 4
model_bands: ['RED', 'NIR_NARROW', 'SWIR_1', 'SWIR_2']Sentinel-2 — B2, B3, B4, B8A, B11, B12:
dataset:
num_bands: 6
model_bands: ['BLUE', 'GREEN', 'RED', 'NIR_NARROW', 'SWIR_1', 'SWIR_2']- Se
model_bandsnão for declarado: o modelo cai no comportamento antigo — copia as primeirasNbandas do HLS na ordem default. Bandas extras (além de 6) são inicializadas com a média dos pesos HLS. - Se um nome em
model_bandsnão existir em_HLS_BAND_INDEX(ex:RED_EDGE): hoje o modelo levantaValueError. Para sensores com bandas fora do HLS (Red Edge do Sentinel-2, Coastal Aerosol), seria preciso estender_HLS_BAND_INDEXcom uma estratégia de inicialização para a banda nova.
Trocar de sensor = mudar
num_bands+model_bandsno YAML, apagar./memmap_output/, e rodar. Sem mexer em código.
python -m scripts.prithvi.train_prithvi --config config/prithvi.yaml
# Retomar de checkpoint
python -m scripts.prithvi.train_prithvi --config config/prithvi.yaml --resume
python -m scripts.prithvi.train_prithvi --config config/prithvi.yaml --resume path/to/checkpoint.pthCaracterísticas do loop:
- Mixed precision (AMP FP16) com gradient clipping (
max_norm=1.0) - Scheduler: Linear warmup + Cosine Annealing
- Loss:
FocalDiceLosscomclass_weights,focal_gamma,focal_weight,dice_weight,ignore_index=255 - LR diferenciado:
learning_ratepara o backbone pré-treinado edecoder_lr(≥ backbone LR) para o decoder/head oversample_rare_classes: true— usaWeightedRandomSamplerpara visitar mais frequentemente patches com classes minoritárias- Métricas:
JaccardIndex(mIoU) eAccuracyvia TorchMetrics
python -m scripts.unet.train_unet --config config/unet.yamlCaracterísticas do loop:
- Mixed precision (AMP FP16) com gradient accumulation (effective batch =
batch_size × grad_accum) - Scheduler: Cosine Annealing
- Loss:
FocalDiceLosscomclass_weights,ignore_index=255 - Augmentação GPU via Kornia: flip, rotação 90°, brilho/contraste, affine, blur gaussiano
torch.channels_lastno modelo para melhor throughput em GPUs com Tensor Cores
Calculados automaticamente por compute_class_weights():
- Frequência inversa com suavização logarítmica
- Clipping para evitar pesos extremos
- Pixels marcados como
IGNORE_INDEX(NoData ou out-of-range) são excluídos da contagem - Normalizados para média 1.0 sobre as classes ativas
python -m scripts.prithvi.predict_prithvi- Carrega modelo TorchScript em FP16 de
./model/prithvi_production_fp16.pt - Processa múltiplas imagens em batch cross-image (sem recarregar o modelo entre arquivos)
- Normalização z-score por banda na GPU antes da inferência
- Saída: GeoTIFF com máscara de classes em
./predicoes/
python -m scripts.unet.predict_unet- Inferência via ONNX Runtime com
CUDAExecutionProvider - Tiling com overlap e média ponderada por mapa gaussiano centrado (elimina artefatos de borda)
- Normalização clamp + z-score na GPU antes de cada batch
- Saída: GeoTIFF comprimido (LZW, tiled) em
./Masks/
python -m scripts.prithvi.prithvi_fp16Gera dois artefatos:
./model/best_prithvi_*_fp16.pth— state dict em FP16./model/prithvi_production_fp16.pt— TorchScript standalone (sem dependência desrc/)
python -m scripts.unet.export_onnx_unet| Diretório | Conteúdo |
|---|---|
./data/img/ |
GeoTIFFs de entrada (N bandas, dimensões definidas em img_size) |
./data/mask/ |
Máscaras *_mask.tif (valores em [0, num_classes) ou NoData) |
./memmap_output/ |
*.dat + stats.npz (gerado automaticamente) |
./model/ |
Pesos Prithvi (.pth, FP16 .pt) |
./output/ |
Melhor checkpoint U-Net, CSV de validação, confusion matrix |
./metrics_prithvi/ |
Curvas de treino do Prithvi (PNG) |
./metrics_unet/ |
Curvas de treino do U-Net (PNG) |
./predicoes/ |
GeoTIFFs de saída da inferência Prithvi |
./Masks/ |
GeoTIFFs de saída da inferência U-Net |
Os dois configs são totalmente genéricos: num_bands, num_classes, img_size, crop_size, hiperparâmetros e arquitetura são definidos no YAML. O código não tem nenhum valor hard-coded de classe ou banda.
paths:
images_dir: "<dir das imagens>"
labels_dir: "<dir das máscaras>"
memmap_dir: "<dir do memmap>"
save_model_path: "<caminho do checkpoint>"
dataset:
img_size: [<H>, <W>]
crop_size: <int>
num_bands: <N> # qualquer N >= 1
num_classes: <K> # qualquer K >= 2 (inclui background se houver)
val_size: <float> # fração de validação
num_cores: <int> # workers do DataLoader
# Opcional: mapeia cada posição da entrada para a banda HLS correspondente
# (ver seção "Mapeamento de Bandas"). Se omitido, usa fallback.
model_bands: ['<HLS_BAND>', ...] # len(model_bands) == num_bands
training:
epochs: <int>
patience: <int> # early stopping
learning_rate: <float> # LR do backbone pré-treinado
decoder_lr: <float> # LR do decoder/head (geralmente ≥ backbone LR)
batch_size: <int>
weight_decay: <float>
focal_gamma: <float> # γ da Focal — foco em pixels difíceis
focal_weight: <float> # peso do termo Focal na loss
dice_weight: <float> # peso do termo Dice na loss
oversample_rare_classes: <bool> # WeightedRandomSampler
model_size: <tiny | 100m | 300m | 600m>paths:
images_dir: "<dir das imagens>"
labels_dir: "<dir das máscaras>"
csv_validation: "<csv de validação>"
checkpoint_path: "<dir de saída>"
best_model_path: "<caminho do melhor modelo>"
memmap_dir: "<dir do memmap>"
dataset:
img_size: [<H>, <W>]
crop_size: <int>
num_bands: <N>
num_classes: <K>
val_size: <float>
num_cores: <int>
training:
epochs: <int>
patience: <int>
learning_rate: <float>
batch_size: <int>
grad_accum: <int> # effective batch = batch_size × grad_accum
model_size: <resnet34 | resnet50>A cada época são registrados e plotados:
| Métrica | Descrição |
|---|---|
train_loss / val_loss |
Loss média por época |
val_miou |
mIoU médio das classes de foreground (exclui background da média final) |
val_acc |
Pixel accuracy |
lr |
Learning rate atual |
time |
Tempo de execução por época (s) |
gpu_mem |
Pico de memória GPU alocada (GB) |
Gráficos salvos em ./metrics_prithvi/ e ./metrics_unet/ após cada época.
Ao final do treino, evaluate_model() gera:
- CSV com precision, recall, F1, IoU e Dice por classe
- Confusion matrix em PNG
- Kappa de Cohen e Weighted IoU globais
src/fix_terratorch.pydeve ser executado uma vez após a instalação — corrige a constanteSENTINEL2_ALL_SOFTCON → SENTINEL2_ALL_MOCOna biblioteca TerraTorch instalada- Todos os scripts são invocados como módulos (
python -m scripts.X.Y), não diretamente - O Prithvi adapta o
patch_embed.projoriginal (6 canais HLS) para o número de bandas declarado emnum_bands, copiando o peso pré-treinado de cada banda HLS correta para a posição certa viamodel_bands IGNORE_INDEX = 255emsrc/dataset.pyé o sentinela para NoData e índices inválidos — mantido consistente entre o builder de memmap, aFocalDiceLosse as métricas- Não há testes unitários; a avaliação é integrada ao loop de treino e produz CSVs e plots automaticamente
Os números abaixo foram coletados em uma configuração específica (11 bandas, 19 classes). Servem como referência relativa entre arquiteturas — valores absolutos variam com o número de bandas, classes, tamanho da imagem e qualidade do dataset.
Dispositivo: cuda
GPU: NVIDIA GeForce RTX 3060 Laptop GPU
batch_size: 4 | n_batches inferência: 20
| Métrica | Prithvi-Tiny | Prithvi-100M | UNet-R34 | UNet-R50 |
|---|---|---|---|---|
| Parâmetros totais (M) | 13.2 | 96.9 | 24.7 | 75.5 |
| Parâmetros treináveis (M) | 13.2 | 96.9 | 24.7 | 75.5 |
| Checkpoint .pth (MB) | 151.08 | 1109.38 | 94.24 | 288.54 |
| Checkpoint presente | OK | OK | OK | OK |
| Métrica | Prithvi-Tiny | Prithvi-100M | UNet-R34 | UNet-R50 |
|---|---|---|---|---|
| ms / batch | 76.8 | 340.9 | 101.9 | 324.3 |
| ms / amostra | 19.2 | 85.2 | 25.5 | 81.1 |
| amostras / segundo | 52.1 | 11.7 | 39.2 | 12.3 |
| pico GPU — inf (GB) | 0.38 | 0.87 | 1.31 | 2.42 |
| speedup vs Prithvi-Tiny | 1.00x | 0.23x | 0.75x | 0.24x |
| Métrica | Prithvi-Tiny | Prithvi-100M | UNet-R34 | UNet-R50 |
|---|---|---|---|---|
| Épocas treinadas | 48 | 38 | 58 | 43 |
| Tempo médio / época (s) | 8.1 | 32.3 | 13.0 | 208.5 |
| Tempo total treino (min) | 6.5 | 20.5 | 12.6 | 149.3 |
| Mem GPU média treino (GB) | 1.95 | 5.23 | 3.26 | 6.77 |
| Mem GPU pico treino (GB) | 1.95 | 5.24 | 3.26 | 6.77 |
| LR final (média 3 ep.) | 7.47e-06 | 1.93e-05 | — | — |
| Early stopping | Não | Não | Sim (ep 58) | Não |
| Métrica | Prithvi-Tiny | Prithvi-100M | UNet-R34 | UNet-R50 |
|---|---|---|---|---|
| Best mIoU | 0.6011 | 0.6495 | 0.2810 | 0.3005 |
| → época | 46 | 37 | 57 | 39 |
| Best pixel accuracy | 0.8975 | 0.9089 | — | — |
| Best val loss | 0.3851 | 0.3457 | 1.0235 | 1.0345 |
| → época | 48 | 38 | 48 | 43 |
| mIoU final (last epoch) | 0.5935 | 0.6489 | 0.2805 | 0.2779 |
| Loss final (last epoch) | 0.3851 | 0.3457 | 1.0339 | 1.0345 |
- Prithvi-Tiny é o mais rápido: 52.1 amostras/s, apenas 0.38 GB VRAM
- UNet-R34 oferece bom compromisso: 39.2 amostras/s, 1.31 GB VRAM
- Prithvi-100M e UNet-R50 são similares em velocidade (~12 amostras/s), mas R50 usa 2.8× mais VRAM
- Prithvi-Tiny: mais rápido (8.1s/época), menor memória (1.95 GB)
- UNet-R34: 13s/época, 3.26 GB — extremamente eficiente
- Prithvi-100M: 32.3s/época, 5.23 GB — moderado
- UNet-R50: 208.5s/época, 6.77 GB — o mais pesado
- Prithvi-100M: melhor qualidade absoluta (mIoU 0.6495)
- Prithvi-Tiny: segundo lugar (mIoU 0.6011)
- UNet-R50: terceiro (mIoU 0.3005)
- UNet-R34: último (mIoU 0.2810)
- Prithvi foundation models dominam em qualidade (≈2× melhor mIoU que ResNets) — diferença esperada quando se aproveita pré-treinamento HLS via
model_bands - ResNets são mais rápidos no treino, mas menos precisos no nosso domínio
- Prithvi-Tiny oferece o melhor custo-benefício geral
- UNet-R50 não compensa o custo: treina 16× mais devagar que R34 para apenas 7% mais mIoU
