was-disaggregation 0.8.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,23 @@
1
+ """Forecast-conditioned stochastic weather generation on xarray grids."""
2
+ from ._version import __version__
3
+ from .data import open_observations, load_probabilities, seasonal_cube, season_dates
4
+ from .model import WeatherGenerator
5
+ from .scalable import generate_dask
6
+ from .nonparametric import ForecastAnalogGenerator, ForecastSchaakeGenerator, schaake_shuffle, enso_rank_weights
7
+ from .attributes import SeasonalTotal, OnsetDate, CessationDate, MaxDrySpell
8
+ from .mre import SeasonalConstraint, constrained_year_weights
9
+ from .diagnostics import attribute_diagnostics, select_members
10
+ from .glm import GLMWeatherGenerator, CoefficientGP
11
+ from .dynamical import (DailyBiasCorrector, ExternalCorrector, NHMM, NHSMM, state_durations, nhmm_predictors, ensemble_copula_coupling,
12
+ preferential_dates, DynamicalDownscaler, synthetic_model_ensemble)
13
+ from .validation import (HindcastExperiment, model_tercile_probabilities, ensemble_normal_scores, rank_histogram,
14
+ reliability_table, daily_scores)
15
+ from .generative import GenerativeDownscaler # needs PyTorch only when fitted
16
+ __all__ = ["WeatherGenerator", "ForecastAnalogGenerator", "ForecastSchaakeGenerator", "generate_dask",
17
+ "open_observations", "load_probabilities", "seasonal_cube", "season_dates", "schaake_shuffle",
18
+ "enso_rank_weights", "SeasonalTotal", "OnsetDate", "CessationDate", "MaxDrySpell",
19
+ "SeasonalConstraint", "constrained_year_weights", "attribute_diagnostics", "select_members",
20
+ "GLMWeatherGenerator", "CoefficientGP", "DailyBiasCorrector", "ExternalCorrector", "NHMM", "NHSMM", "state_durations", "nhmm_predictors",
21
+ "ensemble_copula_coupling", "preferential_dates", "DynamicalDownscaler", "synthetic_model_ensemble",
22
+ "HindcastExperiment", "model_tercile_probabilities", "ensemble_normal_scores",
23
+ "rank_histogram", "reliability_table", "daily_scores", "GenerativeDownscaler"]
@@ -0,0 +1 @@
1
+ __version__ = "0.8.0"
@@ -0,0 +1,358 @@
1
+ """Agro-climatic season attributes computed from daily rainfall.
2
+
3
+ Each attribute turns daily PRCP into one number per season and site: seasonal
4
+ total, onset date, cessation date or maximum dry-spell length. These are the
5
+ quantities forecast by PRESASS / AGRHYMET, and each one can become a
6
+ constraint on the historical year weights (see :mod:`was_disaggregation.mre`).
7
+
8
+ Windows are given as ``("MM-DD", "MM-DD")`` in the calendar year of the season
9
+ (``end < start`` crosses into the next year). February 29 is omitted, matching
10
+ the generators' 365-day Gregorian policy. Date attributes are returned as
11
+ **days after the start of their search window** (0 = search start), which is
12
+ counted on this calendar and orders early -> late. A season whose event is not found
13
+ inside the search window is *censored* at the window length (i.e. classed as
14
+ late); ``censored_`` reports how often this happened.
15
+
16
+ Any missing day inside an attribute's data span makes that season's value NaN
17
+ (missing is never treated as dry).
18
+
19
+ The functions accept a DataArray with a ``T`` dimension and any other
20
+ dimensions (Y, X, member ...), stacked internally into ``site``.
21
+ """
22
+ from __future__ import annotations
23
+
24
+ from dataclasses import dataclass
25
+
26
+ import numpy as np
27
+ import pandas as pd
28
+ import xarray as xr
29
+
30
+
31
+ def _md(text):
32
+ month, day = (int(v) for v in str(text).split("-"))
33
+ return month, day
34
+
35
+
36
+ def _window(year, start, end):
37
+ sm, sd = _md(start)
38
+ em, ed = _md(end)
39
+ first = pd.Timestamp(int(year), sm, sd)
40
+ last = pd.Timestamp(int(year) + (1 if (em, ed) < (sm, sd) else 0), em, ed)
41
+ return first, last
42
+
43
+
44
+ def _window_dates(year, start, end, pad_after=0):
45
+ """Window plus look-ahead on the package's February-29-excluded calendar."""
46
+ first, last = _window(year, start, end)
47
+ dates = pd.date_range(first, last, freq="D")
48
+ dates = dates[~((dates.month == 2) & (dates.day == 29))]
49
+ extra, current = [], last
50
+ while len(extra) < pad_after:
51
+ current += pd.Timedelta(days=1)
52
+ if (current.month, current.day) != (2, 29):
53
+ extra.append(current)
54
+ return dates.append(pd.DatetimeIndex(extra)) if extra else dates
55
+
56
+
57
+ def _window_length(year, start, end):
58
+ return len(_window_dates(year, start, end))
59
+
60
+
61
+ def _stack(prcp):
62
+ """DataArray (T, ...) -> (DataArray (T, site), template of the other dims)."""
63
+ if "T" not in prcp.dims:
64
+ raise ValueError("daily rainfall must have a T dimension")
65
+ if prcp.sizes["T"] == 0:
66
+ raise ValueError("daily rainfall must have a nonempty T dimension")
67
+ other = [d for d in prcp.dims if d != "T"]
68
+ template = prcp.isel(T=0, drop=True)
69
+ if other:
70
+ stacked = prcp.transpose("T", *other).stack(site=other)
71
+ else:
72
+ stacked = prcp.expand_dims(site=[0], axis=1)
73
+ return stacked, template
74
+
75
+
76
+ def extract_windows(prcp, years, start, end, pad_after=0):
77
+ """Daily values for each season window -> (year, day, site), NaN padded.
78
+
79
+ ``pad_after`` extra days are appended after ``end`` (look-ahead for
80
+ false-start checks or post-onset dry-spell windows).
81
+ """
82
+ stacked, _ = _stack(prcp)
83
+ years = list(years)
84
+ if not years or not isinstance(pad_after, (int, np.integer)) or pad_after < 0:
85
+ raise ValueError("years must be nonempty and pad_after a nonnegative integer")
86
+ if not np.issubdtype(stacked["T"].dtype, np.datetime64):
87
+ raise ValueError("T must use Gregorian datetime64 dates")
88
+ times = pd.DatetimeIndex(stacked["T"].values).normalize()
89
+ if times.hasnans or times.has_duplicates:
90
+ raise ValueError("T must contain unique daily dates without NaT")
91
+ stacked = stacked.assign_coords(T=times)
92
+ spans = [_window_dates(y, start, end, pad_after) for y in years]
93
+ lengths = [len(dates) for dates in spans]
94
+ n = max(lengths)
95
+ out = np.full((len(years), n, stacked.sizes["site"]), np.nan)
96
+ for i, dates in enumerate(spans):
97
+ values = stacked.reindex(T=dates).values
98
+ out[i, :lengths[i]] = values
99
+ return out
100
+
101
+
102
+ def _dry_run_ending(wet):
103
+ """Length of the dry run ending at each day (0 on wet days), axis 1 = day."""
104
+ run = np.zeros(wet.shape, dtype=np.int32)
105
+ current = np.zeros(wet.shape[:1] + wet.shape[2:], dtype=np.int32)
106
+ for d in range(wet.shape[1]):
107
+ current = np.where(wet[:, d], 0, current + 1)
108
+ run[:, d] = current
109
+ return run
110
+
111
+
112
+ def _reshape(values, template, leading):
113
+ """(year, site) -> DataArray (leading, *template.dims)."""
114
+ name, coord = leading
115
+ shape = (values.shape[0],) + template.shape
116
+ return xr.DataArray(values.reshape(shape), dims=(name,) + template.dims,
117
+ coords={name: coord, **{d: template[d] for d in template.dims if d in template.coords}})
118
+
119
+
120
+ class Attribute:
121
+ """Base class. Subclasses implement ``values(prcp, years) -> (year, site)``."""
122
+ name = "attribute"
123
+ units = ""
124
+
125
+ def compute(self, prcp: xr.DataArray, years) -> xr.DataArray:
126
+ """Attribute per season year, as DataArray (season_year, <other dims>)."""
127
+ years = [int(y) for y in np.atleast_1d(years)]
128
+ _, template = _stack(prcp)
129
+ values = self.values(prcp, years)
130
+ out = _reshape(values, template, ("season_year", years))
131
+ out.name = self.name
132
+ out.attrs.update(units=self.units, definition=repr(self), calendar_policy="Gregorian; February 29 omitted")
133
+ return out
134
+
135
+ def span(self):
136
+ """(start 'MM-DD', end 'MM-DD', extra days after end) of data needed."""
137
+ raise NotImplementedError
138
+
139
+ def values(self, prcp, years):
140
+ raise NotImplementedError
141
+
142
+
143
+ @dataclass
144
+ class SeasonalTotal(Attribute):
145
+ """Rainfall total over ``window`` (mm). Requires every day of the window."""
146
+ window: tuple = ("07-01", "09-30")
147
+ name: str = "total"
148
+ units: str = "mm"
149
+
150
+ def span(self):
151
+ return self.window[0], self.window[1], 0
152
+
153
+ def values(self, prcp, years):
154
+ x = extract_windows(prcp, years, *self.window)
155
+ valid_len = np.array([_window_length(y, *self.window) for y in years])
156
+ out = np.full((len(years), x.shape[2]), np.nan)
157
+ for i, n in enumerate(valid_len):
158
+ block = x[i, :n]
159
+ out[i] = np.where(np.isfinite(block).all(axis=0), np.nansum(block, axis=0), np.nan)
160
+ return out
161
+
162
+
163
+ @dataclass
164
+ class OnsetDate(Attribute):
165
+ """Agronomic onset (Sivakumar 1988, as used by AGRHYMET / PRESASS).
166
+
167
+ First day ``d`` of the search window such that
168
+ * rain accumulated over days d .. d+accumulation_days-1 >= accumulation_mm,
169
+ * at least ``min_wet_days`` of those days are wet (>= wet_threshold),
170
+ * no dry spell longer than ``dry_spell_days`` occurs in the following
171
+ ``check_days`` days (false-start check).
172
+ Returned as days after ``search[0]``; censored at the window length.
173
+ """
174
+ search: tuple = ("05-01", "09-30")
175
+ accumulation_mm: float = 20.0
176
+ accumulation_days: int = 3
177
+ min_wet_days: int = 1
178
+ dry_spell_days: int = 7
179
+ check_days: int = 30
180
+ wet_threshold: float = 1.0
181
+ name: str = "onset"
182
+ units: str = "days after search start"
183
+
184
+ def __post_init__(self):
185
+ _window(2001, *self.search)
186
+ for name in ("accumulation_days", "min_wet_days", "dry_spell_days", "check_days"):
187
+ value = getattr(self, name)
188
+ minimum = 1 if name in ("accumulation_days", "min_wet_days") else 0
189
+ if not isinstance(value, (int, np.integer)) or value < minimum:
190
+ raise ValueError(f"{name} must be an integer >= {minimum}")
191
+ if self.min_wet_days > self.accumulation_days:
192
+ raise ValueError("min_wet_days cannot exceed accumulation_days")
193
+ if not np.isfinite(self.accumulation_mm) or self.accumulation_mm <= 0:
194
+ raise ValueError("accumulation_mm must be positive")
195
+ if not np.isfinite(self.wet_threshold) or self.wet_threshold <= 0:
196
+ raise ValueError("wet_threshold must be positive")
197
+
198
+ def span(self):
199
+ return self.search[0], self.search[1], self.check_days + self.accumulation_days - 1
200
+
201
+ def values(self, prcp, years):
202
+ onset, _ = self._onset(prcp, years)
203
+ return onset
204
+
205
+ def _onset(self, prcp, years, x=None):
206
+ if x is None:
207
+ x = extract_windows(prcp, years, *self.span()[:2], pad_after=self.span()[2])
208
+ ny, nd, ns = x.shape
209
+ length = np.array([_window_length(y, *self.search) for y in years])
210
+ missing = np.zeros((ny, ns), dtype=bool)
211
+ for i, n in enumerate(length):
212
+ missing[i] = ~np.isfinite(x[i, :n + self.span()[2]]).all(axis=0)
213
+ xf = np.nan_to_num(x)
214
+ wet = xf >= self.wet_threshold
215
+ k = self.accumulation_days
216
+ csum = np.concatenate([np.zeros((ny, 1, ns)), np.cumsum(xf, axis=1)], axis=1)
217
+ cwet = np.concatenate([np.zeros((ny, 1, ns)), np.cumsum(wet, axis=1)], axis=1)
218
+ idx = np.arange(nd - k + 1)
219
+ accum = csum[:, idx + k] - csum[:, idx]
220
+ nwet = cwet[:, idx + k] - cwet[:, idx]
221
+ candidate = (accum >= self.accumulation_mm) & (nwet >= self.min_wet_days)
222
+ # longest dry spell strictly after the accumulation window, within check_days
223
+ run = _dry_run_ending(wet)
224
+ longest = np.zeros(candidate.shape, dtype=np.int32)
225
+ for j in range(1, self.check_days + 1):
226
+ pos = idx + k - 1 + j
227
+ ok = pos < nd
228
+ r = np.zeros(candidate.shape, dtype=np.int32)
229
+ r[:, ok] = np.minimum(run[:, pos[ok]], j)
230
+ longest = np.maximum(longest, r)
231
+ good = candidate & (longest <= self.dry_spell_days)
232
+ onset = np.full((ny, ns), np.nan)
233
+ censored = np.zeros((ny, ns), dtype=bool)
234
+ for i, n in enumerate(length):
235
+ g = good[i, :n]
236
+ found = g.any(axis=0)
237
+ first = np.argmax(g, axis=0)
238
+ onset[i] = np.where(found, first, n)
239
+ censored[i] = ~found
240
+ onset[missing] = np.nan
241
+ self.censored_ = float(np.mean(censored[~missing])) if (~missing).any() else np.nan
242
+ return onset, x
243
+
244
+
245
+ @dataclass
246
+ class CessationDate(Attribute):
247
+ """End of season from a bucket water balance (AGRHYMET-style).
248
+
249
+ Soil water S starts at 0 on ``balance_start`` and evolves as
250
+ S = clip(S + P - et_mm, 0, capacity_mm). Cessation is the first day of
251
+ ``search`` on which S = 0 (days after ``search[0]``); censored if never.
252
+ """
253
+ search: tuple = ("09-01", "11-30")
254
+ balance_start: str = "05-01"
255
+ capacity_mm: float = 70.0
256
+ et_mm: float = 5.0
257
+ name: str = "cessation"
258
+ units: str = "days after search start"
259
+
260
+ def __post_init__(self):
261
+ _window(2001, *self.search)
262
+ _window(2001, self.balance_start, self.search[1])
263
+ if not np.isfinite(self.capacity_mm) or self.capacity_mm <= 0:
264
+ raise ValueError("capacity_mm must be positive")
265
+ if not np.isfinite(self.et_mm) or self.et_mm < 0:
266
+ raise ValueError("et_mm must be finite and nonnegative")
267
+
268
+ def span(self):
269
+ return self.balance_start, self.search[1], 0
270
+
271
+ def values(self, prcp, years):
272
+ x = extract_windows(prcp, years, self.balance_start, self.search[1])
273
+ ny, nd, ns = x.shape
274
+ out = np.full((ny, ns), np.nan)
275
+ for i, y in enumerate(years):
276
+ b0, end = _window(y, self.balance_start, self.search[1])
277
+ s0 = _window(y, *self.search)[0]
278
+ if s0 < b0:
279
+ s0 = s0.replace(year=s0.year + 1)
280
+ dates = _window_dates(y, self.balance_start, self.search[1])
281
+ n = len(dates)
282
+ offset = int((dates < s0).sum())
283
+ block = x[i, :n]
284
+ miss = ~np.isfinite(block).all(axis=0)
285
+ soil = np.zeros(ns)
286
+ result = np.full(ns, float(n - offset))
287
+ found = np.zeros(ns, dtype=bool)
288
+ for d in range(n):
289
+ soil = np.clip(soil + np.nan_to_num(block[d]) - self.et_mm, 0, self.capacity_mm)
290
+ if d >= offset:
291
+ hit = (soil <= 0) & ~found
292
+ result[hit] = d - offset
293
+ found |= hit
294
+ out[i] = np.where(miss, np.nan, result)
295
+ return out
296
+
297
+
298
+ @dataclass
299
+ class MaxDrySpell(Attribute):
300
+ """Longest run of dry days (< wet_threshold) inside a window.
301
+
302
+ Fixed window: ``window=("MM-DD","MM-DD")``. Onset-relative window (PRESASS
303
+ early-season dry spells): ``after_onset=OnsetDate(...)`` and ``length_days``
304
+ -> days [onset, onset + length_days). Seasons without onset are censored
305
+ at ``length_days`` (failed season = longest spell).
306
+ Spells are cut at the window edges.
307
+ """
308
+ window: tuple | None = None
309
+ after_onset: OnsetDate | None = None
310
+ length_days: int = 50
311
+ wet_threshold: float = 1.0
312
+ name: str = "dry_spell"
313
+ units: str = "days"
314
+
315
+ def __post_init__(self):
316
+ if (self.window is None) == (self.after_onset is None):
317
+ raise ValueError("give exactly one of window or after_onset")
318
+ if not isinstance(self.length_days, (int, np.integer)) or self.length_days < 1:
319
+ raise ValueError("length_days must be a positive integer")
320
+ if not np.isfinite(self.wet_threshold) or self.wet_threshold <= 0:
321
+ raise ValueError("wet_threshold must be positive")
322
+
323
+ def span(self):
324
+ if self.window is not None:
325
+ return self.window[0], self.window[1], 0
326
+ s, e, pad = self.after_onset.span()
327
+ return s, e, max(pad, self.length_days - 1)
328
+
329
+ def values(self, prcp, years):
330
+ if self.window is not None:
331
+ x = extract_windows(prcp, years, *self.window)
332
+ n = np.array([_window_length(y, *self.window) for y in years])
333
+ start = np.zeros((len(years), x.shape[2]), dtype=int)
334
+ stop = np.broadcast_to(n[:, None], start.shape)
335
+ else:
336
+ s, e, pad = self.span()
337
+ x = extract_windows(prcp, years, s, e, pad_after=pad)
338
+ onset, _ = self.after_onset._onset(prcp, years, x=x)
339
+ n_search = np.array([_window_length(y, *self.after_onset.search) for y in years])
340
+ no_onset = onset >= n_search[:, None]
341
+ start = np.where(np.isfinite(onset), onset, 0).astype(int)
342
+ stop = start + self.length_days
343
+ start = np.where(no_onset, 0, start)
344
+ ny, nd, ns = x.shape
345
+ day = np.arange(nd)[None, :, None]
346
+ inside = (day >= start[:, None, :]) & (day < np.minimum(stop, nd)[:, None, :])
347
+ miss = (~np.isfinite(x) & inside).any(axis=1)
348
+ dry = inside & (np.nan_to_num(x) < self.wet_threshold)
349
+ run = _dry_run_ending(~dry)
350
+ out = np.where(miss, np.nan, run.max(axis=1).astype(float))
351
+ if self.after_onset is not None:
352
+ # No onset in the search window: a failed season, censored as the
353
+ # longest possible spell (keeps such years in the "long" class).
354
+ out = np.where(np.isfinite(onset), np.where(no_onset, float(self.length_days), out), np.nan)
355
+ return out
356
+
357
+
358
+ __all__ = ["Attribute", "SeasonalTotal", "OnsetDate", "CessationDate", "MaxDrySpell", "extract_windows"]
@@ -0,0 +1,106 @@
1
+ """Small operational entry point; notebooks describe the statistical choices."""
2
+ import argparse
3
+ from pathlib import Path
4
+ from .data import open_observations, load_probabilities, prepare_probabilities
5
+ from .scalable import generate_dask
6
+
7
+
8
+ def main(argv=None):
9
+ p=argparse.ArgumentParser(description="Generate daily gridded weather from a seasonal tercile forecast")
10
+ from ._version import __version__
11
+ p.add_argument("--version", action="version", version=__version__)
12
+ p.add_argument("--obs",required=True,help="Daily PRCP NetCDF")
13
+ p.add_argument("--forecast",required=True,help="NetCDF with PB,PN,PA")
14
+ p.add_argument("--output",required=True)
15
+ p.add_argument("--year",required=True,type=int)
16
+ p.add_argument("--months",nargs="+",type=int,default=[7,8,9])
17
+ p.add_argument("--climatology",nargs=2,type=int,default=[1991,2020])
18
+ p.add_argument("--members",type=int,default=20)
19
+ p.add_argument("--member-batch",type=int,default=20)
20
+ p.add_argument("--tile",nargs=2,type=int,default=[10,10])
21
+ p.add_argument("--seed",type=int,default=42)
22
+ p.add_argument("--wet-threshold",type=float,default=1.0)
23
+ p.add_argument("--spatial",choices=["distance","independent"],default="distance")
24
+ p.add_argument("--conditioning",choices=["mean","mixture"],default="mean",
25
+ help="mean = averaged parameters; mixture = per-member parameter class (adds between-class variability)")
26
+ p.add_argument("--class-draw",choices=["fitted","shared"],default="fitted")
27
+ p.add_argument("--weighting",choices=["tercile","pdf_ratio","mre","croley"],default="tercile",
28
+ help="mre is selected automatically for --total-window/--onset/--dry-spell/--cessation unless mre/croley is explicit")
29
+ p.add_argument("--total-window",nargs=2,metavar=("MM-DD","MM-DD"),
30
+ help="Window of the main (total) forecast when it differs from --months, e.g. 07-01 09-30")
31
+ p.add_argument("--onset",help="NetCDF: onset tercile probabilities (early, normal, late)")
32
+ p.add_argument("--onset-search",nargs=2,default=["05-01","09-30"],metavar=("MM-DD","MM-DD"))
33
+ p.add_argument("--dry-spell",help="NetCDF: post-onset max dry spell probabilities (short, normal, long)")
34
+ p.add_argument("--dry-spell-days",type=int,default=50,help="Post-onset window for --dry-spell")
35
+ p.add_argument("--cessation",help="NetCDF: cessation probabilities (early, normal, late)")
36
+ p.add_argument("--tolerance",type=float,default=0.01,help="Probability tolerance tau of each constraint")
37
+ p.add_argument("--amounts",choices=["gamma","mixed_exponential"],default="gamma")
38
+ p.add_argument("--persistence",choices=["climatology","weighted","independent"],default="climatology")
39
+ p.add_argument("--occurrence",choices=["markov","spell"],default="markov",
40
+ help="spell = run-length dependent (semi-Markov) wet/dry sequence; use with --dry-spell")
41
+ p.add_argument("--no-trace",action="store_true",help="Set sub-threshold rain to zero (v0.1 behaviour)")
42
+ p.add_argument("--bbox",nargs=4,type=float,metavar=("WEST","SOUTH","EAST","NORTH"))
43
+ p.add_argument("--scheduler",choices=["synchronous","threads"],default="synchronous")
44
+ p.add_argument("--workers",type=int,default=2)
45
+ for v in ["TMIN","TMAX","HUMIN","HUMAX","WIND","SOLAR"]:
46
+ p.add_argument("--"+v.lower(),help=f"Daily {v} NetCDF on identical grid/time")
47
+ args=p.parse_args(argv)
48
+ if min(args.members, args.member_batch, args.workers, *args.tile) < 1:
49
+ p.error("members, member-batch, workers and tile sizes must be positive")
50
+ import dask
51
+ paths={"PRCP":args.obs}
52
+ paths.update({v:getattr(args,v.lower()) for v in ["TMIN","TMAX","HUMIN","HUMAX","WIND","SOLAR"] if getattr(args,v.lower())})
53
+ obs=open_observations(paths,chunks={"T":366,"Y":args.tile[0],"X":args.tile[1]})
54
+ if args.bbox:
55
+ w,s,e,n=args.bbox
56
+ obs=obs.where((obs.X>=w)&(obs.X<=e)&(obs.Y>=s)&(obs.Y<=n),drop=True)
57
+ forecast=load_probabilities(args.forecast,target=obs)
58
+ constraints=_constraints(args)
59
+ weighting=args.weighting if not constraints or args.weighting in ("mre","croley") else "mre"
60
+ print(f"Generating {args.members} members on {obs.sizes['Y']} x {obs.sizes['X']} cells; baseline {args.climatology} (must match forecast provider).")
61
+ with dask.config.set(scheduler=args.scheduler,num_workers=args.workers):
62
+ result=generate_dask(obs,forecast,year=args.year,n_members=args.members,months=tuple(args.months),
63
+ climatology=tuple(args.climatology),tile_shape=tuple(args.tile),member_batch=args.member_batch,
64
+ seed=args.seed,wet_threshold=args.wet_threshold,spatial=args.spatial,
65
+ conditioning=args.conditioning,class_draw=args.class_draw,weighting=weighting,
66
+ constraints=constraints or None,occurrence=args.occurrence,
67
+ amount_distribution=args.amounts,persistence=args.persistence,trace_rainfall=not args.no_trace)
68
+ dest=Path(args.output)
69
+ dest.parent.mkdir(parents=True,exist_ok=True)
70
+ result.to_netcdf(dest,encoding={v:{"zlib":True,"complevel":4,
71
+ "dtype":result[v].dtype if result[v].dtype.kind in "iu" else "float32"} for v in result})
72
+ print(f"Written {dest}")
73
+
74
+ def _constraints(args):
75
+ """Build SeasonalConstraints from the optional forecast files."""
76
+ from .attributes import SeasonalTotal, OnsetDate, MaxDrySpell, CessationDate
77
+ from .mre import SeasonalConstraint
78
+ import xarray as xr
79
+ def read(path):
80
+ from .mre import relabel_probabilities
81
+ with xr.open_dataset(path) as ds:
82
+ candidates = []
83
+ for da in ds.data_vars.values():
84
+ try:
85
+ candidates.append(prepare_probabilities(relabel_probabilities(da)))
86
+ except (ValueError, IndexError):
87
+ continue
88
+ if len(candidates) != 1:
89
+ raise ValueError(f"{path}: expected one labelled tercile-probability variable with one forecast time")
90
+ return candidates[0].load()
91
+ out=[]
92
+ if not (args.total_window or args.onset or args.dry_spell or args.cessation):
93
+ return out
94
+ if args.total_window:
95
+ out.append(SeasonalConstraint(SeasonalTotal(tuple(args.total_window)),None,args.tolerance))
96
+ onset=OnsetDate(search=tuple(args.onset_search))
97
+ if args.onset:
98
+ out.append(SeasonalConstraint(onset,read(args.onset),args.tolerance))
99
+ if args.dry_spell:
100
+ out.append(SeasonalConstraint(MaxDrySpell(after_onset=onset,length_days=args.dry_spell_days),read(args.dry_spell),args.tolerance))
101
+ if args.cessation:
102
+ out.append(SeasonalConstraint(CessationDate(),read(args.cessation),args.tolerance))
103
+ return out
104
+
105
+ if __name__ == "__main__":
106
+ main()