diff --git a/unseen/bias_correction.py b/unseen/bias_correction.py index fe52e51..6a46ee0 100644 --- a/unseen/bias_correction.py +++ b/unseen/bias_correction.py @@ -213,12 +213,20 @@ def _parse_command_line(): ) parser.add_argument( "--min_lead", + type=int, default=None, - help="Minimum lead time to include in analysis (int or filename)", + help="Minimum lead time to include in analysis", ) parser.add_argument( - "--min_lead_kwargs", + "--min_lead_file", + type=str, + default=None, + help="Name of file containing the minimum lead time to include in analysis", + ) + parser.add_argument( + "--min_lead_file_kwargs", nargs="*", + default={}, action=general_utils.store_dict, help="Optional fileio.open_dataset kwargs for lead independence (e.g., spatial_agg=median)", ) @@ -270,17 +278,17 @@ def _main(): # Mask lead times below min_lead if args.min_lead: - if isinstance(args.min_lead, str): - # Load min_lead from file - ds_min_lead = fileio.open_dataset(args.min_lead, **args.min_lead_kwargs) - min_lead = ds_min_lead["min_lead"].load() - da_fcst = da_fcst.groupby(f"{args.init_dim}.month").where( - da_fcst[args.lead_dim] >= min_lead - ) - da_fcst = da_fcst.drop_vars("month") - else: - min_lead = args.min_lead - da_fcst = da_fcst.where(da_fcst[args.lead_dim] >= min_lead) + min_lead = int(args.min_lead) + da_fcst = da_fcst.where(da_fcst[args.lead_dim] >= min_lead) + + elif args.min_lead_file: + # Load min_lead from file + ds_min_lead = fileio.open_dataset(args.min_lead_file, **args.min_lead_kwargs) + min_lead = ds_min_lead["min_lead"].load() + da_fcst = da_fcst.groupby(f"{args.init_dim}.month").where( + da_fcst[args.lead_dim] >= min_lead + ) + da_fcst = da_fcst.drop_vars("month") # Calculate bias bias = get_bias( diff --git a/unseen/eva.py b/unseen/eva.py index da1a3d8..963a5ed 100644 --- a/unseen/eva.py +++ b/unseen/eva.py @@ -66,14 +66,13 @@ def fit_gev( core_dim="time", stationary=True, covariate=None, - fitstart="LMM", + fitstart="scipy_fitstart", loc1=0, scale1=0, fc=None, floc=None, fscale=None, bounds=None, - retry_fit=False, assert_good_fit=False, goodness_of_fit_kwargs={}, pick_best_model=False, @@ -103,9 +102,6 @@ def fit_gev( Fixed values for the shape, location and scale parameters. bounds : list of tuples, optional Custom bounds for the shape, loc0, loc1, scale0 and scale1 parameters. - retry_fit : bool, default True - Retry fit with initial estimate generated by passing data[::2] to - fitstart. The best fit is returned. See notes. assert_good_fit : bool, default False Return NaNs if any of the parameters fall outside the support of the distribution (stationary only) @@ -161,8 +157,8 @@ def fit_gev( used to determine if stationary or nonstationary parameters are returned. If the stationary fit is better, the nonstationary parameters are returned with zero trends (see `check_gev_relative_fit` and `get_best_GEV_model_1d`). - - The `retry_fit` may be deprecated in future versions as `use_basinhopping` - may be a better method to improve the fit. + - The `assert_good_fit` will return NaNs if the parameters don't pass a + goodness of fit test (see `check_gev_fit`). - Use `unpack_gev_params` to get the shape, location and scale parameters as a separate array. If nonstationary, the output will also have three parameters that have an extra covariate dimension. @@ -199,7 +195,6 @@ def _fit_1d( floc, fscale, bounds, - retry_fit, assert_good_fit, pick_best_model, alpha, @@ -210,10 +205,11 @@ def _fit_1d( scipy_fit_kwargs, ): """Estimate distribution parameters.""" + + n = 3 if stationary else 5 if np.all(~np.isfinite(data)): # Return NaNs if all input data is infinite - n = 3 if stationary else 5 - return np.array([np.nan] * n) + return np.full(n, np.nan) if np.isnan(data).any(): # Drop NaNs in data @@ -225,6 +221,9 @@ def _fit_1d( # Initial estimates of distribution parameters for MLE if isinstance(fitstart, str): dparams_i = _fitstart_1d(data, fitstart, scipy_fit_kwargs) + if np.isnan(dparams_i).any(): + # If the fitstart method fails, return NaNs + return np.full(n, np.nan) else: # User provided initial estimates dparams_i = fitstart @@ -247,22 +246,6 @@ def _fit_1d( ) dparams = np.array([i for i in dparams], dtype="float64") - if retry_fit: - # Retry fit using alternative fitstart method - _kwargs = kwargs.copy() - _kwargs["stationary"] = True - for k in ["retry_fit", "assert_good_fit", "pick_best_model"]: - _kwargs[k] = False # Avoids recursion - _kwargs["fitstart"] = _fitstart_1d(data[::2], fitstart) - dparams_alt = _fit_1d(data, covariate, **_kwargs) - - # Test if the alternative fit is better - nll_1 = _gev_nllf(dparams, data) - nll_2 = _gev_nllf(dparams_alt, data) - if nll_2 < nll_1: - dparams = dparams_alt - warnings.warn("Better fit estimate using data[::2].") - if not stationary: dparams_ns_i = [dparams_i[0], dparams_i[1], loc1, dparams_i[2], scale1] @@ -549,8 +532,15 @@ def _fitstart_1d(data, method, scipy_fit_kwargs={}): if method == "LMM": # L-moments method - dparams_i = distr.gev.lmom_fit(data) - dparams_i = list(dparams_i.values()) + try: + dparams_i = distr.gev.lmom_fit(data) + dparams_i = list(dparams_i.values()) + except ValueError as e: + warnings.warn( + "L-moments failed to generate a initial guess for MLE. " + + f"Considering changing the fitstart method.{e}" + ) + return e elif method == "scipy_fitstart": # Moments method? @@ -1004,15 +994,16 @@ def gev_confidence_interval( n_resamples=1000, ci=0.95, core_dim="time", - covariate=0, + stationary=True, + covariate=None, + return_covariate=None, fit_kwargs={}, + rng=None, ): - """ - Bootstrapped confidence intervals for return periods or return levels. + """Bootstrapped confidence intervals for return periods or return levels. Parameters: ----------- - data : xarray.DataArray Input data to fit GEV distribution dparams : xarray.DataArray, optional @@ -1028,9 +1019,16 @@ def gev_confidence_interval( ci : float, optional Confidence level (e.g., 0.95 for 95% confidence interval, default: 0.95) core_dim : str, optional - The core dimension along which to apply GEV fitting (default: None, will auto-detect) + The core dimension along which to fit GEV. + stationary : bool, optional + covariate : xarray.DataArray, optional + Covariate for nonstationary GEV fit. + return_coviariate : xarray.DataArray, optional + Covariate values in which to evaluate return levels or return periods. fit_kwargs : dict, optional Additional keyword arguments to pass to `fit_gev` + rng : numpy.random.Generator, optional + Random number generator for reproducibility. If None, a default RNG is used. Returns: -------- @@ -1038,29 +1036,57 @@ def gev_confidence_interval( Confidence intervals with lower and upper bounds along dim 'quantile' """ - # Replace core dim with the one from the fit_kwargs if it exists + # Ensure no duplicate kwargs in fit_kwargs core_dim = fit_kwargs.pop("core_dim", core_dim) + stationary = fit_kwargs.pop("stationary", stationary) covariate = fit_kwargs.pop("covariate", covariate) - rng = np.random.default_rng(seed=0) + if rng is None: + rng = np.random.default_rng(seed=0) + + if (return_period is not None) and (return_level is not None): + raise ValueError("Only one of return_period or return_level can be provided.") + + if not stationary: + if covariate is None: + raise ValueError("Covariate must be provided for a nonstationary fit.") + if return_covariate is None: + return_covariate = covariate + warnings.warn( + "return_coviariate not provided. Evaluating CI at all covariates." + ) + assert hasattr(covariate, core_dim) + assert hasattr(return_covariate, core_dim) + if dparams is None: - dparams = fit_gev(data, covariate, core_dim=core_dim, **fit_kwargs) - shape, loc, scale = unpack_gev_params(dparams, covariate) + dparams = fit_gev( + data, covariate, stationary=stationary, core_dim=core_dim, **fit_kwargs + ) if bootstrap_method == "parametric": # Generate bootstrapped data using the GEV distribution + shape, loc, scale = unpack_gev_params(dparams, covariate) + if stationary: + input_core_dims = [[], [], []] + else: + input_core_dims = [[], [core_dim], [core_dim]] boot_data = apply_ufunc( genextreme.rvs, shape, loc, scale, - input_core_dims=[[], [], []], + input_core_dims=input_core_dims, output_core_dims=[["k", core_dim]], - kwargs=dict(size=(n_resamples, data[core_dim].size)), + kwargs=dict(size=(n_resamples, data[core_dim].size), random_state=rng), vectorize=True, dask="parallelized", + output_dtypes=["float64"], + dask_gufunc_kwargs={ + "output_sizes": {"k": n_resamples, core_dim: data[core_dim].size} + }, ) boot_data = boot_data.transpose("k", core_dim, ...) + boot_covariate = covariate elif bootstrap_method == "non-parametric": # Resample data with replacements @@ -1072,18 +1098,36 @@ def gev_confidence_interval( ) indexer = DataArray(resample_indices, dims=("k", core_dim)) boot_data = data.isel({core_dim: indexer}) + if not stationary: + boot_covariate = covariate.isel({core_dim: indexer}) + else: + boot_covariate = covariate # Fit GEV parameters to resampled data - gev_params_resampled = fit_gev(boot_data, core_dim=core_dim, **fit_kwargs) + gev_params_resampled = fit_gev( + boot_data, + core_dim=core_dim, + stationary=stationary, + covariate=boot_covariate, + **fit_kwargs, + ) if return_period is not None: result = get_return_level( - return_period, gev_params_resampled, core_dim=core_dim, covariate=covariate + return_period, + gev_params_resampled, + core_dim=core_dim, + covariate=return_covariate, ) elif return_level is not None: result = get_return_period( - return_level, gev_params_resampled, core_dim=core_dim, covariate=covariate + return_level, + gev_params_resampled, + core_dim=core_dim, + covariate=return_covariate, ) + else: + return gev_params_resampled # Bounds of confidence intervals ci = ci * 100 # Avoid rounding errors @@ -1126,6 +1170,7 @@ def gev_return_curve( fit_kwargs : dict, optional Additional keyword arguments to pass to `fit_gev` """ + rng = np.random.default_rng(seed=0) # GEV fit to data @@ -1176,7 +1221,7 @@ def gev_return_curve( np.isfinite(boot_event_return_periods) ] event_return_period_lower_ci = np.quantile(boot_event_return_periods, q) - event_return_period_upper_ci = np.quantile(boot_event_return_periods, q - 1) + event_return_period_upper_ci = np.quantile(boot_event_return_periods, 1 - q) event_data = ( event_return_period, event_return_period_lower_ci, @@ -1254,16 +1299,16 @@ def plot_gev_return_curve( ax.plot( curve_return_periods, curve_values, - color="tab:blue", + color="#0000FF", label="GEV fit to data", ) ax.fill_between( curve_return_periods, curve_values_lower_ci, curve_values_upper_ci, - color="tab:blue", + color="#0000FF", alpha=0.2, - label="95% CI on GEV fit", + label="Uncertainty of GEV fit", ) ax.plot( [event_return_period_lower_ci, event_return_period_upper_ci], @@ -1271,16 +1316,16 @@ def plot_gev_return_curve( color="0.5", marker="|", linestyle=":", - label="95% CI for record event", + label="Uncertainty of record event", ) empirical_return_values = np.sort(data, axis=None)[::-1] empirical_return_periods = len(data) / np.arange(1.0, len(data) + 1.0) ax.scatter( empirical_return_periods, empirical_return_values, - color="tab:blue", + color="#0000FF", alpha=0.5, - label="empirical data", + label="Data", ) rp = f"{event_return_period:.0f}" rp_lower = f"{event_return_period_lower_ci:.0f}" @@ -1302,14 +1347,14 @@ def plot_gev_return_curve( handles, labels = ax.get_legend_handles_labels() handles = [handles[3], handles[0], handles[1], handles[2]] labels = [labels[3], labels[0], labels[1], labels[2]] - ax.legend(handles, labels, loc="upper left") + ax.legend(handles, labels) ax.set_xscale("log") ax.set_xlabel("return period (years)") if ylabel: ax.set_ylabel(ylabel) if ylim: ax.set_ylim(ylim) - ax.grid() + ax.grid("both", which="major", linestyle="--", linewidth=0.5) def plot_nonstationary_pdfs( @@ -1320,7 +1365,7 @@ def plot_nonstationary_pdfs( ax=None, title="", units=None, - cmap="rainbow", + cmap="plasma", outfile=None, ): """Plot stationary and nonstationary GEV PDFs. @@ -1348,14 +1393,16 @@ def plot_nonstationary_pdfs( if ax is None: fig, ax = plt.subplots(1, 1, figsize=(10, 7)) - ax.set_title(title, loc="left") + ax.set_title(title) n = covariate.size - colors = colormaps[cmap](np.linspace(0, 1, n)) + colors = colormaps[cmap](np.linspace(0, 0.8, n)) shape, loc, scale = unpack_gev_params(dparams_ns, covariate) # Histogram. - _, bins, _ = ax.hist(data, bins=40, density=True, alpha=0.5, label="Histogram") + _, bins, _ = ax.hist( + data, bins="auto", density=True, alpha=0.2, label="Histogram", color="k" + ) # Stationary GEV PDF shape_s, loc_s, scale_s = dparams_s @@ -1365,13 +1412,13 @@ def plot_nonstationary_pdfs( # Nonstationary GEV PDFs for i, t in enumerate(covariate.values): pdf_ns = genextreme.pdf(bins, shape, loc=loc[i], scale=scale[i]) - ax.plot(bins, pdf_ns, lw=1.6, c=colors[i], zorder=0, label=t) + ax.plot(bins, pdf_ns, lw=2.2, c=colors[i], zorder=0, label=t) ax.set_xlabel(units) ax.set_ylabel("Probability") ax.xaxis.set_minor_locator(AutoMinorLocator()) ax.yaxis.set_minor_locator(AutoMinorLocator()) - ax.legend(loc="upper right", bbox_to_anchor=(1, 1), framealpha=0.3) + ax.legend(loc="upper right", bbox_to_anchor=(1, 1), framealpha=0.3, fontsize=13.5) ax.set_xmargin(1e-3) if outfile: @@ -1601,7 +1648,7 @@ def spatial_plot_gev_parameters( ax.yaxis.set_visible(True) if dataset_name: - fig.suptitle(f"{dataset_name} GEV parameters", y=0.75 if stationary else 0.97) + fig.suptitle(f"{dataset_name} GEV parameters", y=0.77 if stationary else 0.96) # Hide empty subplots for ax in [ax for ax in axes if not ax.collections]: @@ -1644,7 +1691,7 @@ def _parse_command_line(): ) parser.add_argument( "--fitstart", - default="LMM", + default="scipy_fitstart", choices=( "LMM", "scipy", @@ -1657,12 +1704,6 @@ def _parse_command_line(): help="Initial guess method (or estimate) of the GEV parameters", ) - parser.add_argument( - "--retry_fit", - action="store_true", - default=False, - help="Retry fit if it doesn't pass the goodness of fit test", - ) parser.add_argument( "--use_basinhopping", action="store_true", @@ -1694,15 +1735,23 @@ def _parse_command_line(): "--covariate_file", type=str, default=None, help="Covariate file" ) # todo: test this parser.add_argument( - "--min_lead", default=None, help="Minimum lead time (int or filename)" + "--min_lead", + type=int, + default=None, + help="Minimum lead time to include in analysis", ) parser.add_argument( - "--min_lead_kwargs", + "--min_lead_file", type=str, + default=None, + help="Name of file containing the minimum lead time to include in analysis", + ) + parser.add_argument( + "--min_lead_file_kwargs", nargs="*", default={}, action=general_utils.store_dict, - help="Keyword arguments for opening min_lead file", + help="Optional fileio.open_dataset kwargs for lead independence (e.g., spatial_agg=median)", ) parser.add_argument( "--drop_max", @@ -1759,18 +1808,18 @@ def _main(): if args.reference_time_period: ds = time_utils.select_time_period(ds, args.reference_time_period) - # Filter data by minimum lead time + # Mask lead times below min_lead if args.min_lead: - if isinstance(args.min_lead, str): - # Load min_lead from file - ds_min_lead = fileio.open_dataset(args.min_lead, **args.min_lead_kwargs) - min_lead = ds_min_lead["min_lead"].load() - ds = ds.groupby(f"{args.init_dim}.month").where( - ds[args.lead_dim] >= min_lead - ) - ds = ds.drop_vars("month") - else: - ds = ds.where(ds[args.lead_dim] >= args.min_lead) + min_lead = int(args.min_lead) + ds = ds.where(ds[args.lead_dim] >= min_lead) + + elif args.min_lead_file: + # Load min_lead from file + ds_min_lead = fileio.open_dataset(args.min_lead_file, **args.min_lead_kwargs) + min_lead = ds_min_lead["min_lead"].load() + + ds = ds.groupby(f"{args.init_dim}.month").where(ds[args.lead_dim] >= min_lead) + ds = ds.drop_vars("month") # Stack ensemble, init and lead dimensions along new "sample" dimension if args.stack_dims: @@ -1801,7 +1850,6 @@ def _main(): stationary=args.stationary, fitstart=args.fitstart, covariate=covariate, - retry_fit=args.retry_fit, use_basinhopping=args.use_basinhopping, assert_good_fit=args.assert_good_fit, pick_best_model=args.pick_best_model, diff --git a/unseen/fileio.py b/unseen/fileio.py index 30838c0..6f2961e 100644 --- a/unseen/fileio.py +++ b/unseen/fileio.py @@ -512,7 +512,7 @@ def _fix_metadata(ds, metadata_file): with open(metadata_file, "r") as reader: metadata_dict = yaml.load(reader, Loader=yaml.BaseLoader) - valid_keys = ["rename", "drop_coords", "round_coords", "units"] + valid_keys = ["rename", "drop_coords", "round_coords", "units", "custom_code"] for key in metadata_dict.keys(): if key not in valid_keys: raise KeyError(f"Invalid metadata key: {key}") @@ -538,6 +538,14 @@ def _fix_metadata(ds, metadata_file): if var in ds.data_vars: ds[var].attrs["units"] = units + if "custom_code" in metadata_dict: + try: + for code_str in metadata_dict["custom_code"]: + ds = eval(code_str) + except Exception as e: + warnings.warn(f"Custom code failed: {e} (must return ds)") + raise e + return ds @@ -998,6 +1006,7 @@ def _main(): if args.output_chunks: ds = ds.chunk(args.output_chunks) + if args.time_agg_dates: ds = ds.set_coords(("event_time")) ds = ds[kwargs["variables"]] diff --git a/unseen/general_utils.py b/unseen/general_utils.py index e1c829f..82cc9c8 100644 --- a/unseen/general_utils.py +++ b/unseen/general_utils.py @@ -237,6 +237,7 @@ def plot_timeseries_scatter( obs_label=None, units=None, time_dim="time", + add_legend=True, outfile=None, ): """Timeseries scatter plot of ensemble and observed data. @@ -284,20 +285,21 @@ def plot_timeseries_scatter( ax.scatter( da_obs[time_dim], da_obs, - s=20, + s=28, c="k", marker="x", label=obs_label, zorder=10, ) # Plot ensemble data - ax.scatter(da[time_dim], da, s=5, c="deepskyblue", label=label) + ax.scatter(da[time_dim], da, s=8, c="deepskyblue", label=label) ax.set_ylabel(units) ax.set_xmargin(1e-2) ax.xaxis.set_minor_locator(AutoMinorLocator()) ax.yaxis.set_minor_locator(AutoMinorLocator()) - ax.legend() + if add_legend: + ax.legend() if outfile: plt.tight_layout() diff --git a/unseen/independence.py b/unseen/independence.py index 70c4552..9f9ba2e 100644 --- a/unseen/independence.py +++ b/unseen/independence.py @@ -442,6 +442,8 @@ def _get_null_correlation_bounds( n_init_dates = len(da[init_dim]) n_ensembles = len(da[ensemble_dim]) + # # Fixes "tuple indices must be integers or slices, not tuple" error + # da = da.compute() da_stacked = da.stack(sample=(init_dim, lead_dim, ensemble_dim)) null_correlations = [] @@ -526,6 +528,14 @@ def _parse_command_line(): default={}, help="Chunks for writing data to file (e.g. init_date=-1 lead_time=-1)", ) + parser.add_argument( + "--file_kwargs", + type=str, + nargs="*", + default={}, + action=general_utils.store_dict, + help="Keyword arguments for opening the data file (excluding sel and variables)", + ) args = parser.parse_args() return args @@ -541,8 +551,12 @@ def _main(): print(client) ds_fcst = fileio.open_dataset( - args.fcst_file, variables=[args.var], sel=args.spatial_selection + args.fcst_file, + variables=[args.var], + sel=args.spatial_selection, + **args.file_kwargs, ) + da_fcst = ds_fcst[args.var] ds = run_tests( diff --git a/unseen/moments.py b/unseen/moments.py index d94f2bd..eb24fb8 100644 --- a/unseen/moments.py +++ b/unseen/moments.py @@ -330,12 +330,20 @@ def _parse_command_line(): ) parser.add_argument( "--min_lead", + type=int, default=None, - help="Minimum lead time to include in analysis (int or filename)", + help="Minimum lead time to include in analysis", ) parser.add_argument( - "--min_lead_kwargs", + "--min_lead_file", + type=str, + default=None, + help="Name of file containing the minimum lead time to include in analysis", + ) + parser.add_argument( + "--min_lead_file_kwargs", nargs="*", + default={}, action=general_utils.store_dict, help="Optional fileio.open_dataset kwargs for lead independence (e.g., spatial_agg=median)", ) @@ -360,19 +368,19 @@ def _main(): # Mask lead times below min_lead if args.min_lead: - if isinstance(args.min_lead, str): - # Load min_lead from file - ds_min_lead = fileio.open_dataset(args.min_lead, **args.min_lead_kwargs) - min_lead = ds_min_lead["min_lead"].load() - # Assumes min_lead has only one init month - assert min_lead.month.size == 1, "Not implemented for multiple init months" - min_lead = min_lead.drop_vars("month") - if min_lead.size == 1: - min_lead = min_lead.item() - else: - min_lead = args.min_lead + min_lead = int(args.min_lead) da_fcst = da_fcst.where(da_fcst[args.lead_dim] >= min_lead) + elif args.min_lead_file: + # Load min_lead from file + ds_min_lead = fileio.open_dataset(args.min_lead_file, **args.min_lead_kwargs) + min_lead = ds_min_lead["min_lead"].load() + + da_fcst = da_fcst.groupby(f"{args.init_dim}.month").where( + da_fcst[args.lead_dim] >= min_lead + ) + da_fcst = da_fcst.drop_vars("month") + ds_obs = fileio.open_dataset(args.obs_file) da_obs = ds_obs[args.var].dropna("time") diff --git a/unseen/similarity.py b/unseen/similarity.py index 42bc465..4a0a21f 100644 --- a/unseen/similarity.py +++ b/unseen/similarity.py @@ -317,13 +317,19 @@ def _parse_command_line(): ) parser.add_argument( "--min_lead", + type=int, default=None, - help="Minimum lead time to include in analysis (int or filename)", + help="Minimum lead time to include in analysis", ) parser.add_argument( - "--min_lead_kwargs", - nargs="*", + "--min_lead_file", type=str, + default=None, + help="Name of file containing the minimum lead time to include in analysis", + ) + parser.add_argument( + "--min_lead_file_kwargs", + nargs="*", default={}, action=general_utils.store_dict, help="Optional fileio.open_dataset kwargs for lead independence (e.g., spatial_agg=median)", @@ -366,18 +372,20 @@ def _main(): start_date, end_date = args.reference_time_period ds_obs = ds_obs.sel({args.time_dim: slice(start_date, end_date)}) + # Mask lead times below min_lead if args.min_lead: - if isinstance(args.min_lead, str): - # Load min_lead from file - ds_min_lead = fileio.open_dataset(args.min_lead, **args.min_lead_kwargs) - min_lead = ds_min_lead["min_lead"].load() - ds_fcst = ds_fcst.groupby(f"{args.init_dim}.month").where( - ds_fcst[args.lead_dim] >= min_lead - ) - ds_fcst = ds_fcst.drop_vars("month") - else: - min_lead = args.min_lead - ds_fcst = ds_fcst.where(ds_fcst[args.lead_dim] >= min_lead) + min_lead = int(args.min_lead) + ds_fcst = ds_fcst.where(ds_fcst[args.lead_dim] >= min_lead) + + elif args.min_lead_file: + # Load min_lead from file + ds_min_lead = fileio.open_dataset(args.min_lead_file, **args.min_lead_kwargs) + min_lead = ds_min_lead["min_lead"].load() + + ds_fcst = ds_fcst.groupby(f"{args.init_dim}.month").where( + ds_fcst[args.lead_dim] >= min_lead + ) + ds_fcst = ds_fcst.drop_vars("month") ds_similarity = similarity_tests( ds_fcst, diff --git a/unseen/tests/test_eva.py b/unseen/tests/test_eva.py index 6ff5c9a..adc0ab9 100644 --- a/unseen/tests/test_eva.py +++ b/unseen/tests/test_eva.py @@ -5,7 +5,12 @@ import pytest import xarray as xr -from unseen.eva import fit_gev, get_return_period, get_return_level +from unseen.eva import ( + fit_gev, + get_return_period, + get_return_level, + gev_confidence_interval, +) rtol = 0.3 # relative tolerance for testing close values @@ -129,16 +134,6 @@ def test_fit_ns_gev_3d(example_da_gev_3d): assert np.all(dparams.isel(dparams=2) > 0) # Positive trend in location -@pytest.mark.parametrize("example_da_gev", ["xarray", "numpy", "dask"], indirect=True) -def test_fit_gev_1d_retry_fit(example_da_gev): - """Run stationary GEV fit using 1D array & retry_fit.""" - data, dparams_i = example_da_gev - # Set large alpha to force any fit considered bad - dparams = fit_gev(data, stationary=True, retry_fit=True, alpha=1) - # Check fitted params match params used to create data - npt.assert_allclose(dparams, dparams_i, rtol=rtol) - - @pytest.mark.parametrize("example_da_gev", ["xarray", "numpy", "dask"], indirect=True) def test_fit_gev_1d_assert_good_fit(example_da_gev): """Run stationary GEV fit using 1D array & assert_good_fit.""" @@ -268,3 +263,53 @@ def test_get_return_level_3d(example_da_gev_3d): assert return_level.shape == ari.shape assert np.all(np.isfinite(return_level)) + + +@pytest.mark.parametrize("example_da_gev", ["xarray", "dask"], indirect=True) +@pytest.mark.parametrize("stationary", [True, False]) +@pytest.mark.parametrize("ari", [100, np.array([10, 100, 1000])]) +@pytest.mark.parametrize("bootstrap_method", ["parametric", "non-parametric"]) +def test_gev_confidence_interval(example_da_gev, stationary, ari, bootstrap_method): + """Test get_confidence_intervals function.""" + + data, dparams_i = example_da_gev + core_dim = "time" + + if not stationary: + data = add_example_gev_trend(data) + covariate = xr.DataArray(np.arange(data.time.size), dims="time") + elif stationary: + covariate = None + return_covariate = xr.DataArray([0], dims="time") + + if isinstance(ari, int): + ari = np.array([ari]) + ari = xr.DataArray(ari, dims="return_period") + + dparams = fit_gev( + data, + covariate=covariate, + stationary=stationary, + core_dim=core_dim, + ) + + return_level = get_return_level( + ari, dparams, core_dim=core_dim, covariate=return_covariate + ) + + ci_bounds = gev_confidence_interval( + data, + dparams, + return_period=ari, + bootstrap_method=bootstrap_method, + n_resamples=100, + ci=0.95, + core_dim=core_dim, + stationary=stationary, + covariate=covariate, + return_covariate=return_covariate, + ) + + # Check return level is between CI bounds + assert np.all(return_level >= ci_bounds.isel(quantile=0)) + assert np.all(return_level <= ci_bounds.isel(quantile=1))