TSViT โ€” PASTIS24 checkpoints (5-fold)

Per-fold best.pth checkpoints (model weights only, state_dict) from training TSViT on the PASTIS24 benchmark (24x24-patch split of PASTIS), reproducing Tarasiou et al., "ViTs for SITS: Vision Transformers for Satellite Image Time Series", CVPR 2023.

Each checkpoint is the epoch with the best eval IoU during training for that fold, using configs/PASTIS24/TSViT_fold<N>.yaml from the repo (architecture: TSViT, dim=128, 4 temporal + 4 spatial transformer layers, ~1.7M params).

Files

File Fold Test OA Test mIoU
tsvit_fold1_best.pth 1 83.21 64.94
tsvit_fold2_best.pth 2 84.17 66.96
tsvit_fold3_best.pth 3 83.43 65.38
tsvit_fold4_best.pth 4 83.08 63.47
tsvit_fold5_best.pth 5 84.10 67.27
5-fold average 83.60 65.60
Paper (5-fold average) 83.4 65.4

Test-set metrics computed with masked cross-entropy loss (background class masked), overall accuracy (OA, pixel-averaged) and mean IoU (mIoU, class-averaged), matching the paper's evaluation protocol.

Usage

Each file is a raw PyTorch state_dict (not a full checkpoint dict). Load it into a TSViT model instance from the TSViT repo:

import torch
from models.TSViT.TSViTdense import TSViT

model_config = {
    'img_res': 24, 'max_seq_len': 60, 'num_channels': 11, 'num_classes': 19,
    'patch_size': 2, 'dim': 128, 'temporal_depth': 4, 'spatial_depth': 4,
    'heads': 4, 'pool': 'cls', 'dim_head': 32, 'scale_dim': 4,
    'dropout': 0., 'emb_dropout': 0.,
}
net = TSViT(model_config)
net.load_state_dict(torch.load("tsvit_fold1_best.pth", map_location="cpu"))
net.eval()
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support