Download main_method/code/config.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 13.8 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/config.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/config.py
-
curl -L -o config.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/config.py
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", | |
| } | |