File size: 13,795 Bytes
71d64bb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 | """Final method configuration. Sampling and seed match STCLN; preprocessing differences are documented in FINAL_REPORT.md."""
import os
from pathlib import Path
# βββ Paths ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
# Derived from this file's own location so the project moves between machines
# without edits. Current layout on this box:
#
# /home/ubuntu/tsrda/ <- REPO_ROOT = EXP_ROOT (code + checkpoints)
# /home/ubuntu/test/PASTIS/ <- PASTIS_ROOT (dataset)
#
# PASTIS_ROOT is searched in the order below; the first existing hit wins.
# Override either root by env var, e.g.
# PASTIS_ROOT=/mnt/data/PASTIS python3 check_data.py
REPO_ROOT = Path(__file__).resolve().parents[2]
_PASTIS_CANDIDATES = (
# Kaggle notebook locations (PASTIS-R dataset attached as an input)
Path("/home/ubuntu/PASTIS"),
Path("/home/ubuntu/PASTIS"),
REPO_ROOT / "PASTIS-R",
REPO_ROOT / "PASTIS",
REPO_ROOT.parent / "PASTIS",
)
def _find_pastis_root():
env = os.environ.get("PASTIS_ROOT")
if env:
return Path(env)
for cand in _PASTIS_CANDIDATES:
if (cand / "metadata.geojson").is_file():
return cand
raise FileNotFoundError(
"PASTIS dataset not found. Looked for metadata.geojson in:\n "
+ "\n ".join(str(c) for c in _PASTIS_CANDIDATES)
+ "\nSet the PASTIS_ROOT environment variable to the dataset directory."
)
PASTIS_ROOT = _find_pastis_root()
DATA_S2_DIR = PASTIS_ROOT / "DATA_S2"
ANNOT_DIR = PASTIS_ROOT / "ANNOTATIONS"
META_PATH = PASTIS_ROOT / "metadata.geojson"
NORM_PATH = PASTIS_ROOT / "NORM_S2_patch.json"
EXP_ROOT = Path(os.environ.get("EXP_ROOT", REPO_ROOT))
CKPT_DIR_PRE = EXP_ROOT / "checkpoints" / "pretrain"
CKPT_DIR_FT = EXP_ROOT / "checkpoints" / "finetune"
# βββ Reproducibility ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
SEED = 3407 # [OFFICIAL] torch.manual_seed in
# finetuning_STCLN.py:25 and test_STCLN.py:11
VARIANCE_SEEDS = (3407, 42, 1234)
# [RESEARCH] additional seeds used for
# seed-variance runs. Pass with --seed.
# βββ Model dimensions βββββββββββββββββββββββββββββββββββββββββββββββββββββββ
N_CHANNELS = 10 # [OFFICIAL]
D_MODEL = 256 # [OFFICIAL]
N_HEADS = 8 # [OFFICIAL]
N_CLASSES = 20 # [OFFICIAL] 0=background, 19=void, 1-18 scored
T_PAD = 48 # inert; UTAE_ASPP absorbs it for constructor
# compatibility and the encoder self-pads
# T-SRDA schedule: KV reduction ratio per block (replaces Swin/DaViT windows).
R_SCHEDULE = [4, 4, 4] # [TSRDA] the architectural variable
# Legacy temporal-Swin params β kept ONLY so the train scripts' constructor
# call sites compile unchanged. TemporalSRDAEncoder ignores them.
TSE_WINDOWS = [4, 8, 16]
TSE_SHIFTS = [0, 4, 0]
# βββ Protocol: geometry βββββββββββββββββββββββββββββββββββββββββββββββββββββ
PATCH_SIZE = 128 # [OFFICIAL] PASTIS native patch
CROP_SIZE = 32 # [OFFICIAL] PATCH_SIZE // 4
REF_DATE_STR = "2018-09-01" # [OFFICIAL]
USE_INDEX_POSITIONS = True # [OFFICIAL] the model is fed torch.arange(T),
# NOT real day offsets. See
# finetuning_STCLN.py:198 β every forward call
# passes `torch.tensor(range(T))`. Real days are
# still returned by the dataset for analysis.
# βββ Protocol: pretrain split βββββββββββββββββββββββββββββββββββββββββββββββ
PRETRAIN_FOLDS = [5] # [OFFICIAL] 496 patches
PRETRAIN_CROP_GRID = 4 # [OFFICIAL] inner `for i in range(4): for j in
# range(4)` loop -> 16 crops per patch, each a
# SEPARATE optimizer step.
# 496 / 4 = 124 batches x 16 = 1,984 steps/ep
CROPS_PER_PATCH = PRETRAIN_CROP_GRID ** 2 # 16
# βββ Protocol: finetune splits (hardcoded patch IDs, not a fold filter) βββββ
# The official script builds the dataset with folds=[1] / folds=[2] and then
# OVERWRITES dataset.id_patches with these literal lists. Verified against
# metadata.geojson: every train ID is fold 1, every val ID is fold 2, the two
# sets are disjoint, and neither touches the fold-4 test set or fold-5
# pretrain pool. Duplicates are intentional and are KEPT β they make those
# patches count twice per epoch.
FT_CROP_IJ = [(0, 0), (1, 1)] # [OFFICIAL] finetuning_STCLN.py:185 β only the
# first two diagonal crops of the 4x4 grid.
# 76 / 2 = 38 batches x 2 crops = 76 steps/ep
TRAIN_PATCH_IDS = [
10279, 40282, 10174, 40538, 10457, 30088, 10151, 30602, 40275, 20052,
20153, 10056, 20147, 40420, 10450, 20614, 20167, 40446, 20251, 20103,
20613, 20349, 20384, 40337, 40163, 10030, 10068, 10399, 30327, 10289,
30351, 40198, 20167, 20153, 20207, 20203, 10000, 10354, 20214, 10007,
10392, 30055, 10054, 20009, 10040, 10147, 10110, 10105, 20398, 20469,
20203, 20254, 30402, 30013, 40200, 30689, 20115, 20356, 20345, 20466,
40033, 10153, 30153, 30012, 40340, 30140, 40214, 10129, 30211, 20078,
30109, 30013, 20116, 20243, 20085, 40271,
] # [OFFICIAL] 76 entries, 72 unique (20153, 20167, 20203, 30013 appear 2x)
VAL_PATCH_IDS = [
10127, 40438, 30643, 30619, 30436, 10306, 10127, 30226, 40027, 40279,
40085, 40455, 40426, 40444, 40298, 40231, 20498, 40298, 20149, 20315,
20340, 20618, 20412, 20231, 40006, 40138, 40186, 10087, 40430, 30194,
10247, 40383, 20364, 20128, 20060, 20118, 40001, 20080, 40063, 30000,
40063, 40000, 40348, 10021, 40001, 20023, 30000, 20441, 20235, 20344,
20297, 20576, 30014, 30275, 30298, 30083, 20206, 20443, 20423, 20380,
30064, 30084, 30145, 10107, 10225, 40558, 10383, 20131, 30010, 20022,
30295, 20042, 20031, 40046, 20358, 30142,
] # [OFFICIAL] 76 entries, 71 unique (10127, 40298, 40001, 40063, 30000 2x)
# βββ Protocol: test split βββββββββββββββββββββββββββββββββββββββββββββββββββ
TEST_FOLDS = [4] # [OFFICIAL] 482 patches, spatially disjoint
# from folds 1/2/5 with a 1 km buffer
TEST_FULL_PATCH = True # [OFFICIAL] test_STCLN.py:55 does
# `split = W // 1` -> the whole 128x128 patch
# is evaluated, NOT 32x32 crops.
# 482 patches ~ 482 x 16 = 7,712 crop-equiv
# βββ Pretrain hyperparameters βββββββββββββββββββββββββββββββββββββββββββββββ
PRE_EPOCHS = 100 # [OFFICIAL]
PRE_BATCH = 4 # [OFFICIAL] PATCHES per batch, not crops. The
# inner 4x4 loop then does 16 separate
# optimizer steps on 4 crops each. Raising this
# changes the optimization trajectory, not just
# memory β do not tune it for the headline run.
PRE_LR = 1e-4 # [OFFICIAL] flat, no schedule
PRE_WD = 0.0 # [OFFICIAL] AdamW with wd=0 reduces to Adam
# exactly. Set to 0.0 (from 0.01) so the only
# variable vs A_linear is the temporal encoder.
PRE_CLIP = 5.0 # [OFFICIAL]
MASK_RATIO = 0.4 # [OFFICIAL]
NDVI_THRESH = 0.2 # [OFFICIAL]
CLOUD_GATE = 0.9 # [OFFICIAL] STCLN.py:200 `mask[clusterLmean
# <= 0.9] = 1`. clusterLmean averages the NDVI
# indicator over H and W, so this is a PER-FRAME
# gate: a timestep whose frame is <=90%
# vegetated is left ENTIRELY visible. It is NOT
# a per-pixel vegetation filter. Measured on 20
# fold-5 patches: fires on 90.6% of frames,
# leaving 96.2% of all values visible.
PRE_SAVE_EVERY = 20 # [TSRDA] keep a numbered milestone every N
# epochs (0, 19, 39, 59, 79, 99 at N=20).
# Independent of this, `latest.tar` is written
# EVERY epoch with optimizer + scaler + RNG
# state, so a crash costs at most one epoch
# (~26 min) rather than up to N (~8.7 h).
DEEP_SUP = False # [OFFICIAL] the deep-supervision term is
# commented out upstream
# (pretraining_STCLN.py:141
# `l = loss(output, target)#+loss(x2,target)`).
# βββ Finetune hyperparameters βββββββββββββββββββββββββββββββββββββββββββββββ
FT_EPOCHS = 100 # [OFFICIAL]
FT_BATCH = 2 # [OFFICIAL] PATCHES per batch (x2 crop-steps)
FT_LR = 1e-4 # [OFFICIAL] flat Adam, NO scheduler
FT_WD = 0.0 # [OFFICIAL] AdamW with wd=0 == Adam
FT_AUGMENT = False # [OFFICIAL] no augmentation in finetuning_STCLN
FT_SCHEDULER = None # [OFFICIAL] flat LR, no schedule.
FT_PATIENCE = 0 # [OFFICIAL] no early stopping (grep: 0 matches
# in the reference). 0 disables it entirely.
# MUST stay 0: the primary report is the
# FIXED-EPOCH ep99 number. Early stopping on
# 152 val crops manufactured a 147x variance
# effect that vanished at fixed epoch, and if
# it fires at ep58 there is no ep99 checkpoint
# to compare against A_linear at all.
# βββ Evaluation βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
EVAL_BATCH = int(os.environ.get("EVAL_BATCH", 2)) # 2 fits a 16 GB T4 (9.5 GB peak measured)
# [HW] official test_STCLN.py uses 1. Full-patch
# inference is numerically identical at any
# batch size (nothing in eval mode depends on
# the batch), so this is a pure memory/speed
# knob with no effect on the reported metrics.
# MEASURED for T-SRDA on this 23.6 GB L4 against
# the worst case (fold 4 runs to T=61):
# batch 1 -> 5.05 GB peak 482 batches
# batch 2 -> 9.50 GB peak 241 batches
# batch 4 -> 16.86 GB peak 121 batches
# Default 4 gives the protocol's 121 batches and
# fits with ~6.7 GB spare ON A FREE CARD.
# asr-gpu.service is systemd-`enabled` on this
# box and takes ~8 GB when it comes back after a
# reboot, leaving ~15.6 GB β under batch 4. If
# eval OOMs, just rerun with EVAL_BATCH=2; the
# numbers are unchanged.
EVAL_TTA = False # [OFFICIAL] test_STCLN.py has no TTA
# βββ PASTIS class nomenclature ββββββββββββββββββββββββββββββββββββββββββββββ
# THE single source of truth for this repo. The list that shipped in the
# original tsrda scripts was wrong for 12 of 20 classes (it invented
# "Protein peas", "Flax", "Sugar beet", "Vineyards", "Nuts" and mislabelled
# 13 as Soy / 17 as Potatoes). Verified identical to PhenoProto-SSL's
# splits.PASTIS_CLASSES, which preflight.py asserts against the official
# PASTIS benchmark list. Never redefine this locally in another module.
PASTIS_CLASSES = {
0: "Background",
1: "Meadow",
2: "Soft winter wheat",
3: "Corn",
4: "Winter barley",
5: "Winter rapeseed",
6: "Spring barley",
7: "Sunflower",
8: "Grapevine",
9: "Beet",
10: "Winter triticale",
11: "Winter durum wheat",
12: "Fruits, vegetables, flowers",
13: "Potatoes",
14: "Leguminous fodder",
15: "Soybeans",
16: "Orchard",
17: "Mixed cereal",
18: "Sorghum",
19: "Void label",
}
|