TSRDA / main_method /code /config.py
Dhruv1000's picture
Organize complete final models, all ablations, logs and checkpoints with visual guides (part 7)
71d64bb verified
Raw History Blame Contribute Delete
13.8 kB
"""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",
}