From bc5373a07af061d2e8c85757a0ec79635971177a Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Thu, 2 Jul 2026 10:22:35 +0200 Subject: [PATCH 01/26] add config for nudging --- config/varda-single-1.0-nudge.yaml | 144 ++++++++++++++++++ ...tidataset-forecaster-global-ich1-oper.yaml | 11 +- ...oral-downscaler-global_trimedge_multi.yaml | 6 + 3 files changed, 160 insertions(+), 1 deletion(-) create mode 100644 config/varda-single-1.0-nudge.yaml diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml new file mode 100644 index 00000000..bc63bf15 --- /dev/null +++ b/config/varda-single-1.0-nudge.yaml @@ -0,0 +1,144 @@ +# yaml-language-server: $schema=../workflow/tools/config.schema.json +description: | + Evaluate skill of Varda-single-1.0 against ground observations. + +config_label: varda-single-1.0 + +dates: + start: 2025-03-01T00:00 + end: 2025-03-03T00:00 + frequency: 24h + +runs: + - temporal_downscaler: + checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 + label: Varda-single-1.0 + steps: 0/24/1 + config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml + extra_requirements: + - anemoi-datasets==0.5.35 + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@bac161e7d5184335614ea6723987acc30e9d9b1d + - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + forecaster: + checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 + config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml + steps: 0/24/6 + extra_requirements: + - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - baseline: + label: INCA + root: /store_new/mch/msclim/INCA + steps: 0/6/1 + - baseline: + label: ICON-CH2-CTRL + root: /store_new/mch/msopr/osm/ICON-CH2-EPS + steps: 0/24/1 + - baseline: + label: ICON-CH1-CTRL + root: /store_new/mch/msopr/osm/ICON-CH1-EPS + steps: 0/24/1 + +truth: + label: SwissMetNet + root: jretrievedwh:1,2 + # To verify against SwissMetNet observations from the DWH via jretrievedwh, + # set instead (requires jretrievedwh.py on $PATH and $OPR_HOME set): + # Other selectors: root: jretrievedwh:locations=ARO,KLO,LUG + # root: jretrievedwh:bbox=45.8,47.8,5.9,10.5 + # append ;stage=devt to target a non-prod DWH stage + +experiment: + params: + - T_2M + - TD_2M + - U_10M + - V_10M + - TOT_PREC + stratification: + regions: + - mittelland + - berge + - alpennordseite + - alpensuedseite + - jura + root: /store_new/mch/msopr/ml/regions/Prognoseregionen_LV95_20220517 + thresholds: + TOT_PREC: + gt: [0.0, 1, 5] + U_10M: + gt: [2.5, 5.0, 10.0] + V_10M: + gt: [2.5, 5.0, 10.0] + T_2M: + lt: [273.15] + gt: [288.15, 298.15] + dashboard: + stratification: + # - region + # - init_hour + - season + scorecards: + enabled: false + sections: + nowcasting: + baseline: INCA + lead_times: "0/6/1" + stratification: region + variables: + - "U_10M:RMSE,R2,ETS" + - "V_10M:RMSE,R2,ETS" + - "T_2M:RMSE,R2,ETS" + - "TOT_PREC:RMSE,R2,ETS" + short_range: + baseline: ICON-CH1-CTRL + lead_times: "6/24/6" + stratification: region + variables: + - "U_10M:RMSE,R2,ETS" + - "V_10M:RMSE,R2,ETS" + - "T_2M:RMSE,R2,ETS" + - "TOT_PREC:RMSE,R2,ETS" + medium_range: + baseline: ICON-CH2-CTRL + lead_times: "24/120/24" + stratification: region + variables: + - "U_10M:RMSE,R2,ETS" + - "V_10M:RMSE,R2,ETS" + - "T_2M:RMSE,R2,ETS" + - "TOT_PREC:RMSE,R2,ETS" + scoremaps: + enabled: false + +showcase: + params: + - T_2M + - U_10M + - V_10M + meteograms: + enabled: false + stations: [JUN] #, COV, GOR, WFJ, SAE, SAM, DAV, ZER, ANT, VSBAS, BRT, LTB, GOS, CEV, BIA] + animations: + fps: 0.5 + enabled: true + domains: + # - globe + # - europe + - alps + - icon-ch + - switzerland + +locations: + output_root: output/ + +profile: + executor: slurm + global_resources: + gpus: 16 + default_resources: + slurm_partition: "postproc" + cpus_per_task: 1 + mem_mb_per_cpu: 1800 + runtime: "1h" + gpus: 0 + jobs: 50 diff --git a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml index cd0c95a5..24f0fe38 100644 --- a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml +++ b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml @@ -1,4 +1,4 @@ -lead_time: 120h +lead_time: 24h write_initial_state: true allow_nans: true @@ -10,6 +10,15 @@ input: test: use_original_paths: true +# nudging is not performed by the forecaster +# pre_processors: +# - forward_transform_filter: +# nudge_toward_observation: +# path_to_observation: /scratch/mch/llanzila/sruc/evalml/output/data/observation/PeakWeather +# k: 3 +# power: 4.0 +# max_dist: 0.5 + post_processors: - accumulate_from_start_of_forecast: # accumulate tp from start of forecast accumulations: diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml index 1e9e237f..54f5bf0e 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml @@ -12,6 +12,12 @@ input: scale: 0.001 # convert from kg m-2 to m offset: 0 param: tp + - forward_transform_filter: + nudge_toward_observation: + path_to_observation: /scratch/mch/llanzila/sruc/evalml/output/data/observation/PeakWeather + k: 3 + power: 4.0 + max_dist: 0.5 namer: &namer rules: - - shortName: T From 6e13d168188a44107784496a2b88ef694c1b7991 Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Thu, 2 Jul 2026 15:05:51 +0200 Subject: [PATCH 02/26] update configs --- config/varda-single-1.0-nudge.yaml | 28 +++- ...tidataset-forecaster-global-ich1-oper.yaml | 11 +- ...oral-downscaler-global_trimedge_multi.yaml | 6 - ...ownscaler-global_trimedge_multi_nudge.yaml | 144 ++++++++++++++++++ 4 files changed, 167 insertions(+), 22 deletions(-) create mode 100644 resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml index bc63bf15..e7014adb 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-single-1.0-nudge.yaml @@ -1,6 +1,6 @@ # yaml-language-server: $schema=../workflow/tools/config.schema.json description: | - Evaluate skill of Varda-single-1.0 against ground observations. + Evaluate skill of Varda-rapid against ground observations. config_label: varda-single-1.0 @@ -10,6 +10,22 @@ dates: frequency: 24h runs: + - temporal_downscaler: + checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 + label: Varda-rapid-0.1 + steps: 0/24/1 + config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml + extra_requirements: + - anemoi-datasets==0.5.35 + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@0e1b8ac1a9eb3d4459f71f5a7418fc09c6305e1a + - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + forecaster: + checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 + config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml + steps: 0/24/6 + extra_requirements: + - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 label: Varda-single-1.0 @@ -17,7 +33,6 @@ runs: config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml extra_requirements: - anemoi-datasets==0.5.35 - - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@bac161e7d5184335614ea6723987acc30e9d9b1d - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 @@ -25,6 +40,7 @@ runs: steps: 0/24/6 extra_requirements: - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - baseline: label: INCA root: /store_new/mch/msclim/INCA @@ -78,7 +94,7 @@ experiment: # - init_hour - season scorecards: - enabled: false + enabled: true sections: nowcasting: baseline: INCA @@ -100,7 +116,7 @@ experiment: - "TOT_PREC:RMSE,R2,ETS" medium_range: baseline: ICON-CH2-CTRL - lead_times: "24/120/24" + lead_times: "24/24/24" stratification: region variables: - "U_10M:RMSE,R2,ETS" @@ -116,8 +132,8 @@ showcase: - U_10M - V_10M meteograms: - enabled: false - stations: [JUN] #, COV, GOR, WFJ, SAE, SAM, DAV, ZER, ANT, VSBAS, BRT, LTB, GOS, CEV, BIA] + enabled: true + stations: [JUN, KLO, LUG, GVE] #, COV, GOR, WFJ, SAE, SAM, DAV, ZER, ANT, VSBAS, BRT, LTB, GOS, CEV, BIA] animations: fps: 0.5 enabled: true diff --git a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml index 24f0fe38..51d2fb38 100644 --- a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml +++ b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml @@ -9,16 +9,7 @@ env: input: test: use_original_paths: true - -# nudging is not performed by the forecaster -# pre_processors: -# - forward_transform_filter: -# nudge_toward_observation: -# path_to_observation: /scratch/mch/llanzila/sruc/evalml/output/data/observation/PeakWeather -# k: 3 -# power: 4.0 -# max_dist: 0.5 - + post_processors: - accumulate_from_start_of_forecast: # accumulate tp from start of forecast accumulations: diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml index 54f5bf0e..1e9e237f 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml @@ -12,12 +12,6 @@ input: scale: 0.001 # convert from kg m-2 to m offset: 0 param: tp - - forward_transform_filter: - nudge_toward_observation: - path_to_observation: /scratch/mch/llanzila/sruc/evalml/output/data/observation/PeakWeather - k: 3 - power: 4.0 - max_dist: 0.5 namer: &namer rules: - - shortName: T diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml new file mode 100644 index 00000000..7ae39b72 --- /dev/null +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml @@ -0,0 +1,144 @@ +runner: temporal_downscaler + +input: + cutout: + - lam_0: + grib: + path: forecaster/20* + pre_processors: + - forward_transform_filter: + # does not have an effect if the temporal downscaler does not use tp as prognostic variable + rescale: + scale: 0.001 # convert from kg m-2 to m + offset: 0 + param: tp + - forward_transform_filter: + nudge_toward_observation: + backend: jretrieve + jretrieve_bbox: [45.7, 48.0, 5.8, 10.8] + jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src + k: 3 + power: 4.0 + max_dist: 0.5 + # nudge_toward_observation: + # path_to_observation: /scratch/mch/llanzila/sruc/evalml/output/data/observation/PeakWeather + # k: 3 + # power: 4.0 + # max_dist: 0.5 + namer: &namer + rules: + - - shortName: T + - t_{level} + - - shortName: U + - u_{level} + - - shortName: V + - v_{level} + - - shortName: W + - w_{level} + - - shortName: QV + - q_{level} + - - shortName: FI + - z_{level} + - - shortName: PMSL + - msl + - - shortName: FIS + - z + - - shortName: PS + - sp + - - shortName: T_2M + - 2t + - - shortName: TD_2M + - 2d + - - shortName: T_G + - skt + - - shortName: U_10M + - 10u + - - shortName: V_10M + - 10v + - - shortName: FR_LAND + - lsm + - - shortName: TOT_PREC + - tp + - global: + grib: + path: forecaster/ifs* + namer: *namer + +constant_forcings: + test: + use_original_paths: true + +patch_metadata: resources/sgm-temporal-downscaler-ich1-oper-patch.yaml + +post_processors: + - accumulate_from_start_of_forecast: # accumulate tp from start of forecast + accumulations: + - tp + - forward_transform_filter: + rescale: + scale: 1000 # convert units from m to kg m-2 + offset: 0 + param: tp + +output: + tee: + - grib: + path: grib/{dateTime}_{step:03}.grib + encoding: + typeOfGeneratingProcess: 2 + templates: + samples: resources/templates_index_icon.yaml + post_processors: + - extract_mask: # removes global points + mask: "lam_0/cutout_mask" + as_slice: true + # here, the trimedge mask can be specified when available + + - grib: + path: grib/ifs-{dateTime}_{step:03}.grib + encoding: + typeOfGeneratingProcess: 2 + templates: + samples: resources/templates_index_ifs.yaml + post_processors: + - extract_mask: # removes lam points + mask: "lam_0/cutout_mask" + as_slice: true + inverse: true + - assign_mask: # fill local/global overlapping points with nan + mask: "global/cutout_mask" + modifiers: + - patches: + - variable: + U_10M: {"param": 165, "shortName": "10u"} + 10u: {"param": 165, "shortName": "10u"} + V_10M: {"param": 166, "shortName": "10v"} + 10v: {"param": 166, "shortName": "10v"} + TD_2M: {"param": 168, "shortName": "2d"} + 2d: {"param": 168, "shortName": "2d"} + T_2M: {"param": 167, "shortName": "2t"} + 2t: {"param": 167, "shortName": "2t"} + FR_LAND: {"param": 172, "shortName": "lsm"} + lsm: {"param": 172, "shortName": "lsm"} + PMSL: {"param": 151, "shortName": "msl"} + msl: {"param": 151, "shortName": "msl"} + PS: {"param": 134, "shortName": "sp"} + sp: {"param": 134, "shortName": "sp"} + SSO_SIGMA: {"param": 163, "shortName": "slor"} + slor: {"param": 163, "shortName": "slor"} + SSO_STDH: {"param": 160, "shortName": "sdor"} + sdor: {"param": 160, "shortName": "sdor"} + TOT_PREC: {"param": 228, "shortName": "tp"} + tp: {"param": 228, "shortName": "tp"} + z: {"param": 129, "shortName": "z"} + "^q_(\\d+)$": {"param": 133, "shortName": "q"} + "^t_(\\d+)$": {"param": 130, "shortName": "t"} + "^u_(\\d+)$": {"param": 131, "shortName": "u"} + "^v_(\\d+)$": {"param": 132, "shortName": "v"} + "^w_(\\d+)$": {"param": 135, "shortName": "w"} + "^z_(\\d+)$": {"param": 129, "shortName": "z"} + +# silenced due to bug in anemoi-inference for multi-step temporal downscalers, can be removed when fixed +verbosity: 0 +allow_nans: true +output_frequency: "1h" From a81103c4600348643f08f5e127b3282b4d196722 Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Tue, 4 Aug 2026 09:28:16 +0200 Subject: [PATCH 03/26] update config and small edits --- .gitignore | 2 +- config/varda-single-1.0-nudge.yaml | 76 ++++++++++--------- config/varda-single-1.0.yaml | 51 +++++++------ ...tidataset-forecaster-global-ich1-oper.yaml | 22 +++++- ...ownscaler-global_trimedge_multi_nudge.yaml | 16 ++-- src/data_input/jretrieve.py | 3 + src/evalml/cli.py | 1 + src/evalml/config.py | 8 ++ uv.lock | 13 ++-- workflow/rules/plot.smk | 6 +- workflow/rules/verification.smk | 12 +++ 11 files changed, 137 insertions(+), 73 deletions(-) diff --git a/.gitignore b/.gitignore index 50b2ff7d..d0995e3c 100644 --- a/.gitignore +++ b/.gitignore @@ -14,7 +14,7 @@ _dev .vscode .idea .snakemake -output +output*/ _sandbox/ experiment_report.html diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml index e7014adb..5761cb7d 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-single-1.0-nudge.yaml @@ -2,44 +2,49 @@ description: | Evaluate skill of Varda-rapid against ground observations. -config_label: varda-single-1.0 +config_label: varda-single-1.0-nudge dates: - start: 2025-03-01T00:00 - end: 2025-03-03T00:00 - frequency: 24h - + # start: 2025-01-01T06:00 + # end: 2025-12-31T06:00 + # frequency: 24h + - 2025-01-02T00:00 + - 2025-01-02T06:00 + - 2025-01-02T12:00 + - 2025-01-02T18:00 + runs: - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 label: Varda-rapid-0.1 - steps: 0/24/1 + steps: 0/12/1 config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml extra_requirements: - anemoi-datasets==0.5.35 - - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@0e1b8ac1a9eb3d4459f71f5a7418fc09c6305e1a + # - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@4a53741c3d5ba50271b570288253f65187b316bd - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml - steps: 0/24/6 + steps: 0/12/6 extra_requirements: - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - - temporal_downscaler: - checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 - label: Varda-single-1.0 - steps: 0/24/1 - config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml - extra_requirements: - - anemoi-datasets==0.5.35 - - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - forecaster: - checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 - config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml - steps: 0/24/6 - extra_requirements: - - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + # - temporal_downscaler: + # checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 + # label: Varda-single-1.0 + # steps: 0/12/1 + # config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml + # extra_requirements: + # - anemoi-datasets==0.5.35 + # - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + # forecaster: + # checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 + # config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml + # steps: 0/12/6 + # extra_requirements: + # - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - baseline: label: INCA @@ -48,15 +53,16 @@ runs: - baseline: label: ICON-CH2-CTRL root: /store_new/mch/msopr/osm/ICON-CH2-EPS - steps: 0/24/1 + steps: 0/12/1 - baseline: label: ICON-CH1-CTRL root: /store_new/mch/msopr/osm/ICON-CH1-EPS - steps: 0/24/1 + steps: 0/12/1 truth: label: SwissMetNet root: jretrievedwh:1,2 + # root: jretrievedwh:bbox=45.8,47.8,5.9,10.5;seq_type=synop # To verify against SwissMetNet observations from the DWH via jretrievedwh, # set instead (requires jretrievedwh.py on $PATH and $OPR_HOME set): # Other selectors: root: jretrievedwh:locations=ARO,KLO,LUG @@ -105,24 +111,26 @@ experiment: - "V_10M:RMSE,R2,ETS" - "T_2M:RMSE,R2,ETS" - "TOT_PREC:RMSE,R2,ETS" + - "TD_2M:RMSE,R2,ETS" short_range: baseline: ICON-CH1-CTRL - lead_times: "6/24/6" - stratification: region - variables: - - "U_10M:RMSE,R2,ETS" - - "V_10M:RMSE,R2,ETS" - - "T_2M:RMSE,R2,ETS" - - "TOT_PREC:RMSE,R2,ETS" - medium_range: - baseline: ICON-CH2-CTRL - lead_times: "24/24/24" + lead_times: "6/12/6" stratification: region variables: - "U_10M:RMSE,R2,ETS" - "V_10M:RMSE,R2,ETS" - "T_2M:RMSE,R2,ETS" - "TOT_PREC:RMSE,R2,ETS" + - "TD_2M:RMSE,R2,ETS" + # medium_range: + # baseline: ICON-CH2-CTRL + # lead_times: "24/24/24" + # stratification: region + # variables: + # - "U_10M:RMSE,R2,ETS" + # - "V_10M:RMSE,R2,ETS" + # - "T_2M:RMSE,R2,ETS" + # - "TOT_PREC:RMSE,R2,ETS" scoremaps: enabled: false diff --git a/config/varda-single-1.0.yaml b/config/varda-single-1.0.yaml index f39000fe..4b710652 100644 --- a/config/varda-single-1.0.yaml +++ b/config/varda-single-1.0.yaml @@ -5,35 +5,42 @@ description: | config_label: varda-single-1.0 dates: - start: 2025-03-01T00:00 - end: 2025-03-03T00:00 - frequency: 24h + # start: 2025-03-01T00:00 + # end: 2025-03-03T00:00 + # frequency: 24h + - 2025-03-01T00:00 runs: - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 label: Varda-single-1.0 - steps: 0/120/1 + steps: 0/24/1 config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml extra_requirements: - anemoi-datasets==0.5.35 - # - anemoi-inference==0.11.0 + # - git+https://github.com/ecmwf/anemoi-inference.git@eaeb36fae30c754a03b8dc0442fc56373b6bedef + # - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss + - -e /scratch/mch/llanzila/sruc/anemoi-inference forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml - steps: 0/120/6 - - baseline: - label: INCA - root: /store_new/mch/msclim/INCA - steps: 0/6/1 - - baseline: - label: ICON-CH2-CTRL - root: /store_new/mch/msopr/osm/ICON-CH2-EPS - steps: 0/120/1 + steps: 0/24/6 + extra_requirements: + - anemoi-datasets==0.5.35 + # - git+https://github.com/ecmwf/anemoi-inference.git@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - -e /scratch/mch/llanzila/sruc/anemoi-inference + # - baseline: + # label: INCA + # root: /store_new/mch/msclim/INCA + # steps: 0/6/1 + # - baseline: + # label: ICON-CH2-CTRL + # root: /store_new/mch/msopr/osm/ICON-CH2-EPS + # steps: 0/120/1 - baseline: label: ICON-CH1-CTRL root: /store_new/mch/msopr/osm/ICON-CH1-EPS - steps: 0/33/1 + steps: 0/24/1 truth: label: SwissMetNet @@ -75,7 +82,7 @@ experiment: # - init_hour - season scorecards: - enabled: true + enabled: false sections: nowcasting: baseline: INCA @@ -88,7 +95,7 @@ experiment: - "TOT_PREC:RMSE,R2,ETS" short_range: baseline: ICON-CH1-CTRL - lead_times: "6/33/6" + lead_times: "6/24/6" stratification: region variables: - "U_10M:RMSE,R2,ETS" @@ -110,17 +117,18 @@ experiment: showcase: params: - T_2M - - SP_10M - - TOT_PREC + - U_10M + - V_10M meteograms: enabled: false stations: [JUN] #, COV, GOR, WFJ, SAE, SAM, DAV, ZER, ANT, VSBAS, BRT, LTB, GOS, CEV, BIA] animations: + fps: 0.5 enabled: true domains: # - globe - - europe - # - alps + # - europe + - alps - icon-ch - switzerland @@ -137,5 +145,4 @@ profile: mem_mb_per_cpu: 1800 runtime: "1h" gpus: 0 - slurm_account: s83 jobs: 50 diff --git a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml index 51d2fb38..9a952d8f 100644 --- a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml +++ b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml @@ -9,7 +9,27 @@ env: input: test: use_original_paths: true - + +# --- Nudging (disabled by default) --- +# Uncomment to nudge the initial condition toward station observations. +# backend: 'peakweather' reads from a local dataset; 'jretrieve' fetches live from DWH. +# nudge_variables: list of GRIB shortNames to nudge; omit to nudge all (T_2M, U_10M, V_10M, TOT_PREC). +# use_limitation: int (optional, jretrieve only) — passed as --use-limitation; limits obs to +# within ±N minutes of the target time (e.g. 50). +# +# pre_processors: +# - forward_transform_filter: +# nudge_toward_observation: +# backend: jretrieve +# nudge_variables: [T_2M] +# icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files +# jretrieve_bbox: [45.7, 48.0, 5.8, 10.8] +# jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src +# k: 3 +# power: 4.0 +# max_dist: 0.5 +# use_limitation: 50 + post_processors: - accumulate_from_start_of_forecast: # accumulate tp from start of forecast accumulations: diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml index 7ae39b72..7008feba 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml @@ -12,19 +12,21 @@ input: scale: 0.001 # convert from kg m-2 to m offset: 0 param: tp + # --- Nudging (disabled by default) --- + # Uncomment to nudge the initial condition toward station observations. + # backend: 'peakweather' reads from a local dataset; 'jretrieve' fetches live from DWH. + # nudge_variables: list of GRIB shortNames to nudge; omit to nudge all (T_2M, U_10M, V_10M, TOT_PREC). - forward_transform_filter: nudge_toward_observation: - backend: jretrieve + nudge_variables: [T_2M, TD_2M, U_10M, V_10M] + icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files jretrieve_bbox: [45.7, 48.0, 5.8, 10.8] - jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src + jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src/data_input/ k: 3 power: 4.0 max_dist: 0.5 - # nudge_toward_observation: - # path_to_observation: /scratch/mch/llanzila/sruc/evalml/output/data/observation/PeakWeather - # k: 3 - # power: 4.0 - # max_dist: 0.5 + use_limitation: 40 + run_mode: "devt" namer: &namer rules: - - shortName: T diff --git a/src/data_input/jretrieve.py b/src/data_input/jretrieve.py index cbf00ad9..131508d3 100644 --- a/src/data_input/jretrieve.py +++ b/src/data_input/jretrieve.py @@ -246,6 +246,7 @@ def fetch_data( increment_minutes=60, seq_type="surface", stage="prod", + use_limitation: int | None = None, timeout_s=600, ) -> pd.DataFrame: """Fetch observation data; columns: station (int), termin (YYYYMMDDhhmmss), @@ -264,6 +265,8 @@ def fetch_data( "csv", *_stations_to_argv(stations), ] + if use_limitation is not None: + argv += ["--use-limitation", str(use_limitation)] LOG.info("jretrieve data: %s", " ".join(argv)) return _parse_csv(_run_with_retry(argv, env=_build_env(stage), timeout_s=timeout_s)) diff --git a/src/evalml/cli.py b/src/evalml/cli.py index f3df3bf4..d8d016e2 100644 --- a/src/evalml/cli.py +++ b/src/evalml/cli.py @@ -26,6 +26,7 @@ def _base_snakemake_command( command += config.profile.parsable() command += ["--configfile", str(configfile)] command += ["--cores", str(cores)] + command += ["--rerun-incomplete"] return command diff --git a/src/evalml/config.py b/src/evalml/config.py index cfc211ad..2d9e83e8 100644 --- a/src/evalml/config.py +++ b/src/evalml/config.py @@ -305,6 +305,14 @@ class AnimationsConfig(BaseModel): "[lon_min, lon_max, lat_min, lat_max], and optional 'projection'." ), ) + fps: float | None = Field( + default=None, + description=( + "Frames per second for the output GIF. Overrides the default speed " + "(which is derived from the model time step). Use values < 1 for slow animations, " + "e.g. 0.5 = one frame every 2 seconds." + ), + ) class ScorecardConfig(BaseModel): diff --git a/uv.lock b/uv.lock index 1bb467be..57321985 100644 --- a/uv.lock +++ b/uv.lock @@ -10,17 +10,16 @@ resolution-markers = [ [[package]] name = "adjusttext" -version = "1.3.0" +version = "1.4.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "matplotlib" }, { name = "numpy" }, { name = "scipy" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/4c/d4/6585f3b6fdb75648bca294664af4becc8aa2fb3fb08f4e4e9fd27e10d773/adjusttext-1.3.0.tar.gz", hash = "sha256:4ab75cd4453af4828876ac3e964f2c49be642ea834f0c1f7449558d5f12cbca1", size = 15724, upload-time = "2024-10-31T16:45:36.101Z" } +sdist = { url = "https://files.pythonhosted.org/packages/b5/5c/496e506ad3313664df24a79f801719cadcd62af5999fcb299c9b08ff0d4b/adjusttext-1.4.0.tar.gz", hash = "sha256:1f73860ced8cccce3f85ee6989ca133c2579b67a7453f63dbeb38f39bf123154", size = 15852, upload-time = "2026-06-08T16:48:32.726Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/53/1c/8feedd607cc14c5df9aef74fe3af9a99bf660743b842a9b5b1865326b4aa/adjustText-1.3.0-py3-none-any.whl", hash = "sha256:da23d7b24b6db5ffa039bb136bfa556207365e32f48ac74b07ad26dd485bc691", size = 13154, upload-time = "2024-10-31T16:45:35.227Z" }, - { url = "https://files.pythonhosted.org/packages/2d/80/7ad35ee5321a86b842f9e8516c8ae4c86f58db7b40e82ce9759f94517a50/adjusttext-1.3.0-py3-none-any.whl", hash = "sha256:bc6c118cd9d7caf6ae37f9355e51d840a2d7f64b4fb2956b8401de27c5af803b", size = 13264, upload-time = "2026-06-08T16:40:05.041Z" }, + { url = "https://files.pythonhosted.org/packages/b8/2c/897bdd17b05724c894a5b831c6b2e9853adcc2d07a70d6246c0cd5cd3912/adjusttext-1.4.0-py3-none-any.whl", hash = "sha256:6febd6484c0d45c39a22f44b2c1f4a8cd01ef58fada565cab4b629c771df79b5", size = 13262, upload-time = "2026-06-08T16:48:31.765Z" }, ] [[package]] @@ -3226,15 +3225,15 @@ wheels = [ [[package]] name = "plotly" -version = "6.7.0" +version = "6.8.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "narwhals" }, { name = "packaging" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/3a/7f/0f100df1172aadf88a929a9dbb902656b0880ba4b960fe5224867159d8f4/plotly-6.7.0.tar.gz", hash = "sha256:45eea0ff27e2a23ccd62776f77eb43aa1ca03df4192b76036e380bb479b892c6", size = 6911286, upload-time = "2026-04-09T20:36:45.738Z" } +sdist = { url = "https://files.pythonhosted.org/packages/94/fd/d72c292d78aadb93d1a9bcd76bf3c678271040c7cf10abe5788b33040a39/plotly-6.8.0.tar.gz", hash = "sha256:e088e7ddc68d4f70e3d66659224727a45296d71d2b8284181862d3d8f1f0d88f", size = 6915161, upload-time = "2026-06-03T18:33:40.226Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/90/ad/cba91b3bcf04073e4d1655a5c1710ef3f457f56f7d1b79dcc3d72f4dd912/plotly-6.7.0-py3-none-any.whl", hash = "sha256:ac8aca1c25c663a59b5b9140a549264a5badde2e057d79b8c772ae2920e32ff0", size = 9898444, upload-time = "2026-04-09T20:36:39.812Z" }, + { url = "https://files.pythonhosted.org/packages/f9/14/abe5ce876ab5b66ee3c691bf537fcd43d037aea55d447aacf74630a8f31e/plotly-6.8.0-py3-none-any.whl", hash = "sha256:13c5c4a0f70b74cab1913eda0de49b826df5931708eb6f9c3010040614700ec8", size = 9902055, upload-time = "2026-06-03T18:33:34.26Z" }, ] [[package]] diff --git a/workflow/rules/plot.smk b/workflow/rules/plot.smk index e51cbf51..eb70acdf 100644 --- a/workflow/rules/plot.smk +++ b/workflow/rules/plot.smk @@ -155,7 +155,11 @@ rule make_forecast_animation: region="|".join(map(re.escape, SHOWCASE_REGIONS.keys())), localrule: True params: - delay=lambda wc: 10 * int(RUN_CONFIGS[wc.run_id]["steps"].split("/")[2]), + delay=lambda wc: ( + int(100 / config.get("showcase", {}).get("animations", {}).get("fps")) + if config.get("showcase", {}).get("animations", {}).get("fps") + else 10 * int(RUN_CONFIGS[wc.run_id]["steps"].split("/")[2]) + ), shell: """ FRAMES=$(for f in {input}; do [ -s "$f" ] && echo "$f"; done | tr '\\n' ' ') diff --git a/workflow/rules/verification.smk b/workflow/rules/verification.smk index 2b8c80b5..540534c2 100644 --- a/workflow/rules/verification.smk +++ b/workflow/rules/verification.smk @@ -121,6 +121,18 @@ rule verification_metrics_aggregation: init_time=_restrict_reftimes_to_hours(REFTIMES), allow_missing=True, ), + inference_okfiles=lambda wc: expand( + rules.inference_execute.output.okfile, + run_id=( + [wc.run_id] + + ( + [RUN_CONFIGS[wc.run_id]["forecaster"]["run_id"]] + if RUN_CONFIGS[wc.run_id].get("forecaster") is not None + else [] + ) + ), + init_time=_restrict_reftimes_to_hours(REFTIMES), + ), output: OUT_ROOT / f"data/runs/{{run_id}}/verif_aggregated_{TRUTH_HASH}.nc", log: From 6968349f5a6a69112da16ad16d35385dadf8c7b8 Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Tue, 4 Aug 2026 17:33:58 +0200 Subject: [PATCH 04/26] update configs for nudging --- config/varda-single-1.0-nudge.yaml | 2 +- ...ownscaler-global_trimedge_multi_nudge.yaml | 28 ++++++++++++++++--- 2 files changed, 25 insertions(+), 5 deletions(-) diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml index 5761cb7d..9fac8898 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-single-1.0-nudge.yaml @@ -22,7 +22,7 @@ runs: extra_requirements: - anemoi-datasets==0.5.35 # - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss - - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@4a53741c3d5ba50271b570288253f65187b316bd + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@d3786cf393a41b176af94ba3b3fb0d3da8c41b71 - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml index 7008feba..6c087979 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml @@ -14,19 +14,39 @@ input: param: tp # --- Nudging (disabled by default) --- # Uncomment to nudge the initial condition toward station observations. - # backend: 'peakweather' reads from a local dataset; 'jretrieve' fetches live from DWH. + # 1. retrieve_observation: fetches live data from DWH via jretrieve → raw Parquet + # 2. clean_observation: reads raw Parquet, applies QC cleaning → cleaned Parquet + # 3. nudge_toward_observation: reads cleaned Parquet, applies IoR correction # nudge_variables: list of GRIB shortNames to nudge; omit to nudge all (T_2M, U_10M, V_10M, TOT_PREC). + - forward_transform_filter: + retrieve_observation: + obs_path: observation/nudging_station_obs_raw.parquet + jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src/data_input/ + bbox: [45.7, 48.0, 5.8, 10.8] + variables: [T_2M, TD_2M, U_10M, V_10M] + use_limitation: 20 + run_mode: "devt" + - forward_transform_filter: + clean_observation: + obs_path_in: observation/nudging_station_obs_raw.parquet + obs_path_out: observation/nudging_station_obs.parquet - forward_transform_filter: nudge_toward_observation: + obs_path: observation/nudging_station_obs.parquet nudge_variables: [T_2M, TD_2M, U_10M, V_10M] icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files - jretrieve_bbox: [45.7, 48.0, 5.8, 10.8] - jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src/data_input/ k: 3 power: 4.0 max_dist: 0.5 - use_limitation: 40 run_mode: "devt" + holdout_fraction: 0.05 + holdout_seed: 42 + # Station holdout — use one of the two options below, not both. + # Option A: randomly withhold a fraction of stations (e.g. 5% for cross-validation). + # holdout_fraction: 0.05 + # holdout_seed: 42 # optional, defaults to 42 for reproducibility + # Option B: exclude specific stations by their nat_abbr identifier. + # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] namer: &namer rules: - - shortName: T From 930ef8971c3534106e13905faaee038ec65f1c2a Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Tue, 4 Aug 2026 18:38:43 +0200 Subject: [PATCH 05/26] chore: fix pre-commit --- config/varda-single-1.0-nudge.yaml | 2 +- ...oral-downscaler-global_trimedge_multi_nudge.yaml | 4 ++-- workflow/tools/config.schema.json | 13 +++++++++++++ 3 files changed, 16 insertions(+), 3 deletions(-) diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml index 9fac8898..d9ee4b34 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-single-1.0-nudge.yaml @@ -12,7 +12,7 @@ dates: - 2025-01-02T06:00 - 2025-01-02T12:00 - 2025-01-02T18:00 - + runs: - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml index 6c087979..2fec2f0a 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml @@ -40,13 +40,13 @@ input: max_dist: 0.5 run_mode: "devt" holdout_fraction: 0.05 - holdout_seed: 42 + holdout_seed: 42 # Station holdout — use one of the two options below, not both. # Option A: randomly withhold a fraction of stations (e.g. 5% for cross-validation). # holdout_fraction: 0.05 # holdout_seed: 42 # optional, defaults to 42 for reproducibility # Option B: exclude specific stations by their nat_abbr identifier. - # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] + # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] namer: &namer rules: - - shortName: T diff --git a/workflow/tools/config.schema.json b/workflow/tools/config.schema.json index b974e30e..23bed3d6 100644 --- a/workflow/tools/config.schema.json +++ b/workflow/tools/config.schema.json @@ -28,6 +28,19 @@ }, "title": "Domains", "type": "array" + }, + "fps": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "default": null, + "description": "Frames per second for the output GIF. Overrides the default speed (which is derived from the model time step). Use values < 1 for slow animations, e.g. 0.5 = one frame every 2 seconds.", + "title": "Fps" } }, "title": "AnimationsConfig", From 9e43c404fcc86903aa4f0340579f20fd07ff4771 Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Wed, 5 Aug 2026 17:28:35 +0200 Subject: [PATCH 06/26] add cross validation for nudging --- config/varda-single-1.0-nudge.yaml | 6 +- src/evalml/config.py | 9 ++- src/verification/__init__.py | 26 ++++++++ workflow/rules/common.smk | 1 + workflow/rules/verification.smk | 8 +++ .../scripts/report_experiment_dashboard.py | 10 ++++ workflow/scripts/report_scorecard.py | 7 +++ workflow/scripts/verification_metrics.py | 60 +++++++++++++++++++ workflow/scripts/verification_plot_metrics.py | 2 + 9 files changed, 126 insertions(+), 3 deletions(-) diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml index d9ee4b34..512ed67f 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-single-1.0-nudge.yaml @@ -21,8 +21,8 @@ runs: config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml extra_requirements: - anemoi-datasets==0.5.35 - # - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss - - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@d3786cf393a41b176af94ba3b3fb0d3da8c41b71 + - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss + # - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f91315e - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 @@ -99,6 +99,8 @@ experiment: # - region # - init_hour - season + - station_group + cross_validation: true scorecards: enabled: true sections: diff --git a/src/evalml/config.py b/src/evalml/config.py index 2d9e83e8..d07b142b 100644 --- a/src/evalml/config.py +++ b/src/evalml/config.py @@ -399,7 +399,7 @@ class Dashboard(BaseModel): stratification: List[str] = Field( ..., - description="Stratifications to include in the dashboard (any of season, region, init_hour)", + description="Stratifications to include in the dashboard (any of season, region, init_hour, station_group)", ) @@ -425,6 +425,13 @@ class ExperimentConfig(BaseModel): ..., description="Settings for the experiment dashboard.", ) + cross_validation: bool = Field( + default=False, + description=( + "When True, adds a 'station_group' stratification dimension (all/holdout/holdin) " + "to separate metrics for stations used in nudging from those withheld." + ), + ) scorecards: Optional[ExperimentScorecardConfig] = Field( default=None, description="Scorecard generation configuration. Omit or set enabled: false to disable.", diff --git a/src/verification/__init__.py b/src/verification/__init__.py index a4ee8499..92e22048 100644 --- a/src/verification/__init__.py +++ b/src/verification/__init__.py @@ -226,6 +226,19 @@ def _merge_metrics(ds: xr.Dataset, num_workers: int = 4) -> xr.Dataset: return out +def _create_station_group_masks( + values_coord: xr.DataArray, holdout_stations: list[str] +) -> xr.DataArray: + """Boolean masks for all/holdout/holdin station groups over the values dimension.""" + all_nat_abbr = values_coord.values + is_holdout = np.isin(all_nat_abbr, holdout_stations) + return xr.DataArray( + np.stack([np.ones_like(is_holdout), is_holdout, ~is_holdout]), + coords={"station_group": ["all", "holdout", "holdin"], "values": values_coord.values}, + dims=["station_group", "values"], + ) + + def verify( fcst: xr.Dataset, obs: xr.Dataset, @@ -235,6 +248,7 @@ def verify( dim: list[str] | None = None, threshold_dict: dict[str, dict[str, list[float]]] | None = None, num_workers: int | None = None, + holdout_stations: list[str] | None = None, ) -> xr.Dataset: """ Compute verification metrics and statistics comparing forecast and observation datasets. @@ -292,6 +306,13 @@ def verify( lon=obs_aligned["longitude"], lat=obs_aligned["latitude"] ) + station_masks = None + if holdout_stations is not None: + station_masks = _create_station_group_masks(obs_aligned["values"], holdout_stations) + LOG.info("Station group masks created: %d holdout, %d holdin stations", + int(station_masks.sel(station_group="holdout").sum()), + int(station_masks.sel(station_group="holdin").sum())) + scores = [] statistics = [] for param in fcst_aligned.data_vars: @@ -314,6 +335,11 @@ def verify( fcst_param = fcst_aligned[param].where(masks) obs_param = obs_aligned[param].where(masks) + # Apply station group masks — adds "station_group" dim (all/holdout/holdin) + if station_masks is not None: + fcst_param = fcst_param.where(station_masks) + obs_param = obs_param.where(station_masks) + score = _compute_scores( fcst_param, obs_param, diff --git a/workflow/rules/common.smk b/workflow/rules/common.smk index 4a4c29b8..3b0ae057 100644 --- a/workflow/rules/common.smk +++ b/workflow/rules/common.smk @@ -375,6 +375,7 @@ if "jretrieve" in str(config["truth"]["root"]): TRUTH_HASH = truth_hash(config["truth"]) +CROSS_VALIDATION = config.get("experiment", {}).get("cross_validation", False) REGIONS = parse_regions() SHOWCASE_REGIONS = parse_showcase_regions() SHOWCASE_PARAMS = config.get("showcase", {}).get("params", ["T_2M", "SP_10M"]) diff --git a/workflow/rules/verification.smk b/workflow/rules/verification.smk index 540534c2..4df03367 100644 --- a/workflow/rules/verification.smk +++ b/workflow/rules/verification.smk @@ -34,6 +34,7 @@ rule verification_metrics_baseline: regions=REGIONS, experiment_params=",".join(EXPERIMENT_PARAMS), threshold_dict=config["experiment"]["thresholds"], + cross_validation=False, shell: """ export ECCODES_DEFINITION_PATH=$(realpath .venv/share/eccodes-cosmo-resources/definitions) @@ -47,6 +48,7 @@ rule verification_metrics_baseline: --regions "{params.regions}" \ --params "{params.experiment_params}" \ --threshold_dict "{params.threshold_dict}" \ + --cross_validation {params.cross_validation} \ --member "{params.member}" \ --output {output} >{log} 2>&1 """ @@ -89,6 +91,10 @@ rule verification_metrics: ).resolve(), experiment_params=",".join(EXPERIMENT_PARAMS), threshold_dict=config["experiment"]["thresholds"], + cross_validation=CROSS_VALIDATION, + run_workdir=lambda wc: ( + Path(OUT_ROOT) / f"data/runs/{wc.run_id}/{wc.init_time}" + ).resolve(), shell: """ export ECCODES_DEFINITION_PATH=$(realpath .venv/share/eccodes-cosmo-resources/definitions) @@ -102,6 +108,8 @@ rule verification_metrics: --regions "{params.regions}" \ --params "{params.experiment_params}" \ --threshold_dict "{params.threshold_dict}" \ + --cross_validation {params.cross_validation} \ + --run_workdir {params.run_workdir} \ --output {output} >{log} 2>&1 """ diff --git a/workflow/scripts/report_experiment_dashboard.py b/workflow/scripts/report_experiment_dashboard.py index 3edd127e..c18d47e2 100644 --- a/workflow/scripts/report_experiment_dashboard.py +++ b/workflow/scripts/report_experiment_dashboard.py @@ -86,6 +86,13 @@ def main(args): df = df[df["season"] == "all"] if "init_hour" not in stratification: df = df[df["init_hour"] == "all"] + # station_group: baselines have no station_group dim → fill with "all" + if "station_group" not in df.columns: + df["station_group"] = "all" + else: + df["station_group"] = df["station_group"].fillna("all") + if "station_group" not in stratification: + df = df[df["station_group"] == "all"] # create a new column for line styles and shapes in dashboard df.dropna(inplace=True) @@ -98,6 +105,7 @@ def main(args): regions = df["region"].unique() if "region" in stratification else [] seasons = df["season"].unique() if "season" in stratification else [] init_hours = df["init_hour"].unique() if "init_hour" in stratification else [] + station_groups = df["station_group"].unique() if "station_group" in stratification else [] # Columnar JSON: store columns + data array (no repeated keys per row). # region_season_init is a derived column — computed in JS at parse time. @@ -119,6 +127,7 @@ def _round_sig(x, sig=6): "region", "season", "init_hour", + "station_group", ] df_export = df[export_cols].copy() df_export["value"] = df_export["value"].apply( @@ -160,6 +169,7 @@ def _sanitize(v): regions=regions, seasons=seasons, init_hours=init_hours, + station_groups=station_groups, stratification=stratification, header_text=args.header_text, configfile_content=open(args.configfile, "r").read() diff --git a/workflow/scripts/report_scorecard.py b/workflow/scripts/report_scorecard.py index 0e05195b..71f6e67b 100644 --- a/workflow/scripts/report_scorecard.py +++ b/workflow/scripts/report_scorecard.py @@ -218,6 +218,13 @@ def _load_relative_diff(cfg: dict) -> xr.Dataset: model_ds = xr.open_dataset(cfg["model"]["path"]) baseline_ds = xr.open_dataset(cfg["baseline"]["path"]) + # When cross_validation is enabled, model runs carry a station_group dim. + # Scorecards always compare the "all" group — select it and drop the dim. + if "station_group" in model_ds.dims: + model_ds = model_ds.sel(station_group="all", drop=True) + if "station_group" in baseline_ds.dims: + baseline_ds = baseline_ds.sel(station_group="all", drop=True) + for label, ds in [("model", model_ds), ("baseline", baseline_ds)]: if "n_samples" not in ds.data_vars: raise ValueError( diff --git a/workflow/scripts/verification_metrics.py b/workflow/scripts/verification_metrics.py index d5882f00..df0f6cf0 100644 --- a/workflow/scripts/verification_metrics.py +++ b/workflow/scripts/verification_metrics.py @@ -4,6 +4,9 @@ from datetime import datetime from pathlib import Path +import numpy as np +import pandas as pd +import yaml from verification import verify # noqa: E402 from verification.spatial import map_forecast_to_truth # noqa: E402 @@ -30,6 +33,39 @@ class ScriptConfig(Namespace): steps: list[int] = parse_steps("0/120/6") +def _find_nudge_cfg(inf_cfg: dict) -> dict: + """Extract nudge_toward_observation params from the inference config YAML.""" + try: + pre_processors = inf_cfg["input"]["cutout"][0]["lam_0"]["grib"]["pre_processors"] + for pp in pre_processors: + inner = pp.get("forward_transform_filter", {}) + if "nudge_toward_observation" in inner: + return inner["nudge_toward_observation"] + except (KeyError, IndexError, TypeError): + pass + return {} + + +def compute_holdout_stations(obs_parquet: Path, inference_config: Path) -> list[str]: + """Return nat_abbr list of stations withheld from nudging, mirroring the plugin logic.""" + df = pd.read_parquet(obs_parquet) + with open(inference_config) as f: + inf_cfg = yaml.safe_load(f) + nudge_cfg = _find_nudge_cfg(inf_cfg) + + exclude_stations = nudge_cfg.get("exclude_stations") + holdout_fraction = nudge_cfg.get("holdout_fraction") + holdout_seed = nudge_cfg.get("holdout_seed", 42) + + if exclude_stations is not None: + return [s for s in exclude_stations if s in df.index] + if holdout_fraction is not None and 0.0 < float(holdout_fraction) < 1.0: + n_holdout = round(len(df) * float(holdout_fraction)) + rng = np.random.default_rng(holdout_seed) + return list(rng.choice(df.index, size=n_holdout, replace=False)) + return [] + + def program_summary_log(args): """Log a welcome message with the script information.""" LOG.info("=" * 80) @@ -84,6 +120,17 @@ def main(args: ScriptConfig): (datetime.now() - now).total_seconds(), ) + # determine holdout stations for cross-validation station stratification + holdout_stations = None + if args.cross_validation and args.run_workdir is not None: + obs_parquet = args.run_workdir / "observation/nudging_station_obs.parquet" + inference_config = args.run_workdir / "config.yaml" + if obs_parquet.exists() and inference_config.exists(): + holdout_stations = compute_holdout_stations(obs_parquet, inference_config) + LOG.info("Cross-validation holdout: %d stations withheld: %s", len(holdout_stations), holdout_stations) + else: + LOG.warning("cross_validation=True but obs parquet or inference config not found; skipping station stratification.") + # compute metrics and statistics now = datetime.now() results = verify( @@ -93,6 +140,7 @@ def main(args: ScriptConfig): args.truth_label, args.regions, threshold_dict=args.threshold_dict, + holdout_stations=holdout_stations, ) LOG.info( "Computed verification metrics in %s seconds", @@ -169,6 +217,18 @@ def main(args: ScriptConfig): help="Dictionary of thresholds for each parameter in the format '{param: [threshold1, threshold2, ...]}' (default: None).", default=None, ) + parser.add_argument( + "--cross_validation", + type=lambda x: x.lower() == "true", + default=False, + help="When True, compute metrics separately for holdout/holdin station groups.", + ) + parser.add_argument( + "--run_workdir", + type=Path, + default=None, + help="Per-run working directory containing observation parquet and config.yaml.", + ) parser.add_argument( "--member", type=str, diff --git a/workflow/scripts/verification_plot_metrics.py b/workflow/scripts/verification_plot_metrics.py index 2f4b685a..8a199112 100644 --- a/workflow/scripts/verification_plot_metrics.py +++ b/workflow/scripts/verification_plot_metrics.py @@ -77,6 +77,8 @@ def main(args: Namespace) -> None: # remove duplicated but not identical values from analyses (rounding errors) dfs = [xr.open_dataset(f) for f in args.verif_files] + # When cross_validation is enabled, runs carry a station_group dim; select "all" for standard plots. + dfs = [d.sel(station_group="all", drop=True) if "station_group" in d.dims else d for d in dfs] # 1) Ensure each dataset has unique lead_time values dfs = [_ensure_unique_lead_time(d) for d in dfs] # 2) For sources present in multiple datasets, keep the one with most lead_times From 9f48857a7fabb4dee8f1103f711b687357578e96 Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Wed, 5 Aug 2026 17:52:55 +0200 Subject: [PATCH 07/26] fix bug in dashboard plotting --- .../scripts/report_experiment_dashboard.py | 36 ++++++++++++++++--- 1 file changed, 32 insertions(+), 4 deletions(-) diff --git a/workflow/scripts/report_experiment_dashboard.py b/workflow/scripts/report_experiment_dashboard.py index c18d47e2..3fb8d8a5 100644 --- a/workflow/scripts/report_experiment_dashboard.py +++ b/workflow/scripts/report_experiment_dashboard.py @@ -64,6 +64,13 @@ def main(args): _check_n_samples_consistency(dfs, args.verif_files) dfs = [_ensure_unique_lead_time(d) for d in dfs] dfs = _select_best_sources(dfs) + # Normalize station_group dimension before concat: datasets without it (baselines) are + # expanded to station_group=["all"] so xr.concat receives uniform-rank tensors. + if any("station_group" in d.dims for d in dfs): + dfs = [ + d if "station_group" in d.dims else d.expand_dims(station_group=["all"]) + for d in dfs + ] ds = xr.concat(dfs, dim="source", join="outer") LOG.info("Loaded verification netcdf: \n%s", ds) @@ -86,16 +93,37 @@ def main(args): df = df[df["season"] == "all"] if "init_hour" not in stratification: df = df[df["init_hour"] == "all"] - # station_group: baselines have no station_group dim → fill with "all" if "station_group" not in df.columns: df["station_group"] = "all" - else: - df["station_group"] = df["station_group"].fillna("all") if "station_group" not in stratification: df = df[df["station_group"] == "all"] - # create a new column for line styles and shapes in dashboard + # Drop NaN rows before station_group fold so that baselines (NaN at holdin/holdout) + # are not mistakenly counted as multi-group sources. df.dropna(inplace=True) + + # When station_group is in stratification, fold it into the source name for forecast + # sources so each group appears as a separate labelled line in the dashboard. + # Truth/obs sources (e.g. SwissMetNet) only carry stat metrics (mean/std/min/max) + # and are excluded from the fold to avoid spurious "SwissMetNet (holdin)" entries. + if "station_group" in stratification: + _stat_metrics = {"mean", "std", "min", "max"} + _forecast_sources = set( + df.groupby("source")["metric"] + .apply(lambda m: not m.isin(_stat_metrics).all()) + .pipe(lambda s: s[s].index) + ) + _multi_sources = set( + df[df["source"].isin(_forecast_sources)] + .groupby("source")["station_group"] + .nunique() + .pipe(lambda s: s[s > 1].index) + ) + if _multi_sources: + mask = df["source"].isin(_multi_sources) + df.loc[mask, "source"] = ( + df.loc[mask, "source"] + " (" + df.loc[mask, "station_group"] + ")" + ) LOG.info("Loaded verification data frame: \n%s", df) # get unique sources and params From 5b88610a54e6aa26d1ec2e71ec135207337ffaeb Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Thu, 6 Aug 2026 22:30:20 +0200 Subject: [PATCH 08/26] small edits to dashboard --- config/varda-single-1.0-nudge.yaml | 58 +++--- ...caler-global_trimedge_multi_nudge_all.yaml | 167 ++++++++++++++++++ ...-global_trimedge_multi_nudge_holdout.yaml} | 1 + src/evalml/config.py | 2 +- .../scripts/report_experiment_dashboard.py | 44 ++--- 5 files changed, 225 insertions(+), 47 deletions(-) create mode 100644 resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml rename resources/inference/configs/{sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml => sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml} (97%) diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml index 512ed67f..f7c27983 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-single-1.0-nudge.yaml @@ -7,18 +7,35 @@ config_label: varda-single-1.0-nudge dates: # start: 2025-01-01T06:00 # end: 2025-12-31T06:00 - # frequency: 24h + # frequency: 54h - 2025-01-02T00:00 - 2025-01-02T06:00 - 2025-01-02T12:00 - 2025-01-02T18:00 runs: + - temporal_downscaler: + checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 + label: Varda-rapid-0.1-all + steps: 0/12/1 + config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml + extra_requirements: + - anemoi-datasets==0.5.35 + - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss + # - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f91315e + - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + forecaster: + checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 + config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml + steps: 0/12/6 + extra_requirements: + - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 label: Varda-rapid-0.1 steps: 0/12/1 - config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml + config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml extra_requirements: - anemoi-datasets==0.5.35 - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss @@ -31,29 +48,29 @@ runs: extra_requirements: - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - # - temporal_downscaler: - # checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 - # label: Varda-single-1.0 - # steps: 0/12/1 - # config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml - # extra_requirements: - # - anemoi-datasets==0.5.35 - # - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - # forecaster: - # checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 - # config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml - # steps: 0/12/6 - # extra_requirements: - # - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - temporal_downscaler: + checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 + label: Varda-single-1.0 + steps: 0/12/1 + config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml + extra_requirements: + - anemoi-datasets==0.5.35 + - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + forecaster: + checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 + config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml + steps: 0/12/6 + extra_requirements: + - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - baseline: label: INCA root: /store_new/mch/msclim/INCA steps: 0/6/1 - - baseline: - label: ICON-CH2-CTRL - root: /store_new/mch/msopr/osm/ICON-CH2-EPS - steps: 0/12/1 + # - baseline: + # label: ICON-CH2-CTRL + # root: /store_new/mch/msopr/osm/ICON-CH2-EPS + # steps: 0/12/1 - baseline: label: ICON-CH1-CTRL root: /store_new/mch/msopr/osm/ICON-CH1-EPS @@ -99,7 +116,6 @@ experiment: # - region # - init_hour - season - - station_group cross_validation: true scorecards: enabled: true diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml new file mode 100644 index 00000000..844b4e14 --- /dev/null +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml @@ -0,0 +1,167 @@ +runner: temporal_downscaler + +input: + cutout: + - lam_0: + grib: + path: forecaster/20* + pre_processors: + - forward_transform_filter: + # does not have an effect if the temporal downscaler does not use tp as prognostic variable + rescale: + scale: 0.001 # convert from kg m-2 to m + offset: 0 + param: tp + # --- Nudging (disabled by default) --- + # Uncomment to nudge the initial condition toward station observations. + # 1. retrieve_observation: fetches live data from DWH via jretrieve → raw Parquet + # 2. clean_observation: reads raw Parquet, applies QC cleaning → cleaned Parquet + # 3. nudge_toward_observation: reads cleaned Parquet, applies IoR correction + # nudge_variables: list of GRIB shortNames to nudge; omit to nudge all (T_2M, U_10M, V_10M, TOT_PREC). + - forward_transform_filter: + retrieve_observation: + obs_path: observation/nudging_station_obs_raw.parquet + jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src/data_input/ + bbox: [45.7, 48.0, 5.8, 10.8] + variables: [T_2M, TD_2M, U_10M, V_10M] + use_limitation: 20 + run_mode: "devt" + - forward_transform_filter: + clean_observation: + obs_path_in: observation/nudging_station_obs_raw.parquet + obs_path_out: observation/nudging_station_obs.parquet + - forward_transform_filter: + nudge_toward_observation: + obs_path: observation/nudging_station_obs.parquet + nudge_variables: [T_2M, TD_2M, U_10M, V_10M] + icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files + k: 3 + power: 4.0 + max_dist: 0.5 + run_mode: "devt" + holdout_fraction: 0.0 + holdout_seed: 42 + # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] + # Station holdout — use one of the two options below, not both. + # Option A: randomly withhold a fraction of stations (e.g. 5% for cross-validation). + # holdout_fraction: 0.05 + # holdout_seed: 42 # optional, defaults to 42 for reproducibility + # Option B: exclude specific stations by their nat_abbr identifier. + # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] + namer: &namer + rules: + - - shortName: T + - t_{level} + - - shortName: U + - u_{level} + - - shortName: V + - v_{level} + - - shortName: W + - w_{level} + - - shortName: QV + - q_{level} + - - shortName: FI + - z_{level} + - - shortName: PMSL + - msl + - - shortName: FIS + - z + - - shortName: PS + - sp + - - shortName: T_2M + - 2t + - - shortName: TD_2M + - 2d + - - shortName: T_G + - skt + - - shortName: U_10M + - 10u + - - shortName: V_10M + - 10v + - - shortName: FR_LAND + - lsm + - - shortName: TOT_PREC + - tp + - global: + grib: + path: forecaster/ifs* + namer: *namer + +constant_forcings: + test: + use_original_paths: true + +patch_metadata: resources/sgm-temporal-downscaler-ich1-oper-patch.yaml + +post_processors: + - accumulate_from_start_of_forecast: # accumulate tp from start of forecast + accumulations: + - tp + - forward_transform_filter: + rescale: + scale: 1000 # convert units from m to kg m-2 + offset: 0 + param: tp + +output: + tee: + - grib: + path: grib/{dateTime}_{step:03}.grib + encoding: + typeOfGeneratingProcess: 2 + templates: + samples: resources/templates_index_icon.yaml + post_processors: + - extract_mask: # removes global points + mask: "lam_0/cutout_mask" + as_slice: true + # here, the trimedge mask can be specified when available + + - grib: + path: grib/ifs-{dateTime}_{step:03}.grib + encoding: + typeOfGeneratingProcess: 2 + templates: + samples: resources/templates_index_ifs.yaml + post_processors: + - extract_mask: # removes lam points + mask: "lam_0/cutout_mask" + as_slice: true + inverse: true + - assign_mask: # fill local/global overlapping points with nan + mask: "global/cutout_mask" + modifiers: + - patches: + - variable: + U_10M: {"param": 165, "shortName": "10u"} + 10u: {"param": 165, "shortName": "10u"} + V_10M: {"param": 166, "shortName": "10v"} + 10v: {"param": 166, "shortName": "10v"} + TD_2M: {"param": 168, "shortName": "2d"} + 2d: {"param": 168, "shortName": "2d"} + T_2M: {"param": 167, "shortName": "2t"} + 2t: {"param": 167, "shortName": "2t"} + FR_LAND: {"param": 172, "shortName": "lsm"} + lsm: {"param": 172, "shortName": "lsm"} + PMSL: {"param": 151, "shortName": "msl"} + msl: {"param": 151, "shortName": "msl"} + PS: {"param": 134, "shortName": "sp"} + sp: {"param": 134, "shortName": "sp"} + SSO_SIGMA: {"param": 163, "shortName": "slor"} + slor: {"param": 163, "shortName": "slor"} + SSO_STDH: {"param": 160, "shortName": "sdor"} + sdor: {"param": 160, "shortName": "sdor"} + TOT_PREC: {"param": 228, "shortName": "tp"} + tp: {"param": 228, "shortName": "tp"} + z: {"param": 129, "shortName": "z"} + "^q_(\\d+)$": {"param": 133, "shortName": "q"} + "^t_(\\d+)$": {"param": 130, "shortName": "t"} + "^u_(\\d+)$": {"param": 131, "shortName": "u"} + "^v_(\\d+)$": {"param": 132, "shortName": "v"} + "^w_(\\d+)$": {"param": 135, "shortName": "w"} + "^z_(\\d+)$": {"param": 129, "shortName": "z"} + +# silenced due to bug in anemoi-inference for multi-step temporal downscalers, can be removed when fixed +verbosity: 0 +allow_nans: true +output_frequency: "1h" diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml similarity index 97% rename from resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml rename to resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml index 2fec2f0a..0446f0be 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml @@ -41,6 +41,7 @@ input: run_mode: "devt" holdout_fraction: 0.05 holdout_seed: 42 + # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] # Station holdout — use one of the two options below, not both. # Option A: randomly withhold a fraction of stations (e.g. 5% for cross-validation). # holdout_fraction: 0.05 diff --git a/src/evalml/config.py b/src/evalml/config.py index d07b142b..04334bb2 100644 --- a/src/evalml/config.py +++ b/src/evalml/config.py @@ -399,7 +399,7 @@ class Dashboard(BaseModel): stratification: List[str] = Field( ..., - description="Stratifications to include in the dashboard (any of season, region, init_hour, station_group)", + description="Stratifications to include in the dashboard (any of season, region, init_hour)", ) diff --git a/workflow/scripts/report_experiment_dashboard.py b/workflow/scripts/report_experiment_dashboard.py index 3fb8d8a5..b44f505c 100644 --- a/workflow/scripts/report_experiment_dashboard.py +++ b/workflow/scripts/report_experiment_dashboard.py @@ -95,35 +95,32 @@ def main(args): df = df[df["init_hour"] == "all"] if "station_group" not in df.columns: df["station_group"] = "all" - if "station_group" not in stratification: - df = df[df["station_group"] == "all"] # Drop NaN rows before station_group fold so that baselines (NaN at holdin/holdout) # are not mistakenly counted as multi-group sources. df.dropna(inplace=True) - # When station_group is in stratification, fold it into the source name for forecast - # sources so each group appears as a separate labelled line in the dashboard. + # Automatically fold station_group into the source name for forecast sources that + # have multiple station groups (only happens when cross_validation=true). # Truth/obs sources (e.g. SwissMetNet) only carry stat metrics (mean/std/min/max) - # and are excluded from the fold to avoid spurious "SwissMetNet (holdin)" entries. - if "station_group" in stratification: - _stat_metrics = {"mean", "std", "min", "max"} - _forecast_sources = set( - df.groupby("source")["metric"] - .apply(lambda m: not m.isin(_stat_metrics).all()) - .pipe(lambda s: s[s].index) - ) - _multi_sources = set( - df[df["source"].isin(_forecast_sources)] - .groupby("source")["station_group"] - .nunique() - .pipe(lambda s: s[s > 1].index) + # and are excluded to avoid spurious "SwissMetNet (holdin)" entries. + _stat_metrics = {"mean", "std", "min", "max"} + _forecast_sources = set( + df.groupby("source")["metric"] + .apply(lambda m: not m.isin(_stat_metrics).all()) + .pipe(lambda s: s[s].index) + ) + _multi_sources = set( + df[df["source"].isin(_forecast_sources)] + .groupby("source")["station_group"] + .nunique() + .pipe(lambda s: s[s > 1].index) + ) + if _multi_sources: + mask = df["source"].isin(_multi_sources) + df.loc[mask, "source"] = ( + df.loc[mask, "source"] + " (" + df.loc[mask, "station_group"] + ")" ) - if _multi_sources: - mask = df["source"].isin(_multi_sources) - df.loc[mask, "source"] = ( - df.loc[mask, "source"] + " (" + df.loc[mask, "station_group"] + ")" - ) LOG.info("Loaded verification data frame: \n%s", df) # get unique sources and params @@ -133,7 +130,6 @@ def main(args): regions = df["region"].unique() if "region" in stratification else [] seasons = df["season"].unique() if "season" in stratification else [] init_hours = df["init_hour"].unique() if "init_hour" in stratification else [] - station_groups = df["station_group"].unique() if "station_group" in stratification else [] # Columnar JSON: store columns + data array (no repeated keys per row). # region_season_init is a derived column — computed in JS at parse time. @@ -155,7 +151,6 @@ def _round_sig(x, sig=6): "region", "season", "init_hour", - "station_group", ] df_export = df[export_cols].copy() df_export["value"] = df_export["value"].apply( @@ -197,7 +192,6 @@ def _sanitize(v): regions=regions, seasons=seasons, init_hours=init_hours, - station_groups=station_groups, stratification=stratification, header_text=args.header_text, configfile_content=open(args.configfile, "r").read() From 8c4bb4570682de7e2f34b00647109d593a7aebba Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Fri, 7 Aug 2026 19:37:19 +0200 Subject: [PATCH 09/26] add cross validation functionality: stations grouped in all, holdin and holdout --- config/varda-single-1.0-nudge.yaml | 38 +++++---- ...caler-global_trimedge_multi_nudge_all.yaml | 7 +- ...r-global_trimedge_multi_nudge_holdout.yaml | 28 ++++--- resources/report/dashboard/script.js | 18 ++++- .../report/dashboard/template.html.jinja2 | 10 +++ src/evalml/config.py | 31 +++++++- workflow/rules/common.smk | 3 +- workflow/rules/report.smk | 2 + workflow/rules/verification.smk | 13 ++-- workflow/scripts/inference_prepare.py | 2 + .../scripts/report_experiment_dashboard.py | 41 ++++------ workflow/scripts/verification_metrics.py | 77 ++++++++----------- 12 files changed, 155 insertions(+), 115 deletions(-) diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml index f7c27983..2dfae5e4 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-single-1.0-nudge.yaml @@ -5,18 +5,18 @@ description: | config_label: varda-single-1.0-nudge dates: - # start: 2025-01-01T06:00 - # end: 2025-12-31T06:00 - # frequency: 54h - - 2025-01-02T00:00 - - 2025-01-02T06:00 - - 2025-01-02T12:00 - - 2025-01-02T18:00 + start: 2025-01-01T06:00 + end: 2025-12-31T06:00 + frequency: 102h + # - 2025-01-02T00:00 + # - 2025-01-02T06:00 + # - 2025-01-02T12:00 + # - 2025-01-02T18:00 runs: - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 - label: Varda-rapid-0.1-all + label: Varda-rapid-0.1-nudge-all steps: 0/12/1 config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml extra_requirements: @@ -33,7 +33,7 @@ runs: - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 - label: Varda-rapid-0.1 + label: Varda-rapid-0.1-nudge-partial steps: 0/12/1 config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml extra_requirements: @@ -79,12 +79,6 @@ runs: truth: label: SwissMetNet root: jretrievedwh:1,2 - # root: jretrievedwh:bbox=45.8,47.8,5.9,10.5;seq_type=synop - # To verify against SwissMetNet observations from the DWH via jretrievedwh, - # set instead (requires jretrievedwh.py on $PATH and $OPR_HOME set): - # Other selectors: root: jretrievedwh:locations=ARO,KLO,LUG - # root: jretrievedwh:bbox=45.8,47.8,5.9,10.5 - # append ;stage=devt to target a non-prod DWH stage experiment: params: @@ -116,7 +110,19 @@ experiment: # - region # - init_hour - season - cross_validation: true + cross_validation: + # Stations withheld from nudging in the nudge-partial inference run. + exclude_stations: + - FAH + - DEM + - WAE + - GUT + - SAE + - CDM + - PMA + - SIA + - BIA + - SMM scorecards: enabled: true sections: diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml index 844b4e14..05249ee5 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml @@ -22,8 +22,9 @@ input: retrieve_observation: obs_path: observation/nudging_station_obs_raw.parquet jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src/data_input/ - bbox: [45.7, 48.0, 5.8, 10.8] - variables: [T_2M, TD_2M, U_10M, V_10M] + # bbox: [45.7, 48.0, 5.8, 10.8] + group: "1,2" + variables: [T_2M, TD_2M, U_10M, V_10M, TOT_PREC] use_limitation: 20 run_mode: "devt" - forward_transform_filter: @@ -37,7 +38,7 @@ input: icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files k: 3 power: 4.0 - max_dist: 0.5 + max_dist: 0.065 # in degrees, ~6.5 km run_mode: "devt" holdout_fraction: 0.0 holdout_seed: 42 diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml index 0446f0be..ce636a43 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml @@ -22,8 +22,9 @@ input: retrieve_observation: obs_path: observation/nudging_station_obs_raw.parquet jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src/data_input/ - bbox: [45.7, 48.0, 5.8, 10.8] - variables: [T_2M, TD_2M, U_10M, V_10M] + # bbox: [45.7, 48.0, 5.8, 10.8] + group: "1,2" + variables: [T_2M, TD_2M, U_10M, V_10M, TOT_PREC] use_limitation: 20 run_mode: "devt" - forward_transform_filter: @@ -37,17 +38,20 @@ input: icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files k: 3 power: 4.0 - max_dist: 0.5 + max_dist: 0.065 # in degrees, ~6.5 km run_mode: "devt" - holdout_fraction: 0.05 - holdout_seed: 42 - # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] - # Station holdout — use one of the two options below, not both. - # Option A: randomly withhold a fraction of stations (e.g. 5% for cross-validation). - # holdout_fraction: 0.05 - # holdout_seed: 42 # optional, defaults to 42 for reproducibility - # Option B: exclude specific stations by their nat_abbr identifier. - # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] + # Stations to withhold from nudging (held-out for cross-validation). + exclude_stations: + - FAH + - DEM + - WAE + - GUT + - SAE + - CDM + - PMA + - SIA + - BIA + - SMM namer: &namer rules: - - shortName: T diff --git a/resources/report/dashboard/script.js b/resources/report/dashboard/script.js index 238164a1..1dc90e36 100644 --- a/resources/report/dashboard/script.js +++ b/resources/report/dashboard/script.js @@ -38,10 +38,18 @@ initChoices("source-select"); initChoices("metric-select"); initChoices("param-select"); +// Station-group selector: single-select native {% endif %} + {% if cross_validation %} +
+ + +
+ {% endif %}
(may be absent when cross_validation is off) -const stationGroupSelect = document.getElementById("station-group-select"); -if (stationGroupSelect) stationGroupSelect.addEventListener("change", scheduleUpdate); - function getSelected(id) { return choicesInstances[id] ? choicesInstances[id].getValue(true) : []; } -function getStationGroup() { - return stationGroupSelect ? stationGroupSelect.value : "all"; -} - // --------------------------------------------------------------------------- // Data // --------------------------------------------------------------------------- @@ -60,8 +53,9 @@ function getStationGroup() { window.DATA = raw.data.map(row => { const obj = {}; for (let i = 0; i < cols.length; i++) obj[cols[i]] = row[i]; + const sgPart = obj.station_group !== "all" ? ", Group: " + obj.station_group : ""; obj.region_season_init = - "Region: " + obj.region + ", Season: " + obj.season + ", Init: " + obj.init_hour; + "Region: " + obj.region + ", Season: " + obj.season + ", Init: " + obj.init_hour + sgPart; return obj; }); })(); @@ -168,7 +162,7 @@ async function renderLegend(filteredData) { type: "nominal", legend: { orient: "right", - title: "Region / Season / Init", + title: "Region / Season / Init / Group", offset: 8, labelLimit: 400, symbolType: "circle", symbolSize: 120, @@ -211,18 +205,16 @@ async function updateChart() { const selMetrics = getSelected("metric-select"); const selParams = getSelected("param-select"); - const selStationGroup = getStationGroup(); + const selStationGroups = getSelected("station-group-select"); // Filter data by region / season / init / source / station_group // (metric and param are handled per cell) let filtered = DATA; - if (selRegions.length) filtered = filtered.filter(d => selRegions.includes(d.region)); - if (selSeasons.length) filtered = filtered.filter(d => selSeasons.includes(d.season)); - if (selInits.length) filtered = filtered.filter(d => selInits.includes(d.init_hour)); - if (selSources.length) filtered = filtered.filter(d => selSources.includes(d.source)); - // Always filter to the selected station_group so each source appears as one line. - // When cross_validation is off every row has station_group="all" and this is a no-op. - filtered = filtered.filter(d => d.station_group === selStationGroup); + if (selRegions.length) filtered = filtered.filter(d => selRegions.includes(d.region)); + if (selSeasons.length) filtered = filtered.filter(d => selSeasons.includes(d.season)); + if (selInits.length) filtered = filtered.filter(d => selInits.includes(d.init_hour)); + if (selSources.length) filtered = filtered.filter(d => selSources.includes(d.source)); + if (selStationGroups.length) filtered = filtered.filter(d => selStationGroups.includes(d.station_group)); // Show / hide table columns (params) document.querySelectorAll("#chart-table thead th[data-param]").forEach(th => { diff --git a/resources/report/dashboard/template.html.jinja2 b/resources/report/dashboard/template.html.jinja2 index 937eb8f9..201162f5 100644 --- a/resources/report/dashboard/template.html.jinja2 +++ b/resources/report/dashboard/template.html.jinja2 @@ -151,7 +151,7 @@ {% if cross_validation %}
- {% for sg in station_groups %} {% endfor %} From 7ffe364e1ba76d7576e1196eb55b01abbc8581f2 Mon Sep 17 00:00:00 2001 From: Luca Lanzilao Date: Thu, 20 Aug 2026 16:47:50 +0200 Subject: [PATCH 12/26] edit configs --- config/varda-single-1.0-nudge.yaml | 78 +++++++++---------- ...caler-global_trimedge_multi_nudge_all.yaml | 50 ++++++++++-- ...r-global_trimedge_multi_nudge_holdout.yaml | 42 ++++++++-- 3 files changed, 119 insertions(+), 51 deletions(-) diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-single-1.0-nudge.yaml index af66a1cc..60dd1054 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-single-1.0-nudge.yaml @@ -5,31 +5,31 @@ description: | config_label: varda-single-1.0-nudge dates: - # start: 2025-01-01T06:00 - # end: 2025-12-31T06:00 - # frequency: 102h - - 2025-01-02T00:00 - - 2025-04-02T06:00 - - 2025-07-02T12:00 - - 2025-11-02T18:00 + start: 2025-01-01T06:00 + end: 2025-12-31T06:00 + frequency: 18h #294h + # - 2025-01-02T00:00 + # - 2025-04-02T06:00 + # - 2025-07-02T12:00 + # - 2025-11-02T18:00 runs: - - temporal_downscaler: - checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 - label: Varda-rapid-0.1-nudge-all - steps: 0/12/1 - config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml - extra_requirements: - - anemoi-datasets==0.5.35 - - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss - # - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f91315e - - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - forecaster: - checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 - config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml - steps: 0/12/6 - extra_requirements: - - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + # - temporal_downscaler: + # checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 + # label: Varda-rapid-0.1-nudge-all + # steps: 0/12/1 + # config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml + # extra_requirements: + # - anemoi-datasets==0.5.35 + # - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss + # # - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f91315e + # - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + # forecaster: + # checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 + # config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml + # steps: 0/12/6 + # extra_requirements: + # - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 @@ -71,10 +71,10 @@ runs: # label: ICON-CH2-CTRL # root: /store_new/mch/msopr/osm/ICON-CH2-EPS # steps: 0/12/1 - - baseline: - label: ICON-CH1-CTRL - root: /store_new/mch/msopr/osm/ICON-CH1-EPS - steps: 0/12/1 + # - baseline: + # label: ICON-CH1-CTRL + # root: /store_new/mch/msopr/osm/ICON-CH1-EPS + # steps: 0/12/1 truth: label: SwissMetNet @@ -133,7 +133,6 @@ experiment: - BIA - CEV - OTL - scorecards: enabled: true sections: nowcasting: @@ -146,16 +145,16 @@ experiment: - "T_2M:RMSE,R2,ETS" - "TOT_PREC:RMSE,R2,ETS" - "TD_2M:RMSE,R2,ETS" - short_range: - baseline: ICON-CH1-CTRL - lead_times: "6/12/6" - stratification: region - variables: - - "U_10M:RMSE,R2,ETS" - - "V_10M:RMSE,R2,ETS" - - "T_2M:RMSE,R2,ETS" - - "TOT_PREC:RMSE,R2,ETS" - - "TD_2M:RMSE,R2,ETS" + # short_range: + # baseline: ICON-CH1-CTRL + # lead_times: "6/12/6" + # stratification: region + # variables: + # - "U_10M:RMSE,R2,ETS" + # - "V_10M:RMSE,R2,ETS" + # - "T_2M:RMSE,R2,ETS" + # - "TOT_PREC:RMSE,R2,ETS" + # - "TD_2M:RMSE,R2,ETS" # medium_range: # baseline: ICON-CH2-CTRL # lead_times: "24/24/24" @@ -171,10 +170,11 @@ experiment: showcase: params: - T_2M + - TD_2M - U_10M - V_10M meteograms: - enabled: true + enabled: false stations: [JUN, KLO, LUG, GVE] #, COV, GOR, WFJ, SAE, SAM, DAV, ZER, ANT, VSBAS, BRT, LTB, GOS, CEV, BIA] animations: fps: 0.5 diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml index fa90b638..63060a08 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml @@ -25,8 +25,15 @@ input: # bbox: [45.7, 48.0, 5.8, 10.8] group: "1,2" variables: [T_2M, TD_2M, U_10M, V_10M, TOT_PREC] - use_limitation: 20 + use_limitation: 40 run_mode: "devt" + # Extra trim on top of group/bbox above, matching + # notebooks/d_eff_generator.ipynb's station selection exactly — + # keeps the retrieved stations a subset of whatever + # nudge_toward_observation's d_eff_file cache below was built + # from (a station outside that cache raises an error there). + station_filter_mode: domain + domain_bbox: [45.7, 48.0, 5.8, 10.8] - forward_transform_filter: clean_observation: obs_path_in: observation/nudging_station_obs_raw.parquet @@ -36,17 +43,46 @@ input: obs_path: observation/nudging_station_obs.parquet nudge_variables: [T_2M, TD_2M, U_10M, V_10M] icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files - k: 3 - power: 4.0 - max_dist: 0.12 # in degrees, ~12 km + topo_file: /scratch/mch/llanzila/sruc/aux_files/topo_descriptors_icon_R19B08.nc + dem_barrier_file: /store_new/mch/msclim/appclim/data/grids/topodem/v2/topo/radar_100/topo_DEM_1000M.nc + # Barrier-aware effective distances (d_eff) are precomputed offline + # (notebooks/d_eff_generator.ipynb) rather than computed live — this + # filter no longer accepts n_barrier_samples/n_barrier_width_samples/ + # barrier_width_m/elev_scale/elev_diff_scale, since those are now + # baked into the cache file itself. Filename encodes the exact + # settings it was built with: station_filter_mode=domain, + # max_dist=50, n_barrier_samples=50, n_barrier_width_samples=3, + # barrier_width_m=1500, elev_scale=50, elev_diff_scale=100, and the + # 1939-station catalog from group="1,2" + domain_bbox trim above. + d_eff_file: /scratch/mch/llanzila/sruc/aux_files/d_eff_cache_domain_maxdist50km_nbar50x3_bw1500m_elev50_elevdiff100_nsta1939.nc + weight_power: 2.0 + max_dist: 50.0 #km, station influence radius (was 0.4 projected degrees; 0.4 * 111.32 km/deg) + min_topo_w: 0.2 + lim_effective: 0.0 + # Station reliability (v4): shrinks a station's own influence radius + # (rather than uniformly down-weighting its contribution) when its + # residual disagrees with what its neighbours predict for it via a + # leave-one-out spatial-consistency check. number_of_std is the Tukey + # biweight rejection threshold in robust (median/MAD) sigmas; a station + # at or beyond this many sigmas of disagreement gets reliability=0 and + # its radius shrinks to reliability_min_dist_frac * max_dist (never all + # the way to zero). Set use_reliability_check: false to fall back to the + # pre-v4 behaviour (a single global max_dist shared by every station). + use_reliability_check: true + number_of_std: 5.0 + reliability_min_dist_frac: 0.05 + # Saves a 3-panel PNG (station residual, station reliability, gridded + # correction) per nudged variable per reference time. No effect unless + # use_reliability_check is true; never affects the nudging correction + # itself (plotting failures are logged and skipped). + plot_dir: output/nudging_reliability_diagnostics run_mode: "devt" holdout_fraction: 0.0 holdout_seed: 42 - # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] - # Station holdout — use one of the two options below, not both. + # Station holdout — use one of the two options below, not both. # Option A: randomly withhold a fraction of stations (e.g. 5% for cross-validation). # holdout_fraction: 0.05 - # holdout_seed: 42 # optional, defaults to 42 for reproducibility + # holdout_seed: 42 # Option B: exclude specific stations by their nat_abbr identifier. # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] namer: &namer diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml index f2bfcaf6..6a93613d 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml @@ -25,8 +25,15 @@ input: # bbox: [45.7, 48.0, 5.8, 10.8] group: "1,2" variables: [T_2M, TD_2M, U_10M, V_10M, TOT_PREC] - use_limitation: 20 + use_limitation: 40 run_mode: "devt" + # Extra trim on top of group/bbox above, matching + # notebooks/d_eff_generator.ipynb's station selection exactly — + # keeps the retrieved stations a subset of whatever + # nudge_toward_observation's d_eff_file cache below was built + # from (a station outside that cache raises an error there). + station_filter_mode: domain + domain_bbox: [45.7, 48.0, 5.8, 10.8] - forward_transform_filter: clean_observation: obs_path_in: observation/nudging_station_obs_raw.parquet @@ -36,12 +43,37 @@ input: obs_path: observation/nudging_station_obs.parquet nudge_variables: [T_2M, TD_2M, U_10M, V_10M] icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files - k: 3 - power: 4.0 - max_dist: 0.12 # in degrees, ~12 km + topo_file: /scratch/mch/llanzila/sruc/aux_files/topo_descriptors_icon_R19B08.nc + dem_barrier_file: /store_new/mch/msclim/appclim/data/grids/topodem/v2/topo/radar_100/topo_DEM_1000M.nc + weight_power: 2.0 + max_dist: 50.0 #44.528 # km, station influence radius (was 0.4 projected degrees; 0.4 * 111.32 km/deg) + n_barrier_samples: 50 #40 + n_barrier_width_samples: 3 + barrier_width_m: 1500 + elev_scale: 50 #26.949 # m/km, ridge height adding 1 km of effective distance (was 3000 m/deg; 3000 / 111.32 km/deg) + elev_diff_scale: 100 #35.932 # m/km, endpoint elevation difference adding 1 km (was 4000 m/deg; 4000 / 111.32 km/deg) + min_topo_w: 0.2 + lim_effective: 0.0 + # Station reliability (v4): shrinks a station's own influence radius + # (rather than uniformly down-weighting its contribution) when its + # residual disagrees with what its neighbours predict for it via a + # leave-one-out spatial-consistency check. number_of_std is the Tukey + # biweight rejection threshold in robust (median/MAD) sigmas; a station + # at or beyond this many sigmas of disagreement gets reliability=0 and + # its radius shrinks to reliability_min_dist_frac * max_dist (never all + # the way to zero). Set use_reliability_check: false to fall back to the + # pre-v4 behaviour (a single global max_dist shared by every station). + use_reliability_check: true + number_of_std: 5.0 + reliability_min_dist_frac: 0.1 + # Saves a 3-panel PNG (station residual, station reliability, gridded + # correction) per nudged variable per reference time. No effect unless + # use_reliability_check is true; never affects the nudging correction + # itself (plotting failures are logged and skipped). + plot_dir: output/nudging_reliability_diagnostics run_mode: "devt" # Stations to withhold from nudging (held-out for cross-validation). - exclude_stations: + exclude_stations: - COY - FRE - RUE From 05721497b659dfbe9d5d867f17134ca6a8b1fdd6 Mon Sep 17 00:00:00 2001 From: lclanzi Date: Mon, 7 Sep 2026 15:52:05 +0200 Subject: [PATCH 13/26] update cross validation methodology, add filtering to jretrieve output --- ...le-1.0-nudge.yaml => varda-rapid-0.2.yaml} | 70 ++---- ...caler-global_trimedge_multi_nudge_all.yaml | 204 ------------------ ...r-global_trimedge_multi_with_nudging.yaml} | 78 ++----- src/data_input/__init__.py | 67 +++++- src/data_input/jretrieve.py | 37 +++- tests/unit/test_jretrieve.py | 40 ++++ workflow/rules/common.smk | 2 +- workflow/rules/inference.smk | 6 + workflow/rules/verification.smk | 8 +- workflow/scripts/inference_prepare.py | 68 +++++- workflow/scripts/plot_meteogram.py | 2 +- 11 files changed, 252 insertions(+), 330 deletions(-) rename config/{varda-single-1.0-nudge.yaml => varda-rapid-0.2.yaml} (63%) delete mode 100644 resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml rename resources/inference/configs/{sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml => sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml} (62%) diff --git a/config/varda-single-1.0-nudge.yaml b/config/varda-rapid-0.2.yaml similarity index 63% rename from config/varda-single-1.0-nudge.yaml rename to config/varda-rapid-0.2.yaml index 60dd1054..747122d6 100644 --- a/config/varda-single-1.0-nudge.yaml +++ b/config/varda-rapid-0.2.yaml @@ -7,39 +7,21 @@ config_label: varda-single-1.0-nudge dates: start: 2025-01-01T06:00 end: 2025-12-31T06:00 - frequency: 18h #294h - # - 2025-01-02T00:00 + frequency: 294h + # - 2025-03-01T00:00 # - 2025-04-02T06:00 # - 2025-07-02T12:00 # - 2025-11-02T18:00 runs: - # - temporal_downscaler: - # checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 - # label: Varda-rapid-0.1-nudge-all - # steps: 0/12/1 - # config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml - # extra_requirements: - # - anemoi-datasets==0.5.35 - # - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss - # # - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f91315e - # - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - # forecaster: - # checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 - # config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml - # steps: 0/12/6 - # extra_requirements: - # - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef - - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 - label: Varda-rapid-0.1-nudge-partial + label: Varda-rapid-0.1 steps: 0/12/1 - config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml + config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml extra_requirements: - anemoi-datasets==0.5.35 - - -e /scratch/mch/llanzila/sruc/anemoi-plugins-meteoswiss - # - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f91315e + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@6e6b2dd42e66d9c981bd8f5ccca657f0411e4ccc - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 @@ -67,18 +49,15 @@ runs: label: INCA root: /store_new/mch/msclim/INCA steps: 0/6/1 - # - baseline: - # label: ICON-CH2-CTRL - # root: /store_new/mch/msopr/osm/ICON-CH2-EPS - # steps: 0/12/1 - # - baseline: - # label: ICON-CH1-CTRL - # root: /store_new/mch/msopr/osm/ICON-CH1-EPS - # steps: 0/12/1 truth: label: SwissMetNet - root: jretrievedwh:1,2 + # Retrieve over the broad ICON domain (matching notebooks/d_eff_generator.ipynb's + # STATIONS_SEL), then trim to stations within the real Swiss national border + # (filter_mode=switzerland) or within a predefined bounding box (filter_mode=domain;domain_bbox=45.7,48.0,5.8,10.8). + root: jretrievedwh:bbox=40.5,53.0,0.0,17.5;use_limitation=40;filter_mode=switzerland + # root: jretrievedwh:bbox=40.5,53.0,0.0,17.5;use_limitation=40;filter_mode=domain;domain_bbox=45.7,48.0,5.8,10.8 + experiment: params: @@ -87,6 +66,8 @@ experiment: - U_10M - V_10M - TOT_PREC + - PMSL + - PS stratification: regions: - mittelland @@ -107,13 +88,11 @@ experiment: gt: [288.15, 298.15] dashboard: stratification: - # - region - # - init_hour - season cross_validation: # Stations withheld from nudging in the nudge-partial inference run. exclude_stations: - - COY + - CHM - FRE - RUE - BAS @@ -122,7 +101,7 @@ experiment: - HOE - PUY - RAG - - SAE + - EIN - ALT - GOS - COV @@ -145,25 +124,6 @@ experiment: - "T_2M:RMSE,R2,ETS" - "TOT_PREC:RMSE,R2,ETS" - "TD_2M:RMSE,R2,ETS" - # short_range: - # baseline: ICON-CH1-CTRL - # lead_times: "6/12/6" - # stratification: region - # variables: - # - "U_10M:RMSE,R2,ETS" - # - "V_10M:RMSE,R2,ETS" - # - "T_2M:RMSE,R2,ETS" - # - "TOT_PREC:RMSE,R2,ETS" - # - "TD_2M:RMSE,R2,ETS" - # medium_range: - # baseline: ICON-CH2-CTRL - # lead_times: "24/24/24" - # stratification: region - # variables: - # - "U_10M:RMSE,R2,ETS" - # - "V_10M:RMSE,R2,ETS" - # - "T_2M:RMSE,R2,ETS" - # - "TOT_PREC:RMSE,R2,ETS" scoremaps: enabled: false diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml deleted file mode 100644 index 63060a08..00000000 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_all.yaml +++ /dev/null @@ -1,204 +0,0 @@ -runner: temporal_downscaler - -input: - cutout: - - lam_0: - grib: - path: forecaster/20* - pre_processors: - - forward_transform_filter: - # does not have an effect if the temporal downscaler does not use tp as prognostic variable - rescale: - scale: 0.001 # convert from kg m-2 to m - offset: 0 - param: tp - # --- Nudging (disabled by default) --- - # Uncomment to nudge the initial condition toward station observations. - # 1. retrieve_observation: fetches live data from DWH via jretrieve → raw Parquet - # 2. clean_observation: reads raw Parquet, applies QC cleaning → cleaned Parquet - # 3. nudge_toward_observation: reads cleaned Parquet, applies IoR correction - # nudge_variables: list of GRIB shortNames to nudge; omit to nudge all (T_2M, U_10M, V_10M, TOT_PREC). - - forward_transform_filter: - retrieve_observation: - obs_path: observation/nudging_station_obs_raw.parquet - jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src/data_input/ - # bbox: [45.7, 48.0, 5.8, 10.8] - group: "1,2" - variables: [T_2M, TD_2M, U_10M, V_10M, TOT_PREC] - use_limitation: 40 - run_mode: "devt" - # Extra trim on top of group/bbox above, matching - # notebooks/d_eff_generator.ipynb's station selection exactly — - # keeps the retrieved stations a subset of whatever - # nudge_toward_observation's d_eff_file cache below was built - # from (a station outside that cache raises an error there). - station_filter_mode: domain - domain_bbox: [45.7, 48.0, 5.8, 10.8] - - forward_transform_filter: - clean_observation: - obs_path_in: observation/nudging_station_obs_raw.parquet - obs_path_out: observation/nudging_station_obs.parquet - - forward_transform_filter: - nudge_toward_observation: - obs_path: observation/nudging_station_obs.parquet - nudge_variables: [T_2M, TD_2M, U_10M, V_10M] - icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files - topo_file: /scratch/mch/llanzila/sruc/aux_files/topo_descriptors_icon_R19B08.nc - dem_barrier_file: /store_new/mch/msclim/appclim/data/grids/topodem/v2/topo/radar_100/topo_DEM_1000M.nc - # Barrier-aware effective distances (d_eff) are precomputed offline - # (notebooks/d_eff_generator.ipynb) rather than computed live — this - # filter no longer accepts n_barrier_samples/n_barrier_width_samples/ - # barrier_width_m/elev_scale/elev_diff_scale, since those are now - # baked into the cache file itself. Filename encodes the exact - # settings it was built with: station_filter_mode=domain, - # max_dist=50, n_barrier_samples=50, n_barrier_width_samples=3, - # barrier_width_m=1500, elev_scale=50, elev_diff_scale=100, and the - # 1939-station catalog from group="1,2" + domain_bbox trim above. - d_eff_file: /scratch/mch/llanzila/sruc/aux_files/d_eff_cache_domain_maxdist50km_nbar50x3_bw1500m_elev50_elevdiff100_nsta1939.nc - weight_power: 2.0 - max_dist: 50.0 #km, station influence radius (was 0.4 projected degrees; 0.4 * 111.32 km/deg) - min_topo_w: 0.2 - lim_effective: 0.0 - # Station reliability (v4): shrinks a station's own influence radius - # (rather than uniformly down-weighting its contribution) when its - # residual disagrees with what its neighbours predict for it via a - # leave-one-out spatial-consistency check. number_of_std is the Tukey - # biweight rejection threshold in robust (median/MAD) sigmas; a station - # at or beyond this many sigmas of disagreement gets reliability=0 and - # its radius shrinks to reliability_min_dist_frac * max_dist (never all - # the way to zero). Set use_reliability_check: false to fall back to the - # pre-v4 behaviour (a single global max_dist shared by every station). - use_reliability_check: true - number_of_std: 5.0 - reliability_min_dist_frac: 0.05 - # Saves a 3-panel PNG (station residual, station reliability, gridded - # correction) per nudged variable per reference time. No effect unless - # use_reliability_check is true; never affects the nudging correction - # itself (plotting failures are logged and skipped). - plot_dir: output/nudging_reliability_diagnostics - run_mode: "devt" - holdout_fraction: 0.0 - holdout_seed: 42 - # Station holdout — use one of the two options below, not both. - # Option A: randomly withhold a fraction of stations (e.g. 5% for cross-validation). - # holdout_fraction: 0.05 - # holdout_seed: 42 - # Option B: exclude specific stations by their nat_abbr identifier. - # exclude_stations: ['SCU', 'PAY', 'INNSWZ', 'WSLBTF', 'CDF', 'NABDUE', 'MMSAS', 'NABZUE', 'MAS', 'MMERZ', 'THU', 'OBR', 'FLTRL', 'WSLHOB', 'MMSAA'] - namer: &namer - rules: - - - shortName: T - - t_{level} - - - shortName: U - - u_{level} - - - shortName: V - - v_{level} - - - shortName: W - - w_{level} - - - shortName: QV - - q_{level} - - - shortName: FI - - z_{level} - - - shortName: PMSL - - msl - - - shortName: FIS - - z - - - shortName: PS - - sp - - - shortName: T_2M - - 2t - - - shortName: TD_2M - - 2d - - - shortName: T_G - - skt - - - shortName: U_10M - - 10u - - - shortName: V_10M - - 10v - - - shortName: FR_LAND - - lsm - - - shortName: TOT_PREC - - tp - - global: - grib: - path: forecaster/ifs* - namer: *namer - -constant_forcings: - test: - use_original_paths: true - -patch_metadata: resources/sgm-temporal-downscaler-ich1-oper-patch.yaml - -post_processors: - - accumulate_from_start_of_forecast: # accumulate tp from start of forecast - accumulations: - - tp - - forward_transform_filter: - rescale: - scale: 1000 # convert units from m to kg m-2 - offset: 0 - param: tp - -output: - tee: - - grib: - path: grib/{dateTime}_{step:03}.grib - encoding: - typeOfGeneratingProcess: 2 - templates: - samples: resources/templates_index_icon.yaml - post_processors: - - extract_mask: # removes global points - mask: "lam_0/cutout_mask" - as_slice: true - # here, the trimedge mask can be specified when available - - - grib: - path: grib/ifs-{dateTime}_{step:03}.grib - encoding: - typeOfGeneratingProcess: 2 - templates: - samples: resources/templates_index_ifs.yaml - post_processors: - - extract_mask: # removes lam points - mask: "lam_0/cutout_mask" - as_slice: true - inverse: true - - assign_mask: # fill local/global overlapping points with nan - mask: "global/cutout_mask" - modifiers: - - patches: - - variable: - U_10M: {"param": 165, "shortName": "10u"} - 10u: {"param": 165, "shortName": "10u"} - V_10M: {"param": 166, "shortName": "10v"} - 10v: {"param": 166, "shortName": "10v"} - TD_2M: {"param": 168, "shortName": "2d"} - 2d: {"param": 168, "shortName": "2d"} - T_2M: {"param": 167, "shortName": "2t"} - 2t: {"param": 167, "shortName": "2t"} - FR_LAND: {"param": 172, "shortName": "lsm"} - lsm: {"param": 172, "shortName": "lsm"} - PMSL: {"param": 151, "shortName": "msl"} - msl: {"param": 151, "shortName": "msl"} - PS: {"param": 134, "shortName": "sp"} - sp: {"param": 134, "shortName": "sp"} - SSO_SIGMA: {"param": 163, "shortName": "slor"} - slor: {"param": 163, "shortName": "slor"} - SSO_STDH: {"param": 160, "shortName": "sdor"} - sdor: {"param": 160, "shortName": "sdor"} - TOT_PREC: {"param": 228, "shortName": "tp"} - tp: {"param": 228, "shortName": "tp"} - z: {"param": 129, "shortName": "z"} - "^q_(\\d+)$": {"param": 133, "shortName": "q"} - "^t_(\\d+)$": {"param": 130, "shortName": "t"} - "^u_(\\d+)$": {"param": 131, "shortName": "u"} - "^v_(\\d+)$": {"param": 132, "shortName": "v"} - "^w_(\\d+)$": {"param": 135, "shortName": "w"} - "^z_(\\d+)$": {"param": 129, "shortName": "z"} - -# silenced due to bug in anemoi-inference for multi-step temporal downscalers, can be removed when fixed -verbosity: 0 -allow_nans: true -output_frequency: "1h" diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml similarity index 62% rename from resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml rename to resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml index 6a93613d..6b3f4e97 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_nudge_holdout.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml @@ -12,28 +12,18 @@ input: scale: 0.001 # convert from kg m-2 to m offset: 0 param: tp - # --- Nudging (disabled by default) --- - # Uncomment to nudge the initial condition toward station observations. - # 1. retrieve_observation: fetches live data from DWH via jretrieve → raw Parquet - # 2. clean_observation: reads raw Parquet, applies QC cleaning → cleaned Parquet - # 3. nudge_toward_observation: reads cleaned Parquet, applies IoR correction - # nudge_variables: list of GRIB shortNames to nudge; omit to nudge all (T_2M, U_10M, V_10M, TOT_PREC). + # nudge the initial condition toward station observations. - forward_transform_filter: retrieve_observation: obs_path: observation/nudging_station_obs_raw.parquet jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src/data_input/ - # bbox: [45.7, 48.0, 5.8, 10.8] - group: "1,2" - variables: [T_2M, TD_2M, U_10M, V_10M, TOT_PREC] + retrieval_bbox: [40.5, 53.0, 0.0, 17.5] + # station_group: "1,2" + variables: [T_2M, TD_2M, U_10M, V_10M, TOT_PREC, PMSL, PS] use_limitation: 40 run_mode: "devt" - # Extra trim on top of group/bbox above, matching - # notebooks/d_eff_generator.ipynb's station selection exactly — - # keeps the retrieved stations a subset of whatever - # nudge_toward_observation's d_eff_file cache below was built - # from (a station outside that cache raises an error there). station_filter_mode: domain - domain_bbox: [45.7, 48.0, 5.8, 10.8] + trim_bbox: [45.7, 48.0, 5.8, 10.8] - forward_transform_filter: clean_observation: obs_path_in: observation/nudging_station_obs_raw.parquet @@ -41,59 +31,31 @@ input: - forward_transform_filter: nudge_toward_observation: obs_path: observation/nudging_station_obs.parquet - nudge_variables: [T_2M, TD_2M, U_10M, V_10M] + nudge_variables: [T_2M, TD_2M, U_10M, V_10M, PMSL, PS] + temperature_lapse_rate: 0.0065 # K/m + pressure_lapse_rate: 11.5 # Pa/m icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files topo_file: /scratch/mch/llanzila/sruc/aux_files/topo_descriptors_icon_R19B08.nc dem_barrier_file: /store_new/mch/msclim/appclim/data/grids/topodem/v2/topo/radar_100/topo_DEM_1000M.nc + d_eff_file: /scratch/mch/llanzila/sruc/aux_files/d_eff_cache_domain_maxdist50km_nbar50x3_bw1500m_elev50_elevdiff100_nsta1939.nc weight_power: 2.0 - max_dist: 50.0 #44.528 # km, station influence radius (was 0.4 projected degrees; 0.4 * 111.32 km/deg) - n_barrier_samples: 50 #40 - n_barrier_width_samples: 3 - barrier_width_m: 1500 - elev_scale: 50 #26.949 # m/km, ridge height adding 1 km of effective distance (was 3000 m/deg; 3000 / 111.32 km/deg) - elev_diff_scale: 100 #35.932 # m/km, endpoint elevation difference adding 1 km (was 4000 m/deg; 4000 / 111.32 km/deg) + max_dist: 50000.0 # meters (50 km), station influence radius + variable_overrides: + U_10M: + d_eff_file: /scratch/mch/llanzila/sruc/aux_files/d_eff_cache_domain_maxdist50km_nbar50x3_bw1500m_elev50_elevdiff100_nsta1939.nc + max_dist: 50000.0 + V_10M: + d_eff_file: /scratch/mch/llanzila/sruc/aux_files/d_eff_cache_domain_maxdist50km_nbar50x3_bw1500m_elev50_elevdiff100_nsta1939.nc + max_dist: 50000.0 + use_topo_descriptors: false min_topo_w: 0.2 lim_effective: 0.0 - # Station reliability (v4): shrinks a station's own influence radius - # (rather than uniformly down-weighting its contribution) when its - # residual disagrees with what its neighbours predict for it via a - # leave-one-out spatial-consistency check. number_of_std is the Tukey - # biweight rejection threshold in robust (median/MAD) sigmas; a station - # at or beyond this many sigmas of disagreement gets reliability=0 and - # its radius shrinks to reliability_min_dist_frac * max_dist (never all - # the way to zero). Set use_reliability_check: false to fall back to the - # pre-v4 behaviour (a single global max_dist shared by every station). use_reliability_check: true number_of_std: 5.0 - reliability_min_dist_frac: 0.1 - # Saves a 3-panel PNG (station residual, station reliability, gridded - # correction) per nudged variable per reference time. No effect unless - # use_reliability_check is true; never affects the nudging correction - # itself (plotting failures are logged and skipped). + reliability_min_dist_frac: 0.0 + enable_plotting: true plot_dir: output/nudging_reliability_diagnostics run_mode: "devt" - # Stations to withhold from nudging (held-out for cross-validation). - exclude_stations: - - COY - - FRE - - RUE - - BAS - - PLF - - CGI - - HOE - - PUY - - RAG - - SAE - - ALT - - GOS - - COV - - ULR - - BIN - - GUE - - ROE - - BIA - - CEV - - OTL namer: &namer rules: - - shortName: T diff --git a/src/data_input/__init__.py b/src/data_input/__init__.py index d9d101ab..071e9782 100644 --- a/src/data_input/__init__.py +++ b/src/data_input/__init__.py @@ -383,6 +383,51 @@ def _jretrieve_df_to_xarray(df, short_names, catalog) -> xr.Dataset: return xr.Dataset(data_vars=data_vars, coords=coords) +def _trim_stations_xr(ds: xr.Dataset, filter_mode: str, domain_bbox: list | None) -> xr.Dataset: + """Trim the "values" (station) dimension to *filter_mode* — same logic as + RetrieveObservation._trim_stations / notebooks/d_eff_generator.ipynb's + station-trim cell, adapted for an xarray Dataset with latitude/longitude + coords along "values" instead of a pandas DataFrame.""" + lat = ds["latitude"].values + lon = ds["longitude"].values + + if filter_mode == "domain": + lat_min, lat_max, lon_min, lon_max = domain_bbox + mask = (lat >= lat_min) & (lat <= lat_max) & (lon >= lon_min) & (lon <= lon_max) + desc = f"domain bbox {domain_bbox}" + + elif filter_mode == "switzerland": + import cartopy.io.shapereader as shpreader + from shapely.geometry import Point + + shp_path = shpreader.natural_earth( + resolution="10m", category="cultural", name="admin_0_countries" + ) + ch_country = next( + r for r in shpreader.Reader(shp_path).records() + if r.attributes["ADM0_A3"] == "CHE" + ) + swiss_geom = ch_country.geometry + + def _in_switzerland(lat_, lon_): + if pd.isna(lat_) or pd.isna(lon_): + return False + return swiss_geom.contains(Point(lon_, lat_)) + + mask = np.array([_in_switzerland(la, lo) for la, lo in zip(lat, lon)]) + desc = "Swiss national border (Natural Earth)" + + else: + raise ValueError( + f"Unknown filter_mode: {filter_mode!r} (expected 'domain' or 'switzerland')" + ) + + n_before = ds.sizes["values"] + ds = ds.isel(values=mask) + LOG.info("Station filter [%s]: %d -> %d stations", desc, n_before, ds.sizes["values"]) + return ds + + def load_obs_data_from_jretrieve( root, reftime: datetime, steps: list[int], params: list[str] ) -> xr.Dataset: @@ -413,7 +458,7 @@ def load_obs_data_from_jretrieve( from data_input import jretrieve as jr - stations, stage, seq_type = jr.parse_selection(root) + stations, stage, seq_type, use_limitation, filter_mode, domain_bbox = jr.parse_selection(root) jr.check_prerequisites(stage) want_uv = "U_10M" in params or "V_10M" in params @@ -443,9 +488,13 @@ def load_obs_data_from_jretrieve( increment_minutes=step_hours * 60, seq_type=seq_type, stage=stage, + use_limitation=use_limitation, ) raw = _jretrieve_df_to_xarray(df, short_names, catalog) + if filter_mode is not None: + raw = _trim_stations_xr(raw, filter_mode, domain_bbox) + out = xr.Dataset(coords=raw.coords) for icon, short in DWH_PARAM_MAP.items(): if icon in params and short in raw: @@ -465,7 +514,21 @@ def load_obs_data_from_jretrieve( out = out.dropna("values", how="all") times = np.datetime64(reftime) + np.asarray(steps, dtype="timedelta64[h]") - return _select_valid_times(out, times, strict=True) + result = _select_valid_times(out, times, strict=True) + + # Same per-variable station-coverage log as RetrieveObservation, so ground-truth + # coverage can be compared directly against what was actually available to nudge. + _icon_to_short = { + "T_2M": "2t", "TD_2M": "2d", "U_10M": "10u", "V_10M": "10v", + "PMSL": "msl", "TOT_PREC": "tp", "VMAX_10M": "vmax", + } + n_total = result.sizes["values"] + for icon, short in _icon_to_short.items(): + if icon in result.data_vars: + n_valid = int(result[icon].notnull().any("time").sum()) + LOG.info("Stations with valid %s: %d / %d stations", short, n_valid, n_total) + + return result def load_truth_data( diff --git a/src/data_input/jretrieve.py b/src/data_input/jretrieve.py index 131508d3..933706c9 100644 --- a/src/data_input/jretrieve.py +++ b/src/data_input/jretrieve.py @@ -115,20 +115,35 @@ def _stations_to_argv(stations: dict[str, Any]) -> list[str]: raise AssertionError("unreachable") -def parse_selection(root: Any) -> tuple[dict[str, Any], str, str]: - """Parse a truth-root marker into (stations, stage, seq_type). +def parse_selection( + root: Any, +) -> tuple[dict[str, Any], str, str, int | None, str | None, list | None]: + """Parse a truth-root marker into + (stations, stage, seq_type, use_limitation, filter_mode, domain_bbox). Examples (slash-free so they survive ``Path()`` normalisation): ``jretrievedwh:SwissMetNet`` -> group ``jretrievedwh:group=SwissMetNet;stage=devt`` ``jretrievedwh:locations=ARO,KLO`` ``jretrievedwh:bbox=45.8,47.8,5.9,10.5`` + ``jretrievedwh:bbox=40.5,53.0,0.0,17.5;use_limitation=40;filter_mode=switzerland`` + -> retrieve over the given bbox, then trim to stations within the real + Swiss national border (see load_obs_data_from_jretrieve). + ``jretrievedwh:bbox=40.5,53.0,0.0,17.5;filter_mode=domain;domain_bbox=45.7,48.0,5.8,10.8`` + -> retrieve over one bbox, then trim to a second, independent bbox — + mirrors RetrieveObservation's bbox/station_filter_mode+domain_bbox split. + + use_limitation/filter_mode/domain_bbox default to None (no time-window + limit, no extra trim) when not given, unchanged from before they existed. """ _, _, rest = str(root).partition(":") rest = rest.strip() stations: dict[str, Any] = {} stage = "prod" seq_type = "surface" + use_limitation: int | None = None + filter_mode: str | None = None + domain_bbox: list | None = None for i, part in enumerate([p for p in rest.split(";") if p]): if "=" not in part: if i == 0: @@ -143,11 +158,27 @@ def parse_selection(root: Any) -> tuple[dict[str, Any], str, str]: stage = value elif key == "seq_type": seq_type = value + elif key == "use_limitation": + use_limitation = int(value) + elif key == "filter_mode": + if value not in ("domain", "switzerland"): + raise ValueError( + f"filter_mode must be 'domain' or 'switzerland', got {value!r}" + ) + filter_mode = value + elif key == "domain_bbox": + domain_bbox = [float(v) for v in value.split(",") if v] + if len(domain_bbox) != 4: + raise ValueError( + "domain_bbox must be lat_min,lat_max,lon_min,lon_max." + ) else: raise ValueError(f"Unknown jretrieve selector key: {key!r}") if not stations: stations = {"group": DEFAULT_GROUP} - return stations, stage, seq_type + if filter_mode == "domain" and domain_bbox is None: + raise ValueError("domain_bbox is required when filter_mode='domain'.") + return stations, stage, seq_type, use_limitation, filter_mode, domain_bbox def _run(argv: list[str], env: dict[str, str], timeout_s: int) -> str: diff --git a/tests/unit/test_jretrieve.py b/tests/unit/test_jretrieve.py index 2509ebda..1f69d28f 100644 --- a/tests/unit/test_jretrieve.py +++ b/tests/unit/test_jretrieve.py @@ -37,11 +37,17 @@ def test_parse_selection_default_group(): {"group": "SwissMetNet"}, "prod", "surface", + None, + None, + None, ) assert jr.parse_selection("jretrievedwh:SwissMetNet") == ( {"group": "SwissMetNet"}, "prod", "surface", + None, + None, + None, ) @@ -50,6 +56,40 @@ def test_parse_selection_keyvalue_and_stage(): {"locations": "ARO,KLO"}, "devt", "surface", + None, + None, + None, + ) + + +def test_parse_selection_use_limitation_and_filter_mode_switzerland(): + assert jr.parse_selection( + "jretrievedwh:bbox=40.5,53.0,0.0,17.5;use_limitation=40;filter_mode=switzerland" + ) == ( + {"bbox": "40.5,53.0,0.0,17.5"}, + "prod", + "surface", + 40, + "switzerland", + None, + ) + + +def test_parse_selection_filter_mode_domain_requires_domain_bbox(): + with pytest.raises(ValueError, match="domain_bbox"): + jr.parse_selection("jretrievedwh:bbox=40.5,53.0,0.0,17.5;filter_mode=domain") + + +def test_parse_selection_filter_mode_domain_with_bbox(): + assert jr.parse_selection( + "jretrievedwh:bbox=40.5,53.0,0.0,17.5;filter_mode=domain;domain_bbox=45.7,48.0,5.8,10.8" + ) == ( + {"bbox": "40.5,53.0,0.0,17.5"}, + "prod", + "surface", + None, + "domain", + [45.7, 48.0, 5.8, 10.8], ) diff --git a/workflow/rules/common.smk b/workflow/rules/common.smk index 6220e3d3..4fee8ddf 100644 --- a/workflow/rules/common.smk +++ b/workflow/rules/common.smk @@ -370,7 +370,7 @@ def truth_file_dep(_): if "jretrieve" in str(config["truth"]["root"]): from data_input.jretrieve import check_prerequisites, parse_selection - _, _jretrieve_stage, _ = parse_selection(config["truth"]["root"]) + _, _jretrieve_stage, _, _, _, _ = parse_selection(config["truth"]["root"]) check_prerequisites(_jretrieve_stage) diff --git a/workflow/rules/inference.smk b/workflow/rules/inference.smk index 144afc8e..9e5ed873 100644 --- a/workflow/rules/inference.smk +++ b/workflow/rules/inference.smk @@ -251,6 +251,12 @@ rule inference_prepare_temporal_downscaler: if RUN_CONFIGS[wc.run_id].get("forecaster") is None else _get_forecaster_run_id(wc.run_id) ), + # Single source of truth for the holdout station set: experiment.cross_validation + # in the top-level experiment config. Injected into any nudge_toward_observation + # filter found in the inference config (a no-op if the config has none, or if + # cross-validation isn't configured for this experiment) — see + # inference_prepare.py::_inject_nudging_cross_validation. + cross_validation_cfg=CROSS_VALIDATION_CFG, script: "../scripts/inference_prepare.py" diff --git a/workflow/rules/verification.smk b/workflow/rules/verification.smk index 494598d1..3978f227 100644 --- a/workflow/rules/verification.smk +++ b/workflow/rules/verification.smk @@ -41,7 +41,7 @@ rule verification_metrics_baseline: export ECCODES_DEFINITION_PATH=$(realpath .venv/share/eccodes-cosmo-resources/definitions) uv run {input.script} \ --forecast {input.forecast} \ - --truth {params.truth} \ + --truth "{params.truth}" \ --reftime {wildcards.init_time} \ --steps "{params.baseline_steps}" \ --label "{params.baseline_label}" \ @@ -98,7 +98,7 @@ rule verification_metrics: export ECCODES_DEFINITION_PATH=$(realpath .venv/share/eccodes-cosmo-resources/definitions) uv run {input.script} \ --forecast {params.grib_out_dir} \ - --truth {params.truth} \ + --truth "{params.truth}" \ --reftime {wildcards.init_time} \ --steps "{params.fcst_steps}" \ --label "{params.fcst_label}" \ @@ -226,7 +226,7 @@ rule verification_scoremaps: uv run {input.script} \ --run_root {params.run_root} \ --reftimes {params.reftimes} \ - --truth {input.truth} \ + --truth "{input.truth}" \ --step {wildcards.leadtime} \ --steps "{params.fcst_steps}" \ --param {wildcards.param} \ @@ -262,7 +262,7 @@ rule verification_scoremaps_baseline: uv run {input.script} \ --baseline_root {input.forecast} \ --reftimes {params.reftimes} \ - --truth {input.truth} \ + --truth "{input.truth}" \ --step {wildcards.leadtime} \ --steps "{params.baseline_steps}" \ --param {wildcards.param} \ diff --git a/workflow/scripts/inference_prepare.py b/workflow/scripts/inference_prepare.py index 37e182ae..1b6a7d62 100644 --- a/workflow/scripts/inference_prepare.py +++ b/workflow/scripts/inference_prepare.py @@ -8,7 +8,12 @@ from evalml.helpers import setup_logger -def prepare_config(default_config_path: str, output_config_path: str, params: dict): +def prepare_config( + default_config_path: str, + output_config_path: str, + params: dict, + cross_validation_cfg: dict | None = None, +): """Prepare the configuration file for the inference run. Overrides default configuration parameters with those provided in params @@ -22,12 +27,21 @@ def prepare_config(default_config_path: str, output_config_path: str, params: di Path where the updated configuration file will be written. params : dict Dictionary of parameters to override in the default configuration. + cross_validation_cfg : dict, optional + The experiment's ``experiment.cross_validation`` settings (holdout + station selection for verification stratification). When given, its + ``exclude_stations``/``holdout_fraction``/``holdout_seed`` are + injected into every ``nudge_toward_observation`` filter found in the + config — see ``_inject_nudging_cross_validation``. This is the single + source of truth for the holdout station set: the inference config + itself should not hand-maintain its own copy. """ with open(default_config_path, "r") as f: config = yaml.safe_load(f) config = _override_recursive(config, params) + _inject_nudging_cross_validation(config, cross_validation_cfg or {}) with open(output_config_path, "w") as f: yaml.safe_dump(config, f, sort_keys=False) @@ -94,7 +108,13 @@ def prepare_temporal_downscaler(smk): # prepare config overrides = _overrides_from_params(smk) - prepare_config(smk.input.config, smk.output.config, overrides) + cross_validation_cfg = getattr(smk.params, "cross_validation_cfg", None) + prepare_config( + smk.input.config, + smk.output.config, + overrides, + cross_validation_cfg=cross_validation_cfg, + ) LOG.info("Wrote config file at %s", smk.output.config) with open(smk.output.config, "r") as f: @@ -166,6 +186,50 @@ def _override_recursive(original: dict, updates: dict) -> dict: return original +def _inject_nudging_cross_validation(config: dict, cross_validation_cfg: dict) -> None: + """Set exclude_stations/holdout_fraction/holdout_seed on every + nudge_toward_observation filter found anywhere in config, in place, from + the experiment's cross_validation settings — the single source of truth + for the holdout station set (see varda-single-1.0-nudge.yaml's + experiment.cross_validation). The inference config itself should not + hand-maintain its own copy of the list. + + A no-op if cross_validation_cfg has neither exclude_stations nor + holdout_fraction set (e.g. cross-validation isn't configured for this + experiment) — any hand-written value already in the config is then left + untouched. Unlike _override_recursive (dict-into-dict only), this walks + into lists too, since nudge_toward_observation typically sits inside a + pre_processors list. + """ + exclude_stations = cross_validation_cfg.get("exclude_stations") + holdout_fraction = cross_validation_cfg.get("holdout_fraction") + if exclude_stations is None and holdout_fraction is None: + return + holdout_seed = cross_validation_cfg.get("holdout_seed", 42) + + def _walk(node): + if isinstance(node, dict): + block = node.get("nudge_toward_observation") + if isinstance(block, dict): + # Mutually exclusive in NudgeTowardObservation itself — clear both + # before setting the one cross_validation_cfg actually specifies, + # so a stale hand-written value of the other can never linger. + block.pop("exclude_stations", None) + block.pop("holdout_fraction", None) + if exclude_stations is not None: + block["exclude_stations"] = list(exclude_stations) + else: + block["holdout_fraction"] = holdout_fraction + block["holdout_seed"] = holdout_seed + for value in node.values(): + _walk(value) + elif isinstance(node, list): + for item in node: + _walk(item) + + _walk(config) + + def main(smk): """Main function to run the Snakemake workflow.""" if smk.rule == "inference_prepare_forecaster": diff --git a/workflow/scripts/plot_meteogram.py b/workflow/scripts/plot_meteogram.py index 43248d73..3ea0e183 100644 --- a/workflow/scripts/plot_meteogram.py +++ b/workflow/scripts/plot_meteogram.py @@ -144,7 +144,7 @@ def main(): # Load station metadata from DWH LOG.info("Fetching station metadata from jretrieve (SwissMetNet catalog)") - _jr_stations, _jr_stage, _jr_seq_type = jr.parse_selection("jretrievedwh:1,2") + _jr_stations, _jr_stage, _jr_seq_type, _, _, _ = jr.parse_selection("jretrievedwh:1,2") _catalog = jr.StationCatalog.from_meta( jr.fetch_meta( stations=_jr_stations, From 3e8feb87c20f3c77b38dc5a2da5a96e325bda33e Mon Sep 17 00:00:00 2001 From: lclanzi Date: Mon, 7 Sep 2026 17:18:07 +0200 Subject: [PATCH 14/26] fix linting --- ...er-global_trimedge_multi_with_nudging.yaml | 2 +- src/data_input/__init__.py | 28 ++++++++++++++----- src/data_input/jretrieve.py | 4 +-- src/verification/__init__.py | 17 +++++++---- workflow/scripts/inference_prepare.py | 1 - workflow/scripts/verification_metrics.py | 8 ++++-- workflow/scripts/verification_plot_metrics.py | 5 +++- 7 files changed, 45 insertions(+), 20 deletions(-) diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml index 6b3f4e97..0dea09db 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml @@ -52,7 +52,7 @@ input: lim_effective: 0.0 use_reliability_check: true number_of_std: 5.0 - reliability_min_dist_frac: 0.0 + reliability_min_dist_frac: 0.0 enable_plotting: true plot_dir: output/nudging_reliability_diagnostics run_mode: "devt" diff --git a/src/data_input/__init__.py b/src/data_input/__init__.py index 0bed5b7d..ab01d180 100644 --- a/src/data_input/__init__.py +++ b/src/data_input/__init__.py @@ -671,7 +671,9 @@ def _jretrieve_df_to_xarray(df, short_names, catalog) -> xr.Dataset: return xr.Dataset(data_vars=data_vars, coords=coords) -def _trim_stations_xr(ds: xr.Dataset, filter_mode: str, domain_bbox: list | None) -> xr.Dataset: +def _trim_stations_xr( + ds: xr.Dataset, filter_mode: str, domain_bbox: list | None +) -> xr.Dataset: """Trim the "values" (station) dimension to *filter_mode* — same logic as RetrieveObservation._trim_stations / notebooks/d_eff_generator.ipynb's station-trim cell, adapted for an xarray Dataset with latitude/longitude @@ -692,7 +694,8 @@ def _trim_stations_xr(ds: xr.Dataset, filter_mode: str, domain_bbox: list | None resolution="10m", category="cultural", name="admin_0_countries" ) ch_country = next( - r for r in shpreader.Reader(shp_path).records() + r + for r in shpreader.Reader(shp_path).records() if r.attributes["ADM0_A3"] == "CHE" ) swiss_geom = ch_country.geometry @@ -712,7 +715,9 @@ def _in_switzerland(lat_, lon_): n_before = ds.sizes["values"] ds = ds.isel(values=mask) - LOG.info("Station filter [%s]: %d -> %d stations", desc, n_before, ds.sizes["values"]) + LOG.info( + "Station filter [%s]: %d -> %d stations", desc, n_before, ds.sizes["values"] + ) return ds @@ -751,7 +756,9 @@ def load_obs_data_from_jretrieve( from data_input import jretrieve as jr - stations, stage, seq_type, use_limitation, filter_mode, domain_bbox = jr.parse_selection(root) + stations, stage, seq_type, use_limitation, filter_mode, domain_bbox = ( + jr.parse_selection(root) + ) jr.check_prerequisites(stage) want_uv = "U_10M" in params or "V_10M" in params @@ -812,14 +819,21 @@ def load_obs_data_from_jretrieve( # Same per-variable station-coverage log as RetrieveObservation, so ground-truth # coverage can be compared directly against what was actually available to nudge. _icon_to_short = { - "T_2M": "2t", "TD_2M": "2d", "U_10M": "10u", "V_10M": "10v", - "PMSL": "msl", "TOT_PREC": "tp", "VMAX_10M": "vmax", + "T_2M": "2t", + "TD_2M": "2d", + "U_10M": "10u", + "V_10M": "10v", + "PMSL": "msl", + "TOT_PREC": "tp", + "VMAX_10M": "vmax", } n_total = result.sizes["values"] for icon, short in _icon_to_short.items(): if icon in result.data_vars: n_valid = int(result[icon].notnull().any("time").sum()) - LOG.info("Stations with valid %s: %d / %d stations", short, n_valid, n_total) + LOG.info( + "Stations with valid %s: %d / %d stations", short, n_valid, n_total + ) return result diff --git a/src/data_input/jretrieve.py b/src/data_input/jretrieve.py index 5b69183e..9efa736b 100644 --- a/src/data_input/jretrieve.py +++ b/src/data_input/jretrieve.py @@ -222,9 +222,7 @@ def parse_selection( elif key == "domain_bbox": domain_bbox = [float(v) for v in value.split(",") if v] if len(domain_bbox) != 4: - raise ValueError( - "domain_bbox must be lat_min,lat_max,lon_min,lon_max." - ) + raise ValueError("domain_bbox must be lat_min,lat_max,lon_min,lon_max.") else: raise ValueError(f"Unknown jretrieve selector key: {key!r}") if not stations: diff --git a/src/verification/__init__.py b/src/verification/__init__.py index c8d27176..c01cf698 100644 --- a/src/verification/__init__.py +++ b/src/verification/__init__.py @@ -306,7 +306,10 @@ def _create_station_group_masks( is_holdout = np.isin(all_nat_abbr, holdout_stations) return xr.DataArray( np.stack([np.ones_like(is_holdout), is_holdout, ~is_holdout]), - coords={"station_group": ["all", "holdout", "holdin"], "values": values_coord.values}, + coords={ + "station_group": ["all", "holdout", "holdin"], + "values": values_coord.values, + }, dims=["station_group", "values"], ) @@ -402,10 +405,14 @@ def verify( station_masks = None if holdout_stations is not None: - station_masks = _create_station_group_masks(obs_aligned["values"], holdout_stations) - LOG.info("Station group masks created: %d holdout, %d holdin stations", - int(station_masks.sel(station_group="holdout").sum()), - int(station_masks.sel(station_group="holdin").sum())) + station_masks = _create_station_group_masks( + obs_aligned["values"], holdout_stations + ) + LOG.info( + "Station group masks created: %d holdout, %d holdin stations", + int(station_masks.sel(station_group="holdout").sum()), + int(station_masks.sel(station_group="holdin").sum()), + ) scores = [] statistics = [] diff --git a/workflow/scripts/inference_prepare.py b/workflow/scripts/inference_prepare.py index d6e0e7d4..f23e97c7 100644 --- a/workflow/scripts/inference_prepare.py +++ b/workflow/scripts/inference_prepare.py @@ -47,7 +47,6 @@ def prepare_config( yaml.safe_dump(config, f, sort_keys=False) - def prepare_workdir(workdir: Path, resources_root: Path): """Prepare the working directory for the inference run. diff --git a/workflow/scripts/verification_metrics.py b/workflow/scripts/verification_metrics.py index 6ada2e75..a5b58fc0 100644 --- a/workflow/scripts/verification_metrics.py +++ b/workflow/scripts/verification_metrics.py @@ -123,9 +123,13 @@ def main(args: ScriptConfig): len(all_stations), ) else: - LOG.warning("cross_validation_cfg set but no holdout stations selected (check holdout_fraction / exclude_stations).") + LOG.warning( + "cross_validation_cfg set but no holdout stations selected (check holdout_fraction / exclude_stations)." + ) elif cv_cfg: - LOG.warning("cross_validation_cfg set but truth dataset has no 'values' dimension; station stratification skipped.") + LOG.warning( + "cross_validation_cfg set but truth dataset has no 'values' dimension; station stratification skipped." + ) # compute metrics and statistics now = datetime.now() diff --git a/workflow/scripts/verification_plot_metrics.py b/workflow/scripts/verification_plot_metrics.py index fed32f22..39995076 100644 --- a/workflow/scripts/verification_plot_metrics.py +++ b/workflow/scripts/verification_plot_metrics.py @@ -94,7 +94,10 @@ def main(args: Namespace) -> None: # remove duplicated but not identical values from analyses (rounding errors) dfs = [xr.open_dataset(f) for f in args.verif_files] # When cross_validation is enabled, runs carry a station_group dim; select "all" for standard plots. - dfs = [d.sel(station_group="all", drop=True) if "station_group" in d.dims else d for d in dfs] + dfs = [ + d.sel(station_group="all", drop=True) if "station_group" in d.dims else d + for d in dfs + ] # 1) Ensure each dataset has unique lead_time values dfs = [_ensure_unique_lead_time(d) for d in dfs] # 2) For sources present in multiple datasets, keep the one with most lead_times From 65e611fe038e204562f824831d3dfce1908829c0 Mon Sep 17 00:00:00 2001 From: lclanzi Date: Mon, 7 Sep 2026 17:24:40 +0200 Subject: [PATCH 15/26] restore varda-single config --- config/varda-single-1.0.yaml | 46 +++++++++---------- ...tidataset-forecaster-global-ich1-oper.yaml | 24 +--------- 2 files changed, 23 insertions(+), 47 deletions(-) diff --git a/config/varda-single-1.0.yaml b/config/varda-single-1.0.yaml index e9bc416d..f0c614b3 100644 --- a/config/varda-single-1.0.yaml +++ b/config/varda-single-1.0.yaml @@ -5,16 +5,15 @@ description: | config_label: varda-single-1.0 dates: - # start: 2025-03-01T00:00 - # end: 2025-03-03T00:00 - # frequency: 24h - - 2025-03-01T00:00 + start: 2025-03-01T00:00 + end: 2025-03-03T00:00 + frequency: 24h runs: - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 label: Varda-single-1.0 - steps: 0/24/1 + steps: 0/120/1 config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml extra_requirements: - anemoi-datasets==0.5.35 @@ -25,23 +24,19 @@ runs: forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml - steps: 0/24/6 - extra_requirements: - - anemoi-datasets==0.5.35 - # - git+https://github.com/ecmwf/anemoi-inference.git@eaeb36fae30c754a03b8dc0442fc56373b6bedef - - -e /scratch/mch/llanzila/sruc/anemoi-inference - # - baseline: - # label: INCA - # root: /store_new/mch/msclim/INCA - # steps: 0/6/1 - # - baseline: - # label: ICON-CH2-CTRL - # root: /store_new/mch/msopr/osm/ICON-CH2-EPS - # steps: 0/120/1 + steps: 0/120/6 + - baseline: + label: INCA + root: /store_new/mch/msclim/INCA + steps: 0/6/1 + - baseline: + label: ICON-CH2-CTRL + root: /store_new/mch/msopr/osm/ICON-CH2-EPS + steps: 0/120/1 - baseline: label: ICON-CH1-CTRL root: /store_new/mch/msopr/osm/ICON-CH1-EPS - steps: 0/24/1 + steps: 0/33/1 truth: label: SwissMetNet @@ -85,7 +80,7 @@ experiment: # - init_hour - season scorecards: - enabled: false + enabled: true sections: nowcasting: baseline: INCA @@ -98,7 +93,7 @@ experiment: - "TOT_PREC1:RMSE,R2,ETS" short_range: baseline: ICON-CH1-CTRL - lead_times: "6/24/6" + lead_times: "6/33/6" stratification: region variables: - "U_10M:RMSE,R2,ETS" @@ -127,11 +122,11 @@ showcase: stations: [JUN] #, COV, GOR, WFJ, SAE, SAM, DAV, ZER, ANT, VSBAS, BRT, LTB, GOS, CEV, BIA] animations: enabled: true - frames_per_second: 0.5 + frames_per_second: 5 domains: # - globe - # - europe - - alps + - europe + # - alps - icon-ch - switzerland @@ -148,4 +143,5 @@ profile: mem_mb_per_cpu: 1800 runtime: "1h" gpus: 0 - jobs: 50 + slurm_account: s83 + jobs: 50 \ No newline at end of file diff --git a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml index 9a952d8f..dd5f0e5d 100644 --- a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml +++ b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml @@ -1,4 +1,4 @@ -lead_time: 24h +lead_time: 120h write_initial_state: true allow_nans: true @@ -10,26 +10,6 @@ input: test: use_original_paths: true -# --- Nudging (disabled by default) --- -# Uncomment to nudge the initial condition toward station observations. -# backend: 'peakweather' reads from a local dataset; 'jretrieve' fetches live from DWH. -# nudge_variables: list of GRIB shortNames to nudge; omit to nudge all (T_2M, U_10M, V_10M, TOT_PREC). -# use_limitation: int (optional, jretrieve only) — passed as --use-limitation; limits obs to -# within ±N minutes of the target time (e.g. 50). -# -# pre_processors: -# - forward_transform_filter: -# nudge_toward_observation: -# backend: jretrieve -# nudge_variables: [T_2M] -# icon_grid_dir: /scratch/mch/llanzila/sruc/aux_files -# jretrieve_bbox: [45.7, 48.0, 5.8, 10.8] -# jretrieve_src_path: /scratch/mch/llanzila/sruc/evalml/src -# k: 3 -# power: 4.0 -# max_dist: 0.5 -# use_limitation: 50 - post_processors: - accumulate_from_start_of_forecast: # accumulate tp from start of forecast accumulations: @@ -98,4 +78,4 @@ output: "^w_(\\d+)$": {"param": 135, "shortName": "w"} "^z_(\\d+)$": {"param": 129, "shortName": "z"} -patch_metadata: resources/sgm-multidataset-ich1-oper-patch.yaml +patch_metadata: resources/sgm-multidataset-ich1-oper-patch.yaml \ No newline at end of file From 133e3cf7784391e48cc2af9e3f970feb8423b321 Mon Sep 17 00:00:00 2001 From: lclanzi Date: Mon, 7 Sep 2026 17:26:13 +0200 Subject: [PATCH 16/26] fix linting --- config/varda-single-1.0.yaml | 2 +- .../configs/sgm-multidataset-forecaster-global-ich1-oper.yaml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/config/varda-single-1.0.yaml b/config/varda-single-1.0.yaml index f0c614b3..c5875cc7 100644 --- a/config/varda-single-1.0.yaml +++ b/config/varda-single-1.0.yaml @@ -144,4 +144,4 @@ profile: runtime: "1h" gpus: 0 slurm_account: s83 - jobs: 50 \ No newline at end of file + jobs: 50 diff --git a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml index dd5f0e5d..cd0c95a5 100644 --- a/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml +++ b/resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml @@ -78,4 +78,4 @@ output: "^w_(\\d+)$": {"param": 135, "shortName": "w"} "^z_(\\d+)$": {"param": 129, "shortName": "z"} -patch_metadata: resources/sgm-multidataset-ich1-oper-patch.yaml \ No newline at end of file +patch_metadata: resources/sgm-multidataset-ich1-oper-patch.yaml From 5955d8a0533b1fa187862a5b1c76cfb980f2fa37 Mon Sep 17 00:00:00 2001 From: lclanzi Date: Thu, 10 Sep 2026 14:19:57 +0200 Subject: [PATCH 17/26] update configs --- config/varda-rapid-0.2.yaml | 18 +++++++++++------- ...ler-global_trimedge_multi_with_nudging.yaml | 14 ++++++++++++++ 2 files changed, 25 insertions(+), 7 deletions(-) diff --git a/config/varda-rapid-0.2.yaml b/config/varda-rapid-0.2.yaml index 363ae471..63a92fc9 100644 --- a/config/varda-rapid-0.2.yaml +++ b/config/varda-rapid-0.2.yaml @@ -2,7 +2,7 @@ description: | Evaluate skill of Varda-rapid against ground observations. -config_label: varda-single-1.0-nudge +config_label: varda-rapid-0.2 dates: start: 2025-01-01T06:00 @@ -21,14 +21,15 @@ runs: config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml extra_requirements: - anemoi-datasets==0.5.35 - - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@6e6b2dd42e66d9c981bd8f5ccca657f0411e4ccc - - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f62ae9f3eae7bb2d19837d220b10a7bddf35a57c + - git+https://github.com/ecmwf/anemoi-inference@88c1ec13632c570f246cc2580ce324e45f36a03b forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml steps: 0/12/6 extra_requirements: - - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f62ae9f3eae7bb2d19837d220b10a7bddf35a57c + - git+https://github.com/ecmwf/anemoi-inference@88c1ec13632c570f246cc2580ce324e45f36a03b - temporal_downscaler: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-2-interpolator/versions/3 @@ -37,13 +38,15 @@ runs: config: resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi.yaml extra_requirements: - anemoi-datasets==0.5.35 - - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f62ae9f3eae7bb2d19837d220b10a7bddf35a57c + - git+https://github.com/ecmwf/anemoi-inference@88c1ec13632c570f246cc2580ce324e45f36a03b forecaster: checkpoint: https://service.meteoswiss.ch/mlstore#/models/sruc-m-1-forecaster/versions/4 config: resources/inference/configs/sgm-multidataset-forecaster-global-ich1-oper.yaml steps: 0/12/6 extra_requirements: - - git+https://github.com/ecmwf/anemoi-inference@eaeb36fae30c754a03b8dc0442fc56373b6bedef + - git+https://github.com/MeteoSwiss/anemoi-plugins-meteoswiss@f62ae9f3eae7bb2d19837d220b10a7bddf35a57c + - git+https://github.com/ecmwf/anemoi-inference@88c1ec13632c570f246cc2580ce324e45f36a03b - baseline: label: INCA @@ -154,9 +157,10 @@ profile: global_resources: gpus: 16 default_resources: - slurm_partition: "postproc" + slurm_partition: "normal" cpus_per_task: 1 mem_mb_per_cpu: 1800 runtime: "1h" gpus: 0 + slurm_account: s83 jobs: 50 diff --git a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml index 0dea09db..ee7b00e6 100644 --- a/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml +++ b/resources/inference/configs/sgm-temporal-downscaler-global_trimedge_multi_with_nudging.yaml @@ -124,6 +124,14 @@ output: mask: "lam_0/cutout_mask" as_slice: true # here, the trimedge mask can be specified when available + # at overlap steps (multiples of the forecaster stride) replace the + # re-predicted prognostics with the forecaster's boundary values + - forward_transform_filter: + copy-prognostic-from-forecaster: + forecaster_path: forecaster/20* + common_leadtime: 6h + params_to_keep: [tp] # NOTE, if more diagnostics are added, update accordingly. + namer: *namer - grib: path: grib/ifs-{dateTime}_{step:03}.grib @@ -138,6 +146,12 @@ output: inverse: true - assign_mask: # fill local/global overlapping points with nan mask: "global/cutout_mask" + - forward_transform_filter: + copy-prognostic-from-forecaster: + forecaster_path: forecaster/ifs* + common_leadtime: 6h + params_to_keep: [tp] # NOTE, if more diagnostics are added, update accordingly. + namer: *namer modifiers: - patches: - variable: From e97998b908f2549311b24a712ee1ce92bb32bd5c Mon Sep 17 00:00:00 2001 From: lclanzi Date: Fri, 11 Sep 2026 11:29:12 +0200 Subject: [PATCH 18/26] remove filtering from jretrieve and insert a new region (switzerland) for stratification --- config/varda-rapid-0.2.yaml | 66 +++++-------- .../report/dashboard/template.html.jinja2 | 2 +- src/evalml/config.py | 18 +++- workflow/rules/common.smk | 8 +- workflow/rules/inference.smk | 38 +++++--- workflow/rules/report.smk | 4 +- workflow/rules/verification.smk | 8 +- workflow/scripts/inference_prepare.py | 32 +++--- .../scripts/report_experiment_dashboard.py | 10 +- workflow/scripts/report_scorecard.py | 2 +- workflow/scripts/verification_metrics.py | 28 +++--- workflow/scripts/verification_plot_metrics.py | 2 +- workflow/tools/config.schema.json | 97 ++++++++++--------- 13 files changed, 162 insertions(+), 153 deletions(-) diff --git a/config/varda-rapid-0.2.yaml b/config/varda-rapid-0.2.yaml index 63a92fc9..ea90be7b 100644 --- a/config/varda-rapid-0.2.yaml +++ b/config/varda-rapid-0.2.yaml @@ -55,12 +55,7 @@ runs: truth: label: SwissMetNet - # Retrieve over the broad ICON domain (matching notebooks/d_eff_generator.ipynb's - # STATIONS_SEL), then trim to stations within the real Swiss national border - # (filter_mode=switzerland) or within a predefined bounding box (filter_mode=domain;domain_bbox=45.7,48.0,5.8,10.8). - root: jretrievedwh:bbox=40.5,53.0,0.0,17.5;use_limitation=40;filter_mode=switzerland - # root: jretrievedwh:bbox=40.5,53.0,0.0,17.5;use_limitation=40;filter_mode=domain;domain_bbox=45.7,48.0,5.8,10.8 - + root: jretrievedwh:bbox=45.7,48.0,5.8,10.8;use_limitation=40 experiment: params: @@ -73,6 +68,7 @@ experiment: - PS stratification: regions: + - switzerland - mittelland - berge - alpennordseite @@ -92,41 +88,29 @@ experiment: dashboard: stratification: - season - cross_validation: - # Stations withheld from nudging in the nudge-partial inference run. - exclude_stations: - - CHM - - FRE - - RUE - - BAS - - PLF - - CGI - - HOE - - PUY - - RAG - - EIN - - ALT - - GOS - - COV - - ULR - - BIN - - GUE - - ROE - - BIA - - CEV - - OTL - enabled: true - sections: - nowcasting: - baseline: INCA - lead_times: "0/6/1" - stratification: region - variables: - - "U_10M:RMSE,R2,ETS" - - "V_10M:RMSE,R2,ETS" - - "T_2M:RMSE,R2,ETS" - - "TOT_PREC:RMSE,R2,ETS" - - "TD_2M:RMSE,R2,ETS" + - region + # Stations withheld from nudging in the nudge-partial inference run. + station_holdout: + - CHM + - FRE + - RUE + - BAS + - PLF + - CGI + - HOE + - PUY + - RAG + - EIN + - ALT + - GOS + - COV + - ULR + - BIN + - GUE + - ROE + - BIA + - CEV + - OTL scoremaps: enabled: false diff --git a/resources/report/dashboard/template.html.jinja2 b/resources/report/dashboard/template.html.jinja2 index 0cf2901a..0da748f0 100644 --- a/resources/report/dashboard/template.html.jinja2 +++ b/resources/report/dashboard/template.html.jinja2 @@ -154,7 +154,7 @@
{% endif %} - {% if cross_validation %} + {% if station_holdout %}