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",
}