Seeking Advice: Using Prithvi-EO-2.0 for Road Segmentation in VHR (2m) Imagery #580
Environment
Problem DescriptionI'm attempting to use Prithvi-EO-2.0-600M for road detection on high-resolution imagery, but facing extremely poor performance. As shown in the attached metrics and test run, the model isn't detecting roads as good as I was hoping for after over 275 epochs (best model was at epoch 48): Attempted SolutionsI've already identified and addressed class imbalance issues in my validation metrics. Initially, the model was achieving seemingly high overall scores by simply predicting "road" everywhere. I changed the validation monitor from "val/Multiclass_Jaccard_Index" to "val/multiclassjaccardindex_1" to focus specifically on road detection performance rather than overall accuracy. Despite this adjustment to properly monitor road class performance, the model still fails to detect roads in the imagery. The TensorBoard visualization shows a significant gap between training and validation metrics, with training quickly reaching high values while validation remains unstable. ConfigurationI'm using a UNetDecoder with the Prithvi-EO-2.0-600M backbone (configuration file in the bottom). Key settings:
Questions
Thank you for developing and sharing both the Prithvi-EO-2.0 model and the TerraTorch platform :) TerrTorch configuration file:# TerraTorch configuration for road detection with Prithvi-EO-2.0
seed_everything: 42
trainer:
accelerator: auto
strategy: auto
devices: auto
num_nodes: 1
precision: 16-mixed
logger: true
callbacks:
- class_path: RichProgressBar
- class_path: LearningRateMonitor
init_args:
logging_interval: epoch
- class_path: ModelCheckpoint
init_args:
dirpath: output/roads/checkpoints
mode: max
monitor: val/multiclassjaccardindex_1
filename: best-ny-{epoch:02d}
save_top_k: 1
verbose: true
max_epochs: -1
log_every_n_steps: 5
default_root_dir: output/roads/
data:
class_path: GenericNonGeoSegmentationDataModule
init_args:
batch_size: 1
num_workers: 4
dataset_bands: # Dataset bands
- RED
- GREEN
- BLUE
- NIR
output_bands: # Model input bands
- RED
- GREEN
- BLUE
- NIR
rgb_indices:
- 0
- 1
- 2
train_data_root: /data/terratorch/test/output_dir/data
val_data_root: /data/terratorch/test/output_dir/data
test_data_root: /data/terratorch/test/output_dir/data
train_split: /data/terratorch/test/output_dir/data/splits/train.txt
val_split: /data/terratorch/test/output_dir/data/splits/val.txt
test_split: /data/terratorch/test/output_dir/data/splits/test.txt
img_grep: "*_merged.tif"
label_grep: "*.mask.tif"
means:
- 0.013455652347789112 # RED
- 0.019870659345094948 # GREEN
- 0.024954085893903428 # BLUE
- 0.054250626392519756 # NIR
stds:
- 0.027177061600580688 # RED
- 0.036014259893003754 # GREEN
- 0.044061468237158556 # BLUE
- 0.11263794821862147 # NIR
num_classes: 2
train_transform:
- class_path: albumentations.D4
- class_path: ToTensorV2
no_data_replace: 0
no_label_replace: -1
model:
class_path: terratorch.tasks.SemanticSegmentationTask
init_args:
model_factory: EncoderDecoderFactory
model_args:
backbone: prithvi_eo_v2_600
backbone_pretrained: true
backbone_img_size: 1376
backbone_coords_encoding: []
backbone_bands:
- RED
- GREEN
- BLUE
- NIR
necks:
- name: SelectIndices
indices: [7, 15, 23, 31]
- name: ReshapeTokensToImage
- name: LearnedInterpolateToPyramidal
decoder: UNetDecoder
decoder_channels: [512, 256, 128, 64]
head_dropout: 0.1
num_classes: 2
loss: ce
ignore_index: -1
freeze_backbone: false
freeze_decoder: false
optimizer:
class_path: torch.optim.AdamW
init_args:
lr: 5.e-5
weight_decay: 0.1
lr_scheduler:
class_path: ReduceLROnPlateau
init_args:
monitor: val/loss
factor: 0.5
patience: 5 |
Replies: 1 comment
|
Hi, @FerdinandKlingenberg . Thank you for your experiment and feedback. I'm not really a specialist in the task you are trying to execute (certainly, not as you are), but maybe I can give a few suggestions/comments:
|



Hi, @FerdinandKlingenberg . Thank you for your experiment and feedback. I'm not really a specialist in the task you are trying to execute (certainly, not as you are), but maybe I can give a few suggestions/comments:
batch_size=1has some effect.