diff --git a/docs/api_reference.md b/docs/api_reference.md index f3c100d0..51b9848b 100644 --- a/docs/api_reference.md +++ b/docs/api_reference.md @@ -40,6 +40,8 @@ hide: ::: diffwofost.physical_models.test.EngineTestHelper +::: diffwofost.physical_models.weather.to_weather_data_iterator + ::: diffwofost.io.save_model ::: diffwofost.io.load_model diff --git a/pyproject.toml b/pyproject.toml index 22a7a01c..78139183 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -28,6 +28,7 @@ dependencies = [ "safetensors", "torch", "pcse", + "xarray", ] description = "Differentiable WOFOST" keywords = ["wofost", "pytorch", "differentiable", "crop", "optimization"] diff --git a/src/diffwofost/io.py b/src/diffwofost/io.py index d7a38f17..7f6346f4 100644 --- a/src/diffwofost/io.py +++ b/src/diffwofost/io.py @@ -27,6 +27,10 @@ from diffwofost.physical_models.config import Configuration from diffwofost.physical_models.config import load_class +# Use the temporary ParameterProvider. The original PCSE ParameterProvider could be used when +# https://github.com/ajwdewit/pcse/pull/121 is merged an a new version of PCSE released. +from diffwofost.physical_models.parameter_providers import ParameterProvider + def _default_model_filename(model): """Stable filename from class name + init_kwargs hash.""" @@ -224,8 +228,6 @@ def load_model(path, *, model_class=None, device=None, dtype=None): param_data[group_name]["nontensor_params"] = params # Reconstruct ParameterProvider - from pcse.base.parameter_providers import ParameterProvider - init_kwargs = {} for group_name in ("sitedata", "timerdata", "soildata", "cropdata"): key = f"_{group_name}" diff --git a/src/diffwofost/physical_models/crop/assimilation.py b/src/diffwofost/physical_models/crop/assimilation.py index 698af9fc..d6a7a19e 100644 --- a/src/diffwofost/physical_models/crop/assimilation.py +++ b/src/diffwofost/physical_models/crop/assimilation.py @@ -6,14 +6,12 @@ from pcse.base import SimulationObject from pcse.base.parameter_providers import ParameterProvider from pcse.base.variablekiosk import VariableKiosk -from pcse.base.weather import WeatherDataContainer from diffwofost.physical_models.base import TensorParamTemplate from diffwofost.physical_models.base import TensorRatesTemplate from diffwofost.physical_models.config import ComputeConfig from diffwofost.physical_models.traitlets import Tensor from diffwofost.physical_models.utils import AfgenTrait from diffwofost.physical_models.utils import _broadcast_to -from diffwofost.physical_models.utils import _get_drv from diffwofost.physical_models.utils import astro # --------------------------------------------------------------------------- @@ -343,7 +341,7 @@ def initialize( # elements (which share the same weather driver). self._astro_cache: dict = {} - def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None) -> None: + def calc_rates(self, day: datetime.date, drv: dict) -> torch.Tensor: """Compute the potential gross assimilation rate (PGASS).""" p = self.params r = self.rates @@ -356,9 +354,9 @@ def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None lai = _broadcast_to(k["LAI"], self.params.shape, dtype=self.dtype, device=self.device) # Weather drivers - irrad = _get_drv(drv.IRRAD, self.params.shape, dtype=self.dtype, device=self.device) - dtemp = _get_drv(drv.DTEMP, self.params.shape, dtype=self.dtype, device=self.device) - tmin = _get_drv(drv.TMIN, self.params.shape, dtype=self.dtype, device=self.device) + irrad = drv["IRRAD"] + dtemp = drv["DTEMP"] + tmin = drv["TMIN"] # Assimilation is zero before crop emergence (DVS < 0) dvs_mask = dvs >= 0 @@ -373,15 +371,9 @@ def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None # latitude and radiation are passed directly – they may be scalars or # tensors; the function returns torch.Tensor results in all cases. dayl, _daylp, sinld, cosld, difpp, _atmtr, dsinbe, _angot = astro( - day, drv.LAT, drv.IRRAD, dtype=self.dtype, device=self.device + day, drv["LAT"], drv["IRRAD"], dtype=self.dtype, device=self.device ) - dayl_t = _broadcast_to(dayl, self.params.shape, dtype=self.dtype, device=self.device) - sinld_t = _broadcast_to(sinld, self.params.shape, dtype=self.dtype, device=self.device) - cosld_t = _broadcast_to(cosld, self.params.shape, dtype=self.dtype, device=self.device) - difpp_t = _broadcast_to(difpp, self.params.shape, dtype=self.dtype, device=self.device) - dsinbe_t = _broadcast_to(dsinbe, self.params.shape, dtype=self.dtype, device=self.device) - # Parameter tables amax = p.AMAXTB(dvs) amax = amax * p.TMPFTB(dtemp) @@ -389,16 +381,16 @@ def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None eff = p.EFFTB(dtemp) dtga = totass7( - dayl_t, + dayl, amax, eff, lai, kdif, irrad, - difpp_t, - dsinbe_t, - sinld_t, - cosld_t, + difpp, + dsinbe, + sinld, + cosld, epsilon=self._epsilon, dtype=self.dtype, device=self.device, @@ -414,11 +406,11 @@ def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None r.PGASS = pgass * dvs_mask return r.PGASS - def __call__(self, day: datetime.date = None, drv: WeatherDataContainer = None) -> torch.Tensor: + def __call__(self, day: datetime.date, drv: dict) -> torch.Tensor: """Calculate and return the potential gross assimilation rate (PGASS).""" return self.calc_rates(day, drv) - def integrate(self, day: datetime.date = None, delt=1.0) -> None: + def integrate(self, day: datetime.date, delt: float = 1.0) -> None: """No state variables to integrate for this module.""" return diff --git a/src/diffwofost/physical_models/crop/evapotranspiration.py b/src/diffwofost/physical_models/crop/evapotranspiration.py index 6ddf4712..11e2d09d 100644 --- a/src/diffwofost/physical_models/crop/evapotranspiration.py +++ b/src/diffwofost/physical_models/crop/evapotranspiration.py @@ -3,7 +3,6 @@ from pcse.base import SimulationObject from pcse.base.parameter_providers import ParameterProvider from pcse.base.variablekiosk import VariableKiosk -from pcse.base.weather import WeatherDataContainer from pcse.traitlets import Any from pcse.traitlets import Bool from pcse.traitlets import Instance @@ -14,7 +13,6 @@ from diffwofost.physical_models.traitlets import Tensor from diffwofost.physical_models.utils import AfgenTrait from diffwofost.physical_models.utils import _broadcast_to -from diffwofost.physical_models.utils import _get_drv def SWEAF(ET0: torch.Tensor, DEPNR: torch.Tensor) -> torch.Tensor: @@ -102,22 +100,21 @@ def initialize( else: self.etmodule = Evapotranspiration(day, kiosk, parvalues, shape=shape) - def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None): + def calc_rates(self, day: datetime.date, drv: dict): """Delegate rate calculation to the selected evapotranspiration module. Args: - day (datetime.date, optional): The current date of the simulation. - drv (WeatherDataContainer, optional): A dictionary-like container holding - weather data elements as key/value. The values are - arrays or scalars. See PCSE documentation for details. + day (datetime.date): The current date of the simulation. + drv (dict): A container holding weather data elements as key/value. The values are + arrays or scalars. """ return self.etmodule.calc_rates(day, drv) - def __call__(self, day: datetime.date = None, drv: WeatherDataContainer = None): + def __call__(self, day: datetime.date, drv: dict): """Callable interface for rate calculation.""" return self.calc_rates(day, drv) - def integrate(self, day: datetime.date = None, delt=1.0) -> None: + def integrate(self, day: datetime.date, delt: float = 1.0) -> None: """Delegate state integration to the selected evapotranspiration module. Args: @@ -190,11 +187,11 @@ def _initialize_base( self._IDWST = torch.zeros(shape, dtype=self.dtype, device=self.device) self._IDOST = torch.zeros(shape, dtype=self.dtype, device=self.device) - def __call__(self, day: datetime.date = None, drv: WeatherDataContainer = None): + def __call__(self, day: datetime.date, drv: dict): """Callable interface for rate calculation.""" return self.calc_rates(day, drv) - def integrate(self, day: datetime.date = None, delt=1.0) -> None: + def integrate(self, day: datetime.date, delt: float = 1.0) -> None: """Accumulate stress-day counters for water and oxygen stress.""" rfws_stress = (self.rates.RFWS < 1.0).to(dtype=self.dtype) rfos_stress = (self.rates.RFOS < 1.0).to(dtype=self.dtype) @@ -211,11 +208,11 @@ def finalize(self, day: datetime.date) -> None: class _BaseEvapotranspirationNonLayered(_BaseEvapotranspiration): """Shared implementation for non-layered evapotranspiration.""" - def _rf_tramx_co2(self, drv: WeatherDataContainer, et0: torch.Tensor) -> torch.Tensor: + def _rf_tramx_co2(self, drv: dict, et0: torch.Tensor) -> torch.Tensor: """Return CO2 reduction factor for TRAMX (no CO2 effect in base implementation).""" return torch.ones_like(et0) - def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None): + def calc_rates(self, day: datetime.date, drv: dict): p = self.params r = self.rates k = self.kiosk @@ -227,9 +224,9 @@ def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None # TODO see #22 dvs = _broadcast_to(k["DVS"], self.params_shape, dtype=self.dtype, device=self.device) - et0 = _get_drv(drv.ET0, self.params_shape, dtype=self.dtype, device=self.device) - e0 = _get_drv(drv.E0, self.params_shape, dtype=self.dtype, device=self.device) - es0 = _get_drv(drv.ES0, self.params_shape, dtype=self.dtype, device=self.device) + et0 = torch.as_tensor(drv["ET0"]) + e0 = torch.as_tensor(drv["E0"]) + es0 = torch.as_tensor(drv["ES0"]) rf_tramx_co2 = self._rf_tramx_co2(drv, et0) # If DVS < 0, the crop has not yet emerged, so we zero the rates using a mask @@ -483,10 +480,10 @@ def initialize( shape=shape, ) - def _rf_tramx_co2(self, drv: WeatherDataContainer, et0: torch.Tensor) -> torch.Tensor: + def _rf_tramx_co2(self, drv: dict, et0: torch.Tensor) -> torch.Tensor: """Calculate CO2 reduction factor for TRAMX based on atmospheric CO2 concentration.""" - if hasattr(drv, "CO2") and drv.CO2 is not None: - co2 = _get_drv(drv.CO2, self.params_shape, dtype=self.dtype, device=self.device) + if "CO2" in drv and drv["CO2"] is not None: + co2 = drv["CO2"] else: co2 = self.params.CO2 return self.params.CO2TRATB(co2) @@ -634,15 +631,15 @@ def initialize( # Internal DSOS tracker for layered oxygen-stress response self._dsos = torch.zeros(self.params_shape, dtype=self.dtype, device=self.device) - def _rf_tramx_co2(self, drv: WeatherDataContainer, et0: torch.Tensor) -> torch.Tensor: + def _rf_tramx_co2(self, drv: dict, et0: torch.Tensor) -> torch.Tensor: """Calculate CO2 reduction factor for TRAMX using CO2 from driver or parameters.""" - if hasattr(drv, "CO2") and drv.CO2 is not None: - co2 = _get_drv(drv.CO2, self.params_shape, dtype=self.dtype, device=self.device) + if "CO2" in drv and drv["CO2"] is not None: + co2 = drv["CO2"] else: co2 = self.params.CO2 return self.params.CO2TRATB(co2) - def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None): + def calc_rates(self, day: datetime.date, drv: dict): """Calculate daily evapotranspiration rates per soil layer with CO2 effects. Computes transpiration and stress factors for each soil layer based on root @@ -658,9 +655,9 @@ def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None n_layers = self._n_layers - et0 = _get_drv(drv.ET0, self.params_shape, dtype=self.dtype, device=self.device) - e0 = _get_drv(drv.E0, self.params_shape, dtype=self.dtype, device=self.device) - es0 = _get_drv(drv.ES0, self.params_shape, dtype=self.dtype, device=self.device) + et0 = torch.as_tensor(drv["ET0"]) + e0 = torch.as_tensor(drv["E0"]) + es0 = torch.as_tensor(drv["ES0"]) # reduction factor for CO2 on TRAMX rf_tramx_co2 = self._rf_tramx_co2(drv, et0) @@ -786,11 +783,11 @@ def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None r.IDOS = bool(torch.any(r.RFOS < 1.0)) return r.TRA, r.TRAMX - def __call__(self, day: datetime.date = None, drv: WeatherDataContainer = None): + def __call__(self, day: datetime.date, drv: dict): """Callable interface for rate calculation.""" return self.calc_rates(day, drv) - def integrate(self, day: datetime.date = None, delt=1.0) -> None: + def integrate(self, day: datetime.date, delt: float = 1.0) -> None: """Accumulate stress-day counters based on any layer experiencing stress.""" rfws_stress = (self.rates.RFWS < 1.0).any(dim=0).to(dtype=self.dtype) rfos_stress = (self.rates.RFOS < 1.0).any(dim=0).to(dtype=self.dtype) diff --git a/src/diffwofost/physical_models/crop/leaf_dynamics.py b/src/diffwofost/physical_models/crop/leaf_dynamics.py index 9897aad1..e5fe7dba 100644 --- a/src/diffwofost/physical_models/crop/leaf_dynamics.py +++ b/src/diffwofost/physical_models/crop/leaf_dynamics.py @@ -5,14 +5,12 @@ from pcse.base import SimulationObject from pcse.base.parameter_providers import ParameterProvider from pcse.base.variablekiosk import VariableKiosk -from pcse.base.weather import WeatherDataContainer from diffwofost.physical_models.base import TensorParamTemplate from diffwofost.physical_models.base import TensorRatesTemplate from diffwofost.physical_models.base import TensorStatesTemplate from diffwofost.physical_models.config import ComputeConfig from diffwofost.physical_models.traitlets import Tensor from diffwofost.physical_models.utils import AfgenTrait -from diffwofost.physical_models.utils import _get_drv class WOFOST_Leaf_Dynamics(SimulationObject): @@ -248,14 +246,13 @@ def _calc_LAI(self): total_LAI = self.states.LASUM + SAI + PAI return total_LAI - def calc_rates(self, day: datetime.date, drv: WeatherDataContainer) -> None: + def calc_rates(self, day: datetime.date, drv: dict) -> None: """Calculate the rates of change for the leaf dynamics. Args: - day (datetime.date, optional): The current date of the simulation. - drv (WeatherDataContainer, optional): A dictionary-like container holding - weather data elements as key/value. The values are - arrays or scalars. See PCSE documentation for details. + day (datetime.date): The current date of the simulation. + drv (dict): A container holding weather data elements as key/value. The values are + arrays or scalars. """ r = self.rates s = self.states @@ -326,7 +323,7 @@ def calc_rates(self, day: datetime.date, drv: WeatherDataContainer) -> None: r.DRLV = torch.maximum(r.DSLV, r.DALV) # Get the temperature from the drv - TEMP = _get_drv(drv.TEMP, p.shape, self.dtype, self.device) + TEMP = drv["TEMP"] # physiologic ageing of leaves per time step FYSAGE = (TEMP - p.TBASE) / (35.0 - p.TBASE) diff --git a/src/diffwofost/physical_models/crop/phenology.py b/src/diffwofost/physical_models/crop/phenology.py index 474a872f..224e161e 100644 --- a/src/diffwofost/physical_models/crop/phenology.py +++ b/src/diffwofost/physical_models/crop/phenology.py @@ -19,8 +19,6 @@ from diffwofost.physical_models.config import ComputeConfig from diffwofost.physical_models.traitlets import Tensor from diffwofost.physical_models.utils import AfgenTrait -from diffwofost.physical_models.utils import _broadcast_to -from diffwofost.physical_models.utils import _get_drv from diffwofost.physical_models.utils import _restore_state from diffwofost.physical_models.utils import _snapshot_state from diffwofost.physical_models.utils import daylength @@ -182,7 +180,7 @@ def calc_rates(self, day, drv): VERNBASE = params.VERNBASE DVS = self.kiosk["DVS"] - TEMP = _get_drv(drv.TEMP, self.params.shape, self.dtype, self.device) + TEMP = drv["TEMP"] # Operate elementwise only on elements not yet vernalised not_vernalised = ~self.states.ISVERNALISED @@ -505,15 +503,13 @@ def calc_rates(self, day, drv): p = self.params r = self.rates s = self.states - # Day length sensitivity # daylength returns a Tensor directly; broadcast to parameter shape. - DAYLP = daylength(day, drv.LAT, dtype=self.dtype, device=self.device) - DAYLP_t = _broadcast_to(DAYLP, p.shape, dtype=self.dtype, device=self.device) + DAYLP = daylength(day, drv["LAT"], dtype=self.dtype, device=self.device) # Compute DVRED conditionally based on IDSL >= 1 safe_den = p.DLO - p.DLC safe_den = safe_den.sign() * torch.maximum(torch.abs(safe_den), self._epsilon) - dvred_active = torch.clamp((DAYLP_t - p.DLC) / safe_den, 0.0, 1.0) + dvred_active = torch.clamp((DAYLP - p.DLC) / safe_den, 0.0, 1.0) DVRED = torch.where(p.IDSL >= 1, dvred_active, self._ones) # Vernalisation factor - always compute if module exists @@ -529,7 +525,7 @@ def calc_rates(self, day, drv): self._ones, ) - TEMP = _get_drv(drv.TEMP, p.shape, self.dtype, self.device) + TEMP = drv["TEMP"] # Initialize all rate variables r.DTSUME = self._zeros diff --git a/src/diffwofost/physical_models/crop/respiration.py b/src/diffwofost/physical_models/crop/respiration.py index 45e8e5b3..058a35d5 100644 --- a/src/diffwofost/physical_models/crop/respiration.py +++ b/src/diffwofost/physical_models/crop/respiration.py @@ -5,14 +5,12 @@ from pcse.base import SimulationObject from pcse.base.parameter_providers import ParameterProvider from pcse.base.variablekiosk import VariableKiosk -from pcse.base.weather import WeatherDataContainer from diffwofost.physical_models.base import TensorParamTemplate from diffwofost.physical_models.base import TensorRatesTemplate from diffwofost.physical_models.config import ComputeConfig from diffwofost.physical_models.traitlets import Tensor from diffwofost.physical_models.utils import AfgenTrait from diffwofost.physical_models.utils import _broadcast_to -from diffwofost.physical_models.utils import _get_drv class WOFOST_Maintenance_Respiration(SimulationObject): @@ -112,7 +110,7 @@ def initialize( self.rates = self.RateVariables(kiosk, shape=shape) self.kiosk = kiosk - def calc_rates(self, day: datetime.date, drv: WeatherDataContainer): + def calc_rates(self, day: datetime.date, drv: dict): """Calculate maintenance respiration rates. Args: @@ -138,7 +136,7 @@ def calc_rates(self, day: datetime.date, drv: WeatherDataContainer): # TODO see #22 DVS = _broadcast_to(kk["DVS"], p.shape, self.dtype, self.device) - TEMP = _get_drv(drv.TEMP, p.shape, self.dtype, self.device) + TEMP = drv["TEMP"] RMRES = RMR * WRT + RML * WLV + RMS * WST + RMO * WSO RMRES = RMRES * p.RFSETB(DVS) @@ -148,7 +146,7 @@ def calc_rates(self, day: datetime.date, drv: WeatherDataContainer): # No maintenance respiration before emergence (DVS < 0). r.PMRES = torch.where(DVS < 0, torch.zeros_like(PMRES), PMRES) - def __call__(self, day: datetime.date, drv: WeatherDataContainer): + def __call__(self, day: datetime.date, drv: dict): """Calculate and return maintenance respiration (PMRES).""" self.calc_rates(day, drv) return self.rates.PMRES diff --git a/src/diffwofost/physical_models/crop/root_dynamics.py b/src/diffwofost/physical_models/crop/root_dynamics.py index 29262069..39366479 100644 --- a/src/diffwofost/physical_models/crop/root_dynamics.py +++ b/src/diffwofost/physical_models/crop/root_dynamics.py @@ -3,7 +3,6 @@ from pcse.base import SimulationObject from pcse.base.parameter_providers import ParameterProvider from pcse.base.variablekiosk import VariableKiosk -from pcse.base.weather import WeatherDataContainer from diffwofost.physical_models.base import TensorParamTemplate from diffwofost.physical_models.base import TensorRatesTemplate from diffwofost.physical_models.base import TensorStatesTemplate @@ -193,14 +192,13 @@ def initialize( shape=shape, ) - def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None) -> None: + def calc_rates(self, day: datetime.date, drv: dict) -> None: """Calculate the rates of change of the state variables. Args: - day (datetime.date, optional): The current date of the simulation. - drv (WeatherDataContainer, optional): A dictionary-like container holding - weather data elements as key/value. The values are - arrays or scalars. See PCSE documentation for details. + day (datetime.date): The current date of the simulation. + drv (dict): A container holding weather data elements as key/value. The values are + arrays or scalars. """ p = self.params r = self.rates diff --git a/src/diffwofost/physical_models/crop/stem_dynamics.py b/src/diffwofost/physical_models/crop/stem_dynamics.py index 8b0bb370..5b52b061 100644 --- a/src/diffwofost/physical_models/crop/stem_dynamics.py +++ b/src/diffwofost/physical_models/crop/stem_dynamics.py @@ -3,7 +3,6 @@ from pcse.base import SimulationObject from pcse.base.parameter_providers import ParameterProvider from pcse.base.variablekiosk import VariableKiosk -from pcse.base.weather import WeatherDataContainer from diffwofost.physical_models.base import TensorParamTemplate from diffwofost.physical_models.base import TensorRatesTemplate from diffwofost.physical_models.base import TensorStatesTemplate @@ -160,14 +159,13 @@ def initialize( kiosk, publish=["TWST", "WST", "SAI"], WST=WST, DWST=DWST, TWST=TWST, SAI=SAI ) - def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None) -> None: + def calc_rates(self, day: datetime.date, drv: dict) -> None: """Calculate the rates of change of the state variables. Args: day (datetime.date, optional): The current date of the simulation. - drv (WeatherDataContainer, optional): A dictionary-like container holding - weather data elements as key/value. The values are - arrays or scalars. See PCSE documentation for details. + drv (dict, optional): A container holding weather data elements as key/value. The values + are arrays or scalars. """ r = self.rates s = self.states @@ -196,11 +194,11 @@ def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None r.GWST = r.GRST - r.DRST - REALLOC_ST - def integrate(self, day: datetime.date = None, delt=1.0) -> None: + def integrate(self, day: datetime.date, delt: float = 1.0) -> None: """Integrate the state variables using the rates of change. Args: - day (datetime.date, optional): The current date of the simulation. + day (datetime.date): The current date of the simulation. delt (float, optional): The time step for integration. Defaults to 1.0. """ p = self.params diff --git a/src/diffwofost/physical_models/crop/storage_organ_dynamics.py b/src/diffwofost/physical_models/crop/storage_organ_dynamics.py index 5e48979f..7488f0ac 100644 --- a/src/diffwofost/physical_models/crop/storage_organ_dynamics.py +++ b/src/diffwofost/physical_models/crop/storage_organ_dynamics.py @@ -4,7 +4,6 @@ from pcse.base import SimulationObject from pcse.base.parameter_providers import ParameterProvider from pcse.base.variablekiosk import VariableKiosk -from pcse.base.weather import WeatherDataContainer from diffwofost.physical_models.base import TensorParamTemplate from diffwofost.physical_models.base import TensorStatesTemplate from diffwofost.physical_models.config import ComputeConfig @@ -144,13 +143,12 @@ def initialize( kiosk, publish=["TWSO", "WSO", "PAI"], WSO=WSO, DWSO=DWSO, TWSO=TWSO, PAI=PAI ) - def calc_rates(self, day: datetime.date = None, drv: WeatherDataContainer = None) -> None: + def calc_rates(self, day: datetime.date, drv: dict) -> None: """Calculate the rates of change of the state variables. Args: day (datetime.date, optional): The current date of the simulation. - drv (WeatherDataContainer, optional): A dictionary-like container holding - weather data elements as key/value. + drv (dict, optional): A container holding weather data elements as key/value. """ rates = self.rates k = self.kiosk diff --git a/src/diffwofost/physical_models/crop/wofost72.py b/src/diffwofost/physical_models/crop/wofost72.py index 59f1fdd2..bf775362 100644 --- a/src/diffwofost/physical_models/crop/wofost72.py +++ b/src/diffwofost/physical_models/crop/wofost72.py @@ -5,7 +5,6 @@ from pcse.base import SimulationObject from pcse.base.parameter_providers import ParameterProvider from pcse.base.variablekiosk import VariableKiosk -from pcse.base.weather import WeatherDataContainer from pcse.traitlets import Instance from pcse.traitlets import Unicode from diffwofost.physical_models.base import TensorParamTemplate @@ -239,14 +238,13 @@ def _check_carbon_balance(day, DMI, GASS, MRES, CVF, pf): ) raise exc.CarbonBalanceError(msg) - def calc_rates(self, day: datetime.date, drv: WeatherDataContainer) -> None: + def calc_rates(self, day: datetime.date, drv: dict) -> None: """Calculate the rates of change of the state variables. Args: day (datetime.date): The current date of the simulation. - drv (WeatherDataContainer): A dictionary-like container holding - weather data elements as key/value. The values are - arrays or scalars. See PCSE documentation for details. + drv (dict): A container holding weather data elements as key/value. The values are + arrays or scalars. """ p = self.params r = self.rates diff --git a/src/diffwofost/physical_models/engine.py b/src/diffwofost/physical_models/engine.py index e27d412e..9b644ef4 100644 --- a/src/diffwofost/physical_models/engine.py +++ b/src/diffwofost/physical_models/engine.py @@ -6,8 +6,10 @@ """ import gc +from collections.abc import Iterator from collections.abc import MutableMapping from pathlib import Path +from typing import Any import torch from pcse import signals from pcse.base import BaseEngine @@ -28,6 +30,7 @@ class Engine(PcseEngine): """ mconf = Instance(Configuration) + weatherdataprovider = Instance(Iterator) parameterprovider = Instance(MutableMapping) def __init__( @@ -100,8 +103,8 @@ def setup( Args: parameterprovider: Provider with crop and soil parameter values. - weatherdataprovider: Provider used to retrieve daily driving - weather variables. + weatherdataprovider: Daily driving weather variables. It can be an provided as an + iterable or iterator. agromanagement: AgroManagement definition passed to the configured agromanagement component. external_states (list[dict] | None): Optional list of day-keyed @@ -113,7 +116,6 @@ def setup( self._reset_runtime_state() self.parameterprovider = parameterprovider - self._shape = _get_params_shape(self.parameterprovider) # Variable kiosk for registering and publishing variables self.kiosk = VariableKiosk(external_states) @@ -136,9 +138,12 @@ def setup( self.kiosk(self.day) # Driving variables - self.weatherdataprovider = weatherdataprovider + self.weatherdataprovider = _to_iterator(weatherdataprovider) self.drv = self._get_driving_variables(self.day) + # Determine common shape for the parameters and weather data + self._shape = _get_shape(self.parameterprovider, self.drv) + # Call AgroManagement module for management actions at initialization self.agromanager(self.day, None) @@ -217,27 +222,56 @@ def _finish_cropsimulation(self, day): self.crop.finalize(day) self._save_summary_output() + def _get_driving_variables(self, day): + """Get driving variables and return it.""" + drv = next(self.weatherdataprovider) + if "DAY" in drv: + assert drv["DAY"] == day, "Wrong day!" + return drv + + +def _get_shape(parameterprovider: MutableMapping, drivingvariables: dict[str, Any]) -> tuple: + """Infer common tensor shape from the parameter provider and the driving variables. + + Args: + parameterprovider: Parameter provider. + drivingvariables: Weather data. + + Raises: + ValueError: If non-matching shapes are found for the data providers. -def _get_params_shape(parameterprovider): + Returns: + tuple: Shared tensor shape. + """ + params_shape = _get_params_shape(parameterprovider) + weather_shape = _get_params_shape(drivingvariables) + if params_shape and weather_shape and params_shape != weather_shape: + raise ValueError( + "Non-matching shapes between parameter and weather data: " + f"{params_shape} and {weather_shape}" + ) + return params_shape or weather_shape + + +def _get_params_shape(provider: MutableMapping) -> tuple: """Infer the common tensor batch shape from a parameter provider. Afgen table parameters are expected to have an extra trailing dimension for table coordinates, which is ignored when determining the simulation shape. Args: - parameterprovider: Parameter provider containing scalar and tensor - parameters. + provider: Parameter provider containing scalar and tensor parameters. Returns: tuple: Shared tensor shape for all tensor-valued parameters, or an - empty tuple when all parameters are scalar. + empty tuple when all parameters are scalar. Raises: ValueError: If tensor parameters do not share a common shape. """ shape = () - for paramname in parameterprovider._unique_parameters: - param = parameterprovider[paramname] + for paramname in provider.keys(): + param = provider[paramname] if isinstance(param, torch.Tensor): # We need to drop the last dimension from the Afgen table parameters param_shape = param.shape[:-1] if paramname.endswith("TB") else param.shape @@ -248,3 +282,17 @@ def _get_params_shape(parameterprovider): else: raise ValueError("Non-matching shapes found in parameter provider!") return shape + + +def _to_iterator(weatherdata): + """Transform input weather data to an iterator.""" + # if already an iterator, return as is + if hasattr(weatherdata, "__iter__") and hasattr(weatherdata, "__next__"): + return weatherdata + # if iterable, return iterator + elif hasattr(weatherdata, "__iter__"): + return iter(weatherdata) + else: + raise ValueError( + f"Weather data should be provided as an iterable or iterator - got {type(weatherdata)}." + ) diff --git a/src/diffwofost/physical_models/parameter_providers.py b/src/diffwofost/physical_models/parameter_providers.py index 204acc1e..44f968c2 100644 --- a/src/diffwofost/physical_models/parameter_providers.py +++ b/src/diffwofost/physical_models/parameter_providers.py @@ -5,7 +5,8 @@ class ParameterProvider(PcseParameterProvider): """Temporary implementation of PCSE's ParameterProvider. Fixes the `__iter__` method in order to allow for access via dict-like `.items()`, `.keys()`, - and `.values()`. Could be dropped when https://github.com/ajwdewit/pcse/pull/121 is merged. + and `.values()`. Could be dropped when https://github.com/ajwdewit/pcse/pull/121 is merged and + a new version of PCSE released. """ def __iter__(self): diff --git a/src/diffwofost/physical_models/soil/classic_waterbalance.py b/src/diffwofost/physical_models/soil/classic_waterbalance.py index 89ceaefe..e16126af 100644 --- a/src/diffwofost/physical_models/soil/classic_waterbalance.py +++ b/src/diffwofost/physical_models/soil/classic_waterbalance.py @@ -81,8 +81,8 @@ def calc_rates(self, day, drv): # the potential soil/water evaporation rates directly because there is # no shading by the canopy. if "TRA" not in self.kiosk: - r.WTRA = torch.zeros_like(torch.as_tensor(drv.ES0)) - EVSMX = torch.as_tensor(drv.ES0) + r.WTRA = torch.zeros_like(drv["ES0"]) + EVSMX = drv["ES0"] else: r.WTRA = self.kiosk["TRA"] EVSMX = torch.as_tensor(self.kiosk["EVSMX"]) @@ -101,7 +101,7 @@ def calc_rates(self, day, drv): self.DSLR = torch.where(rain_ge_1, torch.ones_like(dslr_inc), dslr_inc) # Hold rainfall amount to keep track of soil surface wetness and reset self.DSLR if needed - self.RAINold = torch.as_tensor(drv.RAIN) + self.RAINold = drv["RAIN"] def integrate(self, day, delt=1.0): """Integrate state variables over one time step.""" @@ -405,9 +405,8 @@ def calc_rates(self, day, drv): Args: day (datetime.date): The current date of the simulation. - drv (WeatherDataContainer): A dictionary-like container holding - weather data elements as key/value. The values are - arrays or scalars. See PCSE documentation for details. + drv (dict): A container holding weather data elements as key/value. The values are + arrays or scalars. """ s = self.states @@ -424,9 +423,9 @@ def calc_rates(self, day, drv): # Transpiration and maximum evaporation rates from crop module. # Before emergence there is no canopy shading yet, so the water balance # must fall back to the weather-driven soil and surface evaporation. - weather_wtra = torch.zeros_like(torch.as_tensor(drv.ES0, dtype=dtype, device=device)) - weather_evwmx = torch.as_tensor(drv.E0, dtype=dtype, device=device) - weather_evsmx = torch.as_tensor(drv.ES0, dtype=dtype, device=device) + weather_wtra = torch.zeros_like(drv["ES0"], dtype=dtype, device=device) + weather_evwmx = drv["E0"] + weather_evsmx = drv["ES0"] if "TRA" not in k: r.WTRA = weather_wtra @@ -472,7 +471,7 @@ def calc_rates(self, day, drv): ) # Potentially infiltrating rainfall - RAIN_t = torch.as_tensor(drv.RAIN, dtype=dtype, device=device) + RAIN_t = drv["RAIN"] RINPRE_fixed = (1.0 - p.NOTINF) * RAIN_t RINPRE_storm = (1.0 - p.NOTINF * self.NINFTB(RAIN_t)) * RAIN_t # IFUNRN: 0 = fixed non-infiltrating fraction, 1 = function of storm size diff --git a/src/diffwofost/physical_models/test.py b/src/diffwofost/physical_models/test.py index 1dbbd7c0..2f425ae4 100644 --- a/src/diffwofost/physical_models/test.py +++ b/src/diffwofost/physical_models/test.py @@ -1,3 +1,4 @@ +import pandas as pd import torch import yaml from pcse import signals @@ -7,6 +8,7 @@ from diffwofost.physical_models.config import ComputeConfig from diffwofost.physical_models.engine import Engine from diffwofost.physical_models.parameter_providers import ParameterProvider +from diffwofost.physical_models.weather import to_weather_data_iterator class EngineTestHelper(Engine): @@ -41,8 +43,23 @@ def _run(self): self._terminate_simulation(self.day) +class WeatherDataContainerTestHelper(WeatherDataContainer): + """A helper class for creating WeatherDataContainer instances. + + It adds support for dict-like indexing of weather data, to provide compatibility with the + interface used within diffWOFOST. + """ + + def __getitem__(self, key): + return getattr(self, key) + + class WeatherDataProviderTestHelper(WeatherDataProvider): - """It stores the weatherdata contained within the YAML tests.""" + """A helper class for creating a WeatherDataProvider instance. + + Needed to provide a data structure that can be used by both PCSE and diffWOFOST, useful for + verifying the compatibility of diffWOFOST modules within PCSE. + """ def __init__(self, yaml_weather, meteo_range_checks=True): super().__init__() @@ -53,12 +70,17 @@ def __init__(self, yaml_weather, meteo_range_checks=True): settings.METEO_RANGE_CHECKS = meteo_range_checks for weather in yaml_weather: weather_inputs = {k: v for k, v in weather.items() if k != "SNOWDEPTH"} - wdc = WeatherDataContainer(**weather_inputs) + wdc = WeatherDataContainerTestHelper(**weather_inputs) self._store_WeatherDataContainer(wdc, wdc.DAY) def prepare_engine_input( - test_data, crop_model_params, device=None, dtype=None, meteo_range_checks=True + test_data, + crop_model_params, + device=None, + dtype=None, + return_weather_data_provider=False, # set True for tests that patch PCSE + meteo_range_checks=True, ): """Prepare the inputs for the engine from the YAML file.""" # If not specified, use default dtype and device @@ -70,33 +92,21 @@ def prepare_engine_input( agro_management_inputs = test_data["AgroManagement"] cropd = test_data["ModelParameters"] - weather_data_provider = WeatherDataProviderTestHelper( - test_data["WeatherVariables"], meteo_range_checks=meteo_range_checks - ) + weather_data = test_data["WeatherVariables"] + if return_weather_data_provider: + # If required, return the PCSE-compatible data structure for weather data + weather_data_provider = WeatherDataProviderTestHelper( + weather_data, meteo_range_checks=meteo_range_checks + ) + else: + weather_data_df = pd.DataFrame(weather_data) + if "DTEMP" not in weather_data_df.columns: + weather_data_df["DTEMP"] = (weather_data_df["TEMP"] + weather_data_df["TMAX"]) / 2.0 + weather_data_iterator = to_weather_data_iterator(weather_data_df, check=meteo_range_checks) + + # create a list out of the iterator, so that the weather data can be reused in several tests + weather_data_provider = list(weather_data_iterator) - # The PCSE WeatherDataContainer stores required variables as Python floats. - # Some of our tests rely on weather inputs being torch.Tensors (e.g. to - # broadcast/batch weather variables). We only do this conversion when - # METEO_RANGE_CHECKS is disabled because the PCSE range checks assume - # scalar floats. - if not meteo_range_checks: - for (_, _), wdc in weather_data_provider.store.items(): - for varname in ( - "IRRAD", - "TMIN", - "TMAX", - "TEMP", - "VAP", - "RAIN", - "WIND", - "E0", - "ES0", - "ET0", - ): - if hasattr(wdc, varname): - value = getattr(wdc, varname) - if not isinstance(value, torch.Tensor): - setattr(wdc, varname, torch.tensor(value, dtype=dtype, device=device)) crop_model_params_provider = ParameterProvider(cropdata=cropd) external_states = test_data.get("ExternalStates") or [] diff --git a/src/diffwofost/physical_models/weather.py b/src/diffwofost/physical_models/weather.py new file mode 100644 index 00000000..840ab92e --- /dev/null +++ b/src/diffwofost/physical_models/weather.py @@ -0,0 +1,217 @@ +from collections.abc import Iterator +from dataclasses import dataclass +from typing import Any +import numpy as np +import pandas as pd +import torch +import xarray as xr +from diffwofost.physical_models.config import ComputeConfig + + +@dataclass(frozen=True) +class WeatherVariable: + unit: str + min: float + max: float + + +# These are the weather variables recognized by diffWOFOST internally, along +# with their units and valid ranges. +WEATHER_VARIABLES = { + "LAT": WeatherVariable("Degrees", -90.0, 90.0), + "LON": WeatherVariable("Degrees", -180.0, 180.0), + "ELEV": WeatherVariable("m", -300, 6000), + "IRRAD": WeatherVariable("J/m2/day", 0.0, 40e6), + "TMIN": WeatherVariable("Celsius", -50.0, 60.0), + "TMAX": WeatherVariable("Celsius", -50.0, 60.0), + "VAP": WeatherVariable("hPa", 0.06, 199.3), + "RAIN": WeatherVariable("cm/day", 0, 25), + "E0": WeatherVariable("cm/day", 0.0, 2.5), + "ES0": WeatherVariable("cm/day", 0.0, 2.5), + "ET0": WeatherVariable("cm/day", 0.0, 2.5), + "SNOWDEPTH": WeatherVariable("cm", 0.0, 250.0), + "TEMP": WeatherVariable("Celsius", -50.0, 60.0), + "TMINRA": WeatherVariable("Celsius", -50.0, 60.0), + "WIND": WeatherVariable("m/s", 0.0, 100.0), + "DTEMP": WeatherVariable("Celsius", -50.0, 60.0), +} + +TIME_DIM_NAMES = {"day", "time", "dates"} + + +def to_weather_data_iterator( + data: pd.DataFrame | xr.Dataset, + check: bool = True, + skipna: bool = True, + time_dim: str | None = None, +) -> Iterator: + """Weather data generator from a Pandas DataFrame or an xarray Dataset. + + This utility function transforms weather data from tabular or nd-array format to an iterator of + torch tensors that can be fed to diffWOFOST's engine. + + Args: + data (pd.DataFrame | xr.Dataset): DataFrame or Dataset containing weather data. Weather + variables should be listed as columns (DataFrame) or data variables (Dataset). In order + to be interpreted as weather variables, they should be named as the keys of + `diffwofost.physical_models.weather.WEATHER_VARIABLES`. Rows/elements are expected to + represent daily time steps (an optional column/1D-coordinate named "DAY" should list the + corresponding dates). + check (bool, optional): Optionally carry out validity checks for the dataset. Defaults to + True. + skipna (bool, optional): How to handle NaN values when `check` is True. If True, allow NaN + values as part of the weather data. + time_dim (str, optional): name of the dimension to iterate over. Only relevant if `data` is + a xr.Dataset object. If not provided, the function will try to guess it from: + * the name of the dimension of the "DAY" coordinate (if present). + * the first element of `diffwofost.physical_models.weather.TIME_DIM_NAMES` in + `data.dims`. + + Yields: + dict[str, typing.Any]: Weather variables as key-value pairs. Variables will be converted + to torch tensors, using dtype and device as configured in `ComputeConfig`. + + Examples: + >>> import pandas as pd + >>> weather_data = pd.DataFrame({ + ... "DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"], + ... "TEMP": [10., 11., 9., 12.], + ... }) + >>> weather_data_iter = iterator_from_dataframe(weather_data) + >>> next(weather_data_iter) + {'DAY': datetime.date(2020, 4, 1), 'TEMP': tensor(10.)} + >>> weather_data_faulty = pd.DataFrame({ + ... "DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"], + ... "TEMP": [10., 1000., 9., 12.], # unrealistic temperature + ... }) + >>> iterator_from_dataframe(weather_data_faulty) + ValueError: Values for `TEMP` outside the range [-50.0, 60.0] (expected unit is Celsius). + + """ + _check_weather_variable_keys(data) + + iter_dim = _get_iterator_dimension(data, time_dim) + iter_length = _get_iterator_length(data, iter_dim) + + dates = _extract_dates_if_present(data) + + if check: + _check_weather_variables_range(data, skipna=skipna) + if dates is not None: + _check_dates(dates) + + variables = {} + + # if present, include the dates to the returned variables + if dates is not None: + variables["DAY"] = dates + + variables.update(_to_dict_of_tensors(data, iter_dim)) + + return _iterate(variables, length=iter_length) + + +def _check_weather_variable_keys(data: pd.DataFrame | xr.Dataset) -> None: + num_variables = len([var_name for var_name in WEATHER_VARIABLES if var_name in data]) + if num_variables < 1: + raise ValueError( + "No weather variable found. Variables should be named as the keys of " + "`diffwofost.physical_models.weather.WEATHER_VARIABLES`." + ) + + +def _get_iterator_dimension( + data: pd.DataFrame | xr.Dataset, time_dim: str | None = None +) -> str | None: + """Determine the dimension to iterate over. + + This is only relevant for xr.Dataset objects whose variables have >= 2 dimensions: if the + weather variables are 1D, the only available dimension will be used for iteration. + """ + if isinstance(data, pd.DataFrame) or _are_all_weather_variables_less_than_2d(data): + # there is only one dimension to iterate over. + return None + elif time_dim: + # if the time dimension is provided, only check if it's a valid dimension name + assert time_dim in data.dims, f"Dimension {time_dim} missing from dimensions {data.dims}." + return time_dim + elif "DAY" in data: + # if the time dimension is not provided, check the dimension of the "DAY" coordinate + # (if present) + day = data["DAY"] + assert day.ndim == 1, "Daily dates should be provided as a 1D-coordinate." + return day.dims[0] + else: + # if none of the above, check sensible names for the time dimension + for guess in TIME_DIM_NAMES: + if guess in data.dims: + return guess + raise ValueError( + f"Cannot determine which dimension to iterate over for dataset with dims: {data.dims}." + ) + + +def _get_iterator_length(data: pd.DataFrame | xr.Dataset, iter_dim: str | None) -> int: + return len(data[iter_dim]) if iter_dim is not None else len(data) + + +def _are_all_weather_variables_less_than_2d(data: xr.Dataset) -> bool: + return all( + [var.ndim < 2 for var_name, var in data.variables.items() if var_name in WEATHER_VARIABLES] + ) + + +def _extract_dates_if_present(data: pd.DataFrame | xr.Dataset) -> np.ndarray | None: + if "DAY" not in data: + return None + else: + day = data["DAY"].values + return pd.to_datetime(day).date + + +def _check_weather_variables_range(data: pd.DataFrame | xr.Dataset, skipna: bool = True) -> None: + for var_name, var_range in WEATHER_VARIABLES.items(): + if var_name in data: + var = data[var_name] + is_null = var.isnull() + if not skipna and is_null.any(): + raise ValueError(f"{var_name} includes {int(is_null.sum())} NaN values.") + outside_range = (var < var_range.min) | (var > var_range.max) + is_invalid = outside_range.where(~is_null, other=False) + if is_invalid.any(): + raise ValueError( + f"Values for `{var_name}` outside the range [{var_range.min}, {var_range.max}] " + f"(expected unit is {var_range.unit})." + ) + + +def _check_dates(dates: np.ndarray) -> None: + expected = pd.date_range(start=dates[0], periods=len(dates), freq="D") + if not (dates == expected).all(): + raise ValueError( + "Column `DAY` must contain consecutive daily dates with no gaps or duplicates." + ) + + +def _to_dict_of_tensors( + data: pd.DataFrame | xr.Dataset, + iter_dim: str | None = None, +) -> dict[str, torch.Tensor]: + return { + var_name: _to_tensor(data[var_name], iter_dim) + for var_name in WEATHER_VARIABLES + if var_name in data + } + + +def _to_tensor(data: pd.Series | xr.DataArray, iter_dim: str | None = None) -> torch.Tensor: + device = ComputeConfig.get_device() + dtype = ComputeConfig.get_dtype() + if isinstance(data, xr.DataArray) and iter_dim is not None: + data = data.transpose(iter_dim, ...) + return torch.tensor(data.to_numpy(), device=device, dtype=dtype) + + +def _iterate(variables: dict[str, Any], length): + for n in range(length): + yield {k: v[n] for k, v in variables.items()} diff --git a/tests/physical_models/crop/test_assimilation.py b/tests/physical_models/crop/test_assimilation.py index 2231911f..9243b1fb 100644 --- a/tests/physical_models/crop/test_assimilation.py +++ b/tests/physical_models/crop/test_assimilation.py @@ -249,12 +249,6 @@ def test_assimilation_with_multiple_parameter_arrays(self, device): repeated = crop_model_params_provider[param].repeat(30, 5, 1) crop_model_params_provider.set_override(param, repeated, check=False) - # Make weather drivers match (30, 5) so _get_drv validates/broadcasts. - for (_, _), wdc in weather_data_provider.store.items(): - wdc.IRRAD = torch.ones((30, 5), device=device, dtype=torch.float64) * wdc.IRRAD - wdc.TEMP = torch.ones((30, 5), device=device, dtype=torch.float64) * wdc.TEMP - wdc.TMIN = torch.ones((30, 5), device=device, dtype=torch.float64) * wdc.TMIN - engine = EngineTestHelper(config=assimilation_config) engine.setup( crop_model_params_provider, @@ -294,8 +288,8 @@ def test_assimilation_with_incompatible_parameter_vectors(self): "EFFTB", crop_model_params_provider["EFFTB"].repeat(5, 1), check=False ) + engine = EngineTestHelper(config=assimilation_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=assimilation_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -317,14 +311,27 @@ def test_assimilation_with_incompatible_weather_parameter_vectors(self): crop_model_params_provider.set_override( "AMAXTB", crop_model_params_provider["AMAXTB"].repeat(10, 1), check=False ) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(5, dtype=torch.float64) * wdc.TEMP + # Broadcast weather variables to a shape that does not match the parameters + shape = (5,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + broadcasted = broadcast(weather_data_provider) + + engine = EngineTestHelper(config=assimilation_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=assimilation_config) engine.setup( crop_model_params_provider, - weather_data_provider, + broadcasted, agro_management_inputs, external_states, ) @@ -334,7 +341,7 @@ def test_wofost_pp_with_assimilation(self, test_data_url): test_data = get_test_data(test_data_url) crop_model_params = ["AMAXTB", "EFFTB", "KDIFTB", "TMPFTB", "TMNFTB"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] diff --git a/tests/physical_models/crop/test_evapotranspiration.py b/tests/physical_models/crop/test_evapotranspiration.py index 1a83b088..913afc13 100644 --- a/tests/physical_models/crop/test_evapotranspiration.py +++ b/tests/physical_models/crop/test_evapotranspiration.py @@ -268,23 +268,25 @@ def test_evapotranspiration_with_one_parameter_vector(self, param, device): ) if param == "ET0": - for (_, _), wdc in weather_data_provider.store.items(): - wdc.ET0 = torch.ones(10, dtype=torch.float64, device=wdc.ET0.device) * wdc.ET0 - with pytest.raises(ValueError): - engine = EngineTestHelper(config=evapotranspiration_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - return - - if param == "KDIFTB": + shape = (10,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) + elif param == "KDIFTB": repeated = crop_model_params_provider[param].repeat(10, 1) + crop_model_params_provider.set_override(param, repeated, check=False) else: repeated = crop_model_params_provider[param].repeat(10) - crop_model_params_provider.set_override(param, repeated, check=False) + crop_model_params_provider.set_override(param, repeated, check=False) engine = EngineTestHelper(config=evapotranspiration_config) engine.setup( @@ -455,11 +457,6 @@ def test_evapotranspiration_with_multiple_parameter_arrays(self, device): repeated = crop_model_params_provider[param].broadcast_to(batch_shape) crop_model_params_provider.set_override(param, repeated, check=False) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.ET0 = torch.ones(batch_shape, dtype=torch.float64, device=wdc.ET0.device) * wdc.ET0 - wdc.E0 = torch.ones(batch_shape, dtype=torch.float64, device=wdc.E0.device) * wdc.E0 - wdc.ES0 = torch.ones(batch_shape, dtype=torch.float64, device=wdc.ES0.device) * wdc.ES0 - engine = EngineTestHelper(config=evapotranspiration_config) engine.setup( crop_model_params_provider, @@ -508,8 +505,8 @@ def test_evapotranspiration_with_incompatible_parameter_vectors(self): "DEPNR", crop_model_params_provider["DEPNR"].repeat(5), check=False ) + engine = EngineTestHelper(config=evapotranspiration_config) with pytest.raises(ValueError, match="Non-matching shapes found in parameter provider!"): - engine = EngineTestHelper(config=evapotranspiration_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -541,11 +538,24 @@ def test_evapotranspiration_with_incompatible_weather_parameter_vectors(self): crop_model_params_provider.set_override( "CFET", crop_model_params_provider["CFET"].repeat(10), check=False ) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.ET0 = torch.ones(5, dtype=torch.float64, device=wdc.ET0.device) * wdc.ET0 + # Broadcast weather variables to a shape that does not match the parameters + shape = (5,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) + + engine = EngineTestHelper(config=evapotranspiration_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=evapotranspiration_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -568,7 +578,12 @@ def test_wofost_pp_with_evapotranspiration(self, test_data_url): "SMFCF", ] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params, meteo_range_checks=False) + prepare_engine_input( + test_data, + crop_model_params, + return_weather_data_provider=True, + meteo_range_checks=False, + ) ) expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] @@ -659,7 +674,7 @@ def _kiosk_with_states(): kiosk.set_variable(oid, "SM", torch.tensor(0.25, dtype=torch.float64, device=device)) return kiosk - drv = SimpleNamespace( + drv = dict( ET0=torch.tensor(0.5, dtype=torch.float64, device=device), E0=torch.tensor(0.6, dtype=torch.float64, device=device), ES0=torch.tensor(0.55, dtype=torch.float64, device=device), diff --git a/tests/physical_models/crop/test_leaf_dynamics.py b/tests/physical_models/crop/test_leaf_dynamics.py index 6e28f607..4ad95317 100644 --- a/tests/physical_models/crop/test_leaf_dynamics.py +++ b/tests/physical_models/crop/test_leaf_dynamics.py @@ -135,9 +135,20 @@ def test_leaf_dynamics_with_one_parameter_vector(self, param, device): # Setting a vector (with one value) for the selected parameter if param == "TEMP": - # Vectorize weather variable - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(10, device=device, dtype=torch.float64) * wdc.TEMP + # Broadcast weather variables + shape = (10,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) elif param in ["KDIFTB", "SLATB"]: # AfgenTrait parameters need to have shape (N, M) repeated = crop_model_params_provider[param].repeat(10, 1) @@ -146,48 +157,32 @@ def test_leaf_dynamics_with_one_parameter_vector(self, param, device): repeated = crop_model_params_provider[param].repeat(10) crop_model_params_provider.set_override(param, repeated, check=False) - if param == "TEMP": - # Expect error due to incompatible shapes - # (By defaults parameters are not reshaped following weather variables) - with pytest.raises(ValueError): - engine = EngineTestHelper(config=leaf_dynamics_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - actual_results = engine.get_output() - else: - engine = EngineTestHelper(config=leaf_dynamics_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - actual_results = engine.get_output() + engine = EngineTestHelper(config=leaf_dynamics_config) + engine.setup( + crop_model_params_provider, + weather_data_provider, + agro_management_inputs, + external_states, + ) + engine.run_till_terminate() + actual_results = engine.get_output() - # get expected results from YAML test data - expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] + # get expected results from YAML test data + expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] - assert len(actual_results) == len(expected_results) + assert len(actual_results) == len(expected_results) - for reference, model in zip(expected_results, actual_results, strict=False): - assert reference["DAY"] == model["day"] - # Verify output is on the correct device - for var in expected_precision.keys(): - assert model[var].device.type == device, f"{var} should be on {device}" - # Move to CPU for comparison - model_cpu = { - k: v.cpu() if isinstance(v, torch.Tensor) else v for k, v in model.items() - } - assert all( - all(abs(reference[var] - model_cpu[var]) < precision) - for var, precision in expected_precision.items() - ) + for reference, model in zip(expected_results, actual_results, strict=False): + assert reference["DAY"] == model["day"] + # Verify output is on the correct device + for var in expected_precision.keys(): + assert model[var].device.type == device, f"{var} should be on {device}" + # Move to CPU for comparison + model_cpu = {k: v.cpu() if isinstance(v, torch.Tensor) else v for k, v in model.items()} + assert all( + all(abs(reference[var] - model_cpu[var]) < precision) + for var, precision in expected_precision.items() + ) @pytest.mark.parametrize( "param,delta", @@ -316,9 +311,6 @@ def test_leaf_dynamics_with_multiple_parameter_arrays(self, device): repeated = crop_model_params_provider[param].broadcast_to((30, 5)) crop_model_params_provider.set_override(param, repeated, check=False) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones((30, 5), dtype=torch.float64, device=device) * wdc.TEMP - engine = EngineTestHelper(config=leaf_dynamics_config) engine.setup( crop_model_params_provider, @@ -365,8 +357,8 @@ def test_leaf_dynamics_with_incompatible_parameter_vectors(self): "SPAN", crop_model_params_provider["SPAN"].repeat(5), check=False ) + engine = EngineTestHelper(config=leaf_dynamics_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=leaf_dynamics_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -390,11 +382,23 @@ def test_leaf_dynamics_with_incompatible_weather_parameter_vectors(self): crop_model_params_provider.set_override( "TDWI", crop_model_params_provider["TDWI"].repeat(10), check=False ) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(5, dtype=torch.float64) * wdc.TEMP + # Broadcast weather variables to a shape that does not match the parameters + shape = (5,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + weather_data_provider = broadcast(weather_data_provider) + + engine = EngineTestHelper(config=leaf_dynamics_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=leaf_dynamics_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -408,7 +412,7 @@ def test_wofost_pp_with_leaf_dynamics(self, test_data_url): test_data = get_test_data(test_data_url) crop_model_params = ["SPAN", "TDWI", "TBASE", "PERDL", "RGRLAI", "KDIFTB", "SLATB"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) # get expected results from YAML test data diff --git a/tests/physical_models/crop/test_partitioning.py b/tests/physical_models/crop/test_partitioning.py index d7940829..463ef22c 100644 --- a/tests/physical_models/crop/test_partitioning.py +++ b/tests/physical_models/crop/test_partitioning.py @@ -278,8 +278,8 @@ def test_partitioning_with_incompatible_parameter_vectors(self): crop_model_params_provider.set_override("FRTB", [[0.0, 0.3, 2.0, 0.1]] * 4, check=False) crop_model_params_provider.set_override("FLTB", [[0.0, 0.3, 2.0, 0.1]] * 2, check=False) + engine = EngineTestHelper(config=partitioning_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=partitioning_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -287,13 +287,13 @@ def test_partitioning_with_incompatible_parameter_vectors(self): external_states, ) - @pytest.mark.parametrize("test_data_url", wofost72_data_urls[:1]) + @pytest.mark.parametrize("test_data_url", wofost72_data_urls) def test_wofost_pp_with_partitioning(self, test_data_url): # prepare model input test_data = get_test_data(test_data_url) crop_model_params = ["FRTB", "FLTB", "FSTB", "FOTB"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) # get expected results from YAML test data diff --git a/tests/physical_models/crop/test_phenology.py b/tests/physical_models/crop/test_phenology.py index b25afe5a..b1befeae 100644 --- a/tests/physical_models/crop/test_phenology.py +++ b/tests/physical_models/crop/test_phenology.py @@ -233,10 +233,19 @@ def test_phenology_with_one_parameter_vector(self, param, device): if param == "TEMP": if device == "cuda": pytest.skip("Weather parameter vector tests are CPU-only") - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(10, dtype=torch.float64, device=device) * torch.as_tensor( - wdc.TEMP, dtype=torch.float64, device=device - ) + shape = (10,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) elif param == "DTSMTB": repeated = crop_model_params_provider[param].repeat(10, 1) crop_model_params_provider.set_override(param, repeated, check=False) @@ -244,30 +253,19 @@ def test_phenology_with_one_parameter_vector(self, param, device): repeated = crop_model_params_provider[param].repeat(10) crop_model_params_provider.set_override(param, repeated, check=False) - if param == "TEMP": - with pytest.raises(ValueError): - engine = EngineTestHelper(config=phenology_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - ) - engine.run_till_terminate() - _ = engine.get_output() - else: - engine = EngineTestHelper(config=phenology_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - ) - engine.run_till_terminate() - actual_results = engine.get_output() - expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] + engine = EngineTestHelper(config=phenology_config) + engine.setup( + crop_model_params_provider, + weather_data_provider, + agro_management_inputs, + ) + engine.run_till_terminate() + actual_results = engine.get_output() + expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] - assert len(actual_results) == len(expected_results) - for reference, model in zip(expected_results, actual_results, strict=False): - assert_reference_match(reference, model, expected_precision) + assert len(actual_results) == len(expected_results) + for reference, model in zip(expected_results, actual_results, strict=False): + assert_reference_match(reference, model, expected_precision) @pytest.mark.parametrize( "param,delta", @@ -443,9 +441,6 @@ def test_phenology_with_multiple_parameter_arrays(self, device): repeated = crop_model_params_provider[param].broadcast_to((30, 5)) crop_model_params_provider.set_override(param, repeated, check=False) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones((30, 5), device=device, dtype=torch.float64) * wdc.TEMP - engine = EngineTestHelper(config=phenology_config) engine.setup( crop_model_params_provider, @@ -498,8 +493,8 @@ def test_phenology_with_incompatible_parameter_vectors(self): "TSUM2", crop_model_params_provider["TSUM2"].repeat(5), check=False ) + engine = EngineTestHelper(config=phenology_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=phenology_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -535,11 +530,24 @@ def test_phenology_with_incompatible_weather_parameter_vectors(self): crop_model_params_provider.set_override( "TSUM1", crop_model_params_provider["TSUM1"].repeat(10), check=False ) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(5, dtype=torch.float64) * wdc.TEMP + # Broadcast weather variables to a shape that does not match the parameters + shape = (5,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) + + engine = EngineTestHelper(config=phenology_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=phenology_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -566,7 +574,7 @@ def test_wofost_pp_with_phenology(self, test_data_url, monkeypatch): "VERNDVS", ] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] diff --git a/tests/physical_models/crop/test_respiration.py b/tests/physical_models/crop/test_respiration.py index 0edb19ef..b0cb6ac0 100644 --- a/tests/physical_models/crop/test_respiration.py +++ b/tests/physical_models/crop/test_respiration.py @@ -128,25 +128,25 @@ def test_respiration_with_one_parameter_vector(self, param, device): ) = prepare_engine_input(test_data, crop_model_params, meteo_range_checks=False) if param == "TEMP": - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(10, dtype=torch.float64, device=device) * wdc.TEMP - with pytest.raises(ValueError): - engine = EngineTestHelper(config=respiration_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - _ = engine.get_output() - return - - if param == "RFSETB": + shape = (10,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) + elif param == "RFSETB": repeated = crop_model_params_provider[param].repeat(10, 1) + crop_model_params_provider.set_override(param, repeated, check=False) else: repeated = crop_model_params_provider[param].repeat(10) - crop_model_params_provider.set_override(param, repeated, check=False) + crop_model_params_provider.set_override(param, repeated, check=False) engine = EngineTestHelper(config=respiration_config) engine.setup( @@ -271,9 +271,6 @@ def test_respiration_with_multiple_parameter_arrays(self, device): "RFSETB", crop_model_params_provider["RFSETB"].repeat(30, 5, 1), check=False ) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones((30, 5), dtype=torch.float64, device=device) * wdc.TEMP - engine = EngineTestHelper(config=respiration_config) engine.setup( crop_model_params_provider, @@ -313,8 +310,8 @@ def test_respiration_with_incompatible_parameter_vectors(self): "RML", crop_model_params_provider["RML"].repeat(5), check=False ) + engine = EngineTestHelper(config=respiration_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=respiration_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -336,11 +333,22 @@ def test_respiration_with_incompatible_weather_parameter_vectors(self): crop_model_params_provider.set_override( "RMR", crop_model_params_provider["RMR"].repeat(10), check=False ) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(5, dtype=torch.float64) * wdc.TEMP + shape = (5,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + weather_data_provider = broadcast(weather_data_provider) + + engine = EngineTestHelper(config=respiration_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=respiration_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -354,7 +362,7 @@ def test_wofost_pp_with_respiration(self, test_data_url): test_data = get_test_data(test_data_url) crop_model_params = ["Q10", "RMR", "RML", "RMS", "RMO", "RFSETB"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) # get expected results from YAML test data diff --git a/tests/physical_models/crop/test_root_dynamics.py b/tests/physical_models/crop/test_root_dynamics.py index ecead5d0..3fd237be 100644 --- a/tests/physical_models/crop/test_root_dynamics.py +++ b/tests/physical_models/crop/test_root_dynamics.py @@ -336,8 +336,8 @@ def test_root_dynamics_with_incompatible_parameter_vectors(self, device): "RRI", crop_model_params_provider["RRI"].repeat(5), check=False ) + engine = EngineTestHelper(config=root_dynamics_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=root_dynamics_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -351,7 +351,7 @@ def test_wofost_pp_with_root_dynamics(self, test_data_url): test_data = get_test_data(test_data_url) crop_model_params = ["RDI", "RRI", "RDMCR", "RDMSOL", "TDWI", "IAIRDU", "RDRRTB"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) # get expected results from YAML test data diff --git a/tests/physical_models/crop/test_stem_dynamics.py b/tests/physical_models/crop/test_stem_dynamics.py index 2c5268ff..f1983aa3 100644 --- a/tests/physical_models/crop/test_stem_dynamics.py +++ b/tests/physical_models/crop/test_stem_dynamics.py @@ -49,13 +49,13 @@ def _prepare_common_stem_inputs(test_data_url, device, meteo_range_checks=True): if "RDRSTB" not in crop_model_params_provider: crop_model_params_provider.set_override( "RDRSTB", - torch.tensor([[0.0, 0.0, 2.5, 0.0]], dtype=torch.float64, device=device), + torch.tensor([0.0, 0.0, 2.5, 0.0], dtype=torch.float64, device=device), check=False, ) if "SSATB" not in crop_model_params_provider: crop_model_params_provider.set_override( "SSATB", - torch.tensor([[0.0, 0.0003, 2.5, 0.0003]], dtype=torch.float64, device=device), + torch.tensor([0.0, 0.0003, 2.5, 0.0003], dtype=torch.float64, device=device), check=False, ) if "TDWI" not in crop_model_params_provider: @@ -185,9 +185,20 @@ def test_stem_dynamics_with_one_parameter_vector(self, param, device): # Setting a vector (with one value) for the selected parameter if param == "TEMP": - # Vectorize weather variable - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(10, dtype=torch.float64, device=device) * wdc.TEMP + # Broadcast weather variables + shape = (10,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) else: # Broadcast all parameters to match the batch size of 10 # This ensures compatibility for all parameters including table traits @@ -203,35 +214,21 @@ def test_stem_dynamics_with_one_parameter_vector(self, param, device): p_name, p_val.repeat(10, 1), check=False ) - if param == "TEMP": - # Vectorize weather variable - # We expect the model to handle scalar parameters with vectorized weather - # via implicit broadcasting or explicit checks passing. - engine = EngineTestHelper(config=stem_dynamics_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - actual_results = engine.get_output() - else: - engine = EngineTestHelper(config=stem_dynamics_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - actual_results = engine.get_output() + engine = EngineTestHelper(config=stem_dynamics_config) + engine.setup( + crop_model_params_provider, + weather_data_provider, + agro_management_inputs, + external_states, + ) + engine.run_till_terminate() + actual_results = engine.get_output() - # get expected results from YAML test data - expected_results = test_data["ModelResults"] + # get expected results from YAML test data + expected_results = test_data["ModelResults"] - # Assertions on values removed as test data is not appropriate for this module - assert len(actual_results) == len(expected_results) + # Assertions on values removed as test data is not appropriate for this module + assert len(actual_results) == len(expected_results) @pytest.mark.parametrize( "param,delta", @@ -258,8 +255,7 @@ def test_stem_dynamics_with_different_parameter_values(self, param, delta, devic if param in {"RDRSTB", "SSATB"}: # AfgenTrait parameters need to have shape (N, M) non_zeros_mask = test_value != 0 - # Use cat to get (2, 4) instead of stack (2, 1, 4) - param_vec = torch.cat([test_value + non_zeros_mask * delta, test_value], dim=0) + param_vec = torch.stack([test_value + non_zeros_mask * delta, test_value]) target_batch_size = 2 else: param_vec = torch.tensor( @@ -359,9 +355,6 @@ def test_stem_dynamics_with_multiple_parameter_arrays(self, device): repeated = crop_model_params_provider[param].broadcast_to((30, 5)) crop_model_params_provider.set_override(param, repeated, check=False) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones((30, 5), dtype=torch.float64, device=device) * wdc.TEMP - engine = EngineTestHelper(config=stem_dynamics_config) engine.setup( crop_model_params_provider, @@ -398,8 +391,8 @@ def test_stem_dynamics_with_incompatible_parameter_vectors(self): "RDRSTB", crop_model_params_provider["RDRSTB"].repeat(5, 1), check=False ) + engine = EngineTestHelper(config=stem_dynamics_config) with pytest.raises((AssertionError, ValueError)): - engine = EngineTestHelper(config=stem_dynamics_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -422,11 +415,24 @@ def test_stem_dynamics_with_incompatible_weather_parameter_vectors(self): crop_model_params_provider.set_override( "TDWI", crop_model_params_provider["TDWI"].repeat(10), check=False ) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(5, dtype=torch.float64) * wdc.TEMP + # Broadcast weather variables + shape = (5,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) + + engine = EngineTestHelper(config=stem_dynamics_config) with pytest.raises((AssertionError, ValueError)): - engine = EngineTestHelper(config=stem_dynamics_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -440,7 +446,7 @@ def test_wofost_pp_with_stem_dynamics(self, test_data_url): test_data = get_test_data(test_data_url) crop_model_params = ["TDWI", "RDRSTB", "SSATB"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) # get expected results from YAML test data diff --git a/tests/physical_models/crop/test_storage_organ_dynamics.py b/tests/physical_models/crop/test_storage_organ_dynamics.py index 1d7f5ef7..2a3f2320 100644 --- a/tests/physical_models/crop/test_storage_organ_dynamics.py +++ b/tests/physical_models/crop/test_storage_organ_dynamics.py @@ -181,9 +181,20 @@ def test_storage_dynamics_with_one_parameter_vector(self, param, device): # Setting a vector (with one value) for the selected parameter if param == "TEMP": - # Vectorize weather variable - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(10, dtype=torch.float64, device=device) * wdc.TEMP + # Broadcast weather variable + shape = (10,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) else: # Broadcast all parameters to match the batch size of 10 for p_name in ["TDWI", "SPA"]: @@ -198,35 +209,21 @@ def test_storage_dynamics_with_one_parameter_vector(self, param, device): p_name, p_val.repeat(10, 1), check=False ) - if param == "TEMP": - # Vectorize weather variable - # We expect the model to handle scalar parameters with vectorized weather - # via implicit broadcasting or explicit checks passing. - engine = EngineTestHelper(config=storage_dynamics_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - actual_results = engine.get_output() - else: - engine = EngineTestHelper(config=storage_dynamics_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - actual_results = engine.get_output() + engine = EngineTestHelper(config=storage_dynamics_config) + engine.setup( + crop_model_params_provider, + weather_data_provider, + agro_management_inputs, + external_states, + ) + engine.run_till_terminate() + actual_results = engine.get_output() - # get expected results from YAML test data - expected_results = test_data["ModelResults"] + # get expected results from YAML test data + expected_results = test_data["ModelResults"] - # Assertions on values removed as test data is not appropriate for this module - assert len(actual_results) == len(expected_results) + # Assertions on values removed as test data is not appropriate for this module + assert len(actual_results) == len(expected_results) @pytest.mark.parametrize( "param,delta", @@ -338,9 +335,6 @@ def test_storage_dynamics_with_multiple_parameter_arrays(self, device): repeated = crop_model_params_provider[param].broadcast_to((30, 5)) crop_model_params_provider.set_override(param, repeated, check=False) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones((30, 5), dtype=torch.float64, device=device) * wdc.TEMP - engine = EngineTestHelper(config=storage_dynamics_config) engine.setup( crop_model_params_provider, @@ -377,8 +371,8 @@ def test_storage_dynamics_with_incompatible_parameter_vectors(self): "SPA", crop_model_params_provider["SPA"].repeat(5), check=False ) + engine = EngineTestHelper(config=storage_dynamics_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=storage_dynamics_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -392,7 +386,7 @@ def test_wofost_pp_with_storage_dynamics(self, test_data_url): test_data = get_test_data(test_data_url) crop_model_params = ["TDWI", "SPA"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) # get expected results from YAML test data diff --git a/tests/physical_models/crop/test_wofost72.py b/tests/physical_models/crop/test_wofost72.py index 215b10ca..8e069af2 100644 --- a/tests/physical_models/crop/test_wofost72.py +++ b/tests/physical_models/crop/test_wofost72.py @@ -1,4 +1,3 @@ -import copy import datetime import warnings from unittest.mock import patch @@ -92,7 +91,7 @@ def get_test_diff_wofost72_model(): test_data = get_test_data(test_data_url) _wofost72_template_inputs = _build_wofost72_template_inputs(test_data) (crop_model_params_provider, weather_data_provider, agro_management_inputs, external_states) = ( - copy.deepcopy(_wofost72_template_inputs) + _wofost72_template_inputs ) return DiffWofost72( crop_model_params_provider, @@ -267,14 +266,20 @@ def test_wofost72_with_one_parameter_vector(self, param, device): # Setting a vector (with one value) for the selected parameter if param == "TEMP": - # Vectorize weather variable - for (_, _), wdc in weather_data_provider.store.items(): - base = wdc.TEMP - if isinstance(base, torch.Tensor): - ones = torch.ones(10, dtype=base.dtype, device=base.device) - wdc.TEMP = ones * base - else: - wdc.TEMP = torch.ones(10, dtype=torch.float64) * base + # Broadcast weather variable + shape = (10,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) elif param in ["KDIFTB", "SLATB"]: # AfgenTrait parameters need to have shape (N, M) repeated = crop_model_params_provider[param].repeat(10, 1) @@ -283,48 +288,32 @@ def test_wofost72_with_one_parameter_vector(self, param, device): repeated = crop_model_params_provider[param].repeat(10) crop_model_params_provider.set_override(param, repeated, check=False) - if param == "TEMP": - # Expect error due to incompatible shapes - # (By defaults parameters are not reshaped following weather variables) - with pytest.raises(ValueError): - engine = EngineTestHelper(config=wofost72_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - actual_results = engine.get_output() - else: - engine = EngineTestHelper(config=wofost72_config) - engine.setup( - crop_model_params_provider, - weather_data_provider, - agro_management_inputs, - external_states, - ) - engine.run_till_terminate() - actual_results = engine.get_output() + engine = EngineTestHelper(config=wofost72_config) + engine.setup( + crop_model_params_provider, + weather_data_provider, + agro_management_inputs, + external_states, + ) + engine.run_till_terminate() + actual_results = engine.get_output() - # get expected results from YAML test data - expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] + # get expected results from YAML test data + expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] - assert len(actual_results) == len(expected_results) + assert len(actual_results) == len(expected_results) - for reference, model in zip(expected_results, actual_results, strict=False): - assert reference["DAY"] == model["day"] - # Verify output is on the correct device - for var in expected_precision.keys(): - assert model[var].device.type == device, f"{var} should be on {device}" - # Move to CPU for comparison - model_cpu = { - k: v.cpu() if isinstance(v, torch.Tensor) else v for k, v in model.items() - } - assert all( - all(abs(reference[var] - model_cpu[var]) < precision) - for var, precision in expected_precision.items() - ) + for reference, model in zip(expected_results, actual_results, strict=False): + assert reference["DAY"] == model["day"] + # Verify output is on the correct device + for var in expected_precision.keys(): + assert model[var].device.type == device, f"{var} should be on {device}" + # Move to CPU for comparison + model_cpu = {k: v.cpu() if isinstance(v, torch.Tensor) else v for k, v in model.items()} + assert all( + all(abs(reference[var] - model_cpu[var]) < precision) + for var, precision in expected_precision.items() + ) @pytest.mark.parametrize( "param,delta", @@ -453,14 +442,6 @@ def test_wofost72_with_multiple_parameter_arrays(self, device): repeated = crop_model_params_provider[param].broadcast_to((30, 5)) crop_model_params_provider.set_override(param, repeated, check=False) - for (_, _), wdc in weather_data_provider.store.items(): - base = wdc.TEMP - if isinstance(base, torch.Tensor): - ones = torch.ones((30, 5), dtype=base.dtype, device=base.device) - wdc.TEMP = ones * base - else: - wdc.TEMP = torch.ones((30, 5), dtype=torch.float64) * base - engine = EngineTestHelper(config=wofost72_config) engine.setup( crop_model_params_provider, @@ -507,8 +488,8 @@ def test_wofost72_with_incompatible_parameter_vectors(self): "SPAN", crop_model_params_provider["SPAN"].repeat(5), check=False ) + engine = EngineTestHelper(config=wofost72_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=wofost72_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -532,11 +513,23 @@ def test_wofost72_with_incompatible_weather_parameter_vectors(self): crop_model_params_provider.set_override( "TDWI", crop_model_params_provider["TDWI"].repeat(10), check=False ) - for (_, _), wdc in weather_data_provider.store.items(): - wdc.TEMP = torch.ones(5, dtype=torch.float64) * wdc.TEMP + shape = (5,) + + def broadcast(wdp): + for weather_data in wdp: + out = {} + for k, v in weather_data.items(): + if isinstance(v, torch.Tensor): + out[k] = torch.broadcast_to(v, shape) + else: + out[k] = v + yield out + + weather_data_provider = broadcast(weather_data_provider) + + engine = EngineTestHelper(config=wofost72_config) with pytest.raises(ValueError): - engine = EngineTestHelper(config=wofost72_config) engine.setup( crop_model_params_provider, weather_data_provider, @@ -685,7 +678,7 @@ def test_wofost72_against_pcse_pp(self, test_data_url): test_data = get_test_data(test_data_url) crop_model_params = ["SPAN", "TDWI", "TBASE", "PERDL", "RGRLAI", "KDIFTB", "SLATB"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) # get expected results from YAML test data diff --git a/tests/physical_models/soil/test_waterbalance.py b/tests/physical_models/soil/test_waterbalance.py index 58f478ab..59d4782c 100644 --- a/tests/physical_models/soil/test_waterbalance.py +++ b/tests/physical_models/soil/test_waterbalance.py @@ -291,7 +291,7 @@ def test_wofost72_pp_with_waterbalance(self, test_data_url): test_data = get_test_data(test_data_url) crop_model_params = ["SMFCF"] (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( - prepare_engine_input(test_data, crop_model_params) + prepare_engine_input(test_data, crop_model_params, return_weather_data_provider=True) ) expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] @@ -312,6 +312,38 @@ def test_wofost72_pp_with_waterbalance(self, test_data_url): for var, precision in expected_precision.items() ) + @pytest.mark.parametrize("test_data_url", waterbalance_data_urls) + def test_diffwofost_wofost72_pp_with_waterbalance(self, test_data_url): + """WaterbalancePP with diffWOFOST's Wofost72 reproduces PCSE reference results.""" + test_data = get_test_data(test_data_url) + crop_model_params = ["SMFCF"] + (crop_model_params_provider, weather_data_provider, agro_management_inputs, _) = ( + prepare_engine_input(test_data, crop_model_params) + ) + + expected_results, expected_precision = test_data["ModelResults"], test_data["Precision"] + + waterbalance_config = Configuration( + CROP=Wofost72, + SOIL=WaterbalancePP, + OUTPUT_VARS=[key for key in expected_results[0].keys() if key != "DAY"], + ) + engine = EngineTestHelper(config=waterbalance_config) + engine.setup( + crop_model_params_provider, + weather_data_provider, + agro_management_inputs, + ) + engine.run_till_terminate() + actual_results = engine.get_output() + + for reference, model_out in zip(expected_results, actual_results, strict=False): + assert reference["DAY"] == model_out["day"] + assert all( + abs(reference[var] - model_out[var]) < precision + for var, precision in expected_precision.items() + ) + @pytest.mark.usefixtures("fast_mode") class TestDiffWaterbalancePPGradients: diff --git a/tests/physical_models/test_utils.py b/tests/physical_models/test_utils.py index ef70bebe..67d5badb 100644 --- a/tests/physical_models/test_utils.py +++ b/tests/physical_models/test_utils.py @@ -4,15 +4,10 @@ import pytest import torch from diffwofost.physical_models.config import ComputeConfig -from diffwofost.physical_models.test import WeatherDataProviderTestHelper -from diffwofost.physical_models.test import get_test_data -from diffwofost.physical_models.test import prepare_engine_input from diffwofost.physical_models.utils import Afgen from diffwofost.physical_models.utils import AfgenTrait -from diffwofost.physical_models.utils import _get_drv from diffwofost.physical_models.utils import astro from diffwofost.physical_models.utils import daylength -from . import phy_data_folder ComputeConfig.set_dtype(torch.float64) DTYPE = ComputeConfig.get_dtype() @@ -711,91 +706,6 @@ def test_batched_gradient_at_boundaries(self): assert torch.isclose(x2.grad, torch.tensor(2.0, dtype=DTYPE), atol=1e-5) -@pytest.mark.usefixtures("fast_mode") -class TestGetDrvParam: - """Tests for _get_drv function.""" - - def test_prepare_engine_input_tensorizes_weather_when_checks_disabled(self): - test_data_url = f"{phy_data_folder}/test_leafdynamics_wofost72_05.yaml" - test_data = get_test_data(test_data_url) - - _, weather_provider, _, external_states = prepare_engine_input( - test_data, - ["RGRLAI"], - meteo_range_checks=False, - ) - - weather = weather_provider(weather_provider.first_date) - assert isinstance(weather.TEMP, torch.Tensor) - assert isinstance(weather.IRRAD, torch.Tensor) - assert external_states - assert isinstance(external_states[0]["DVS"], torch.Tensor) - - def test_prepare_engine_input_returns_empty_external_states_when_missing(self): - test_data_url = f"{phy_data_folder}/test_leafdynamics_wofost72_05.yaml" - test_data = get_test_data(test_data_url) - test_data = dict(test_data) - test_data.pop("ExternalStates", None) - - _, _, _, external_states = prepare_engine_input(test_data, ["RGRLAI"]) - - assert external_states == [] - - def test_weather_provider_does_not_mutate_input_weather(self): - test_data_url = f"{phy_data_folder}/test_phenology_wofost72_05.yaml" - test_data = get_test_data(test_data_url) - weather_inputs = test_data["WeatherVariables"] - - assert any("SNOWDEPTH" in item for item in weather_inputs) - - WeatherDataProviderTestHelper(weather_inputs) - - assert any("SNOWDEPTH" in item for item in weather_inputs) - - def test_float_broadcast(self): - expected_shape = (3, 2) - test_data_url = f"{phy_data_folder}/test_leafdynamics_wofost72_05.yaml" - test_data = get_test_data(test_data_url) - provider = WeatherDataProviderTestHelper(test_data["WeatherVariables"]) - wdc = provider(provider.first_date) - scalar = wdc.TEMP - out = _get_drv(scalar, expected_shape, dtype=DTYPE) - assert out.shape == expected_shape - assert torch.allclose(out, torch.full(expected_shape, scalar, dtype=DTYPE)) - - def test_scalar_broadcast(self): - expected_shape = (3, 2) - test_data_url = f"{phy_data_folder}/test_leafdynamics_wofost72_05.yaml" - test_data = get_test_data(test_data_url) - provider = WeatherDataProviderTestHelper(test_data["WeatherVariables"]) - wdc = provider(provider.first_date) - scalar = torch.tensor(wdc.IRRAD, dtype=DTYPE) # 0-d tensor - out = _get_drv(scalar, expected_shape, dtype=DTYPE) - assert out.shape == expected_shape - assert torch.allclose(out, torch.full(expected_shape, scalar.item(), dtype=DTYPE)) - - def test_matching_shape_pass_through(self): - expected_shape = (3, 2) - base_val = torch.tensor(12.34, dtype=DTYPE) - var = torch.ones(expected_shape, dtype=DTYPE) * base_val - out = _get_drv(var, expected_shape, dtype=DTYPE) - assert out.shape == expected_shape - # Should be the same object (no copy) - assert out.data_ptr() == var.data_ptr() - - def test_wrong_shape_raises(self): - expected_shape = (3, 2) - wrong = torch.ones(2, 3, dtype=DTYPE) - with pytest.raises(ValueError, match="incompatible shape"): - _get_drv(wrong, expected_shape, dtype=DTYPE) - - def test_one_dim_shape_raises(self): - expected_shape = (3, 2) - one_dim = torch.ones(3, dtype=DTYPE) - with pytest.raises(ValueError, match="incompatible shape"): - _get_drv(one_dim, expected_shape, dtype=DTYPE) - - # --------------------------------------------------------------------------- # Shared helpers for astro / daylength tests # --------------------------------------------------------------------------- diff --git a/tests/physical_models/test_weather.py b/tests/physical_models/test_weather.py new file mode 100644 index 00000000..46d9a278 --- /dev/null +++ b/tests/physical_models/test_weather.py @@ -0,0 +1,261 @@ +import datetime +import types +import numpy as np +import pandas as pd +import pytest +import torch +import xarray as xr +from diffwofost.physical_models.weather import to_weather_data_iterator + + +class TestToWeatherDataIteratorFromDataFrame: + def test_returns_iterator_of_weather_variables(self): + weather_data = pd.DataFrame( + { + "DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"], + "TEMP": [10.0, 11.0, 9.0, 12.0], + } + ) + weather_data_iter = to_weather_data_iterator(weather_data) + + assert isinstance(weather_data_iter, types.GeneratorType) + first = next(weather_data_iter) + assert first["DAY"] == datetime.date(2020, 4, 1) + assert torch.equal(first["TEMP"], torch.tensor(10.0)) + + def test_raises_on_values_outside_valid_range(self): + weather_data_faulty = pd.DataFrame( + { + "DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"], + "TEMP": [10.0, 1000.0, 9.0, 12.0], # unrealistic temperature + } + ) + with pytest.raises(ValueError, match="outside the range"): + to_weather_data_iterator(weather_data_faulty) + + def test_succeed_if_values_outside_valid_range_and_check_disabled(self): + weather_data_faulty = pd.DataFrame( + { + "DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"], + "TEMP": [10.0, 1000.0, 9.0, 12.0], # unrealistic temperature + } + ) + weather_data_iter = to_weather_data_iterator(weather_data_faulty, check=False) + assert isinstance(weather_data_iter, types.GeneratorType) + + def test_raises_on_wrong_date_intervals(self): + weather_data_faulty = pd.DataFrame( + { + "DAY": ["2020-04-01", "2020-04-03", "2020-04-04", "2020-04-05"], # one day missing + "TEMP": [10.0, 11.0, 9.0, 12.0], + } + ) + with pytest.raises(ValueError, match="consecutive daily dates"): + to_weather_data_iterator(weather_data_faulty) + + def test_succeed_if_wrong_date_intervals_and_check_disabled(self): + weather_data_faulty = pd.DataFrame( + { + "DAY": ["2020-04-01", "2020-04-03", "2020-04-04", "2020-04-05"], # one day missing + "TEMP": [10.0, 11.0, 9.0, 12.0], + } + ) + weather_data_iter = to_weather_data_iterator(weather_data_faulty, check=False) + assert isinstance(weather_data_iter, types.GeneratorType) + + def test_nan_values_are_allowed_by_default(self): + weather_data_with_nan = pd.DataFrame( + { + "DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"], + "TEMP": [10.0, 11.0, 9.0, np.nan], + } + ) + weather_data_iter = to_weather_data_iterator(weather_data_with_nan) + assert isinstance(weather_data_iter, types.GeneratorType) + + def test_raises_on_nan_values_if_skipna(self): + weather_data_with_nan = pd.DataFrame( + { + "DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"], + "TEMP": [10.0, 11.0, 9.0, np.nan], + } + ) + with pytest.raises(ValueError, match="NaN"): + to_weather_data_iterator(weather_data_with_nan, skipna=False) + + def test_day_column_is_optional(self): + weather_data = pd.DataFrame({"TEMP": [10.0, 11.0, 9.0, 12.0]}) + weather_data_iter = to_weather_data_iterator(weather_data) + first = next(weather_data_iter) + assert len(first) == 1 + assert "TEMP" in first + + +class TestToWeatherDataIteratorFromDataset: + def test_raises_if_no_key_is_recognized(self): + weather_data = xr.Dataset( + { + "MYTEMP": ("DAY", [10.0, 11.0, 9.0, 12.0]), + "MYRAIN": ("DAY", [0.0, 1.0, 24.0, 1.0]), + }, + coords={"DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"]}, + ) + with pytest.raises(ValueError): + to_weather_data_iterator(weather_data) + + def test_returns_iterator_of_weather_variables(self): + weather_data = xr.Dataset( + {"TEMP": ("DAY", [10.0, 11.0, 9.0, 12.0])}, + coords={"DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"]}, + ) + weather_data_iter = to_weather_data_iterator(weather_data) + + assert isinstance(weather_data_iter, types.GeneratorType) + first = next(weather_data_iter) + assert first["DAY"] == datetime.date(2020, 4, 1) + assert torch.equal(first["TEMP"], torch.tensor(10.0)) + + def test_returns_iterator_inferring_dimension_names_from_day_coords(self): + # the time dimension will be identified from the "DAY" variable or coordinate + days = ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"] + dim_name = "my_time" + weather_data = xr.Dataset( + { + "TEMP": ( + ("location", dim_name), + [ + [10.0, 11.0, 9.0, 12.0], + [11.0, 7.0, 8.0, 13.0], + [9.0, 3.0, 11.0, 11.0], + ], + ) + }, + coords={"DAY": (dim_name, days)}, + ) + weather_data_iter = to_weather_data_iterator(weather_data) + assert isinstance(weather_data_iter, types.GeneratorType) + + weather_data_list = list(weather_data_iter) + assert len(weather_data_list) == len(days) + first = weather_data_list[0] + assert first["DAY"] == datetime.date(2020, 4, 1) + assert torch.equal(first["TEMP"], torch.tensor([10.0, 11.0, 9.0])) + + @pytest.mark.parametrize("dim_name", ["time", "day", "dates"]) + def test_returns_iterator_with_known_dimension_names(self, dim_name): + # if the DAY variable/coordinate is missing, the time dimension can be recognized from some + # default names + weather_data = xr.Dataset( + { + "TEMP": ( + ("location", dim_name), + [ + [10.0, 11.0, 9.0, 12.0], + [11.0, 7.0, 8.0, 13.0], + [9.0, 3.0, 11.0, 11.0], + ], + ) + }, + ) + weather_data_iter = to_weather_data_iterator(weather_data) + + assert isinstance(weather_data_iter, types.GeneratorType) + weather_data_list = list(weather_data_iter) + assert len(weather_data_list) == 4 + first = weather_data_list[0] + assert torch.equal(first["TEMP"], torch.tensor([10.0, 11.0, 9.0])) + + def test_raises_if_time_dimension_not_recognized(self): + # if the dimension name is not known and the "DAY" variable or coordinate is not present, + # an error will be raised + weather_data = xr.Dataset( + { + "TEMP": ( + ("location", "my_time_dim"), + [ + [10.0, 11.0, 9.0, 12.0], + [11.0, 7.0, 8.0, 13.0], + [9.0, 3.0, 11.0, 11.0], + ], + ) + }, + ) + with pytest.raises(ValueError): + to_weather_data_iterator(weather_data) + + def test_returns_iterator_if_time_dimension_specified_on_input(self): + dim_name = "my_time_dim" + weather_data = xr.Dataset( + { + "TEMP": ( + ("location", dim_name), + [ + [10.0, 11.0, 9.0, 12.0], + [11.0, 7.0, 8.0, 13.0], + [9.0, 3.0, 11.0, 11.0], + ], + ) + }, + ) + weather_data_iter = to_weather_data_iterator(weather_data, time_dim=dim_name) + assert isinstance(weather_data_iter, types.GeneratorType) + + weather_data_list = list(weather_data_iter) + assert len(weather_data_list) == 4 + first = weather_data_list[0] + assert torch.equal(first["TEMP"], torch.tensor([10.0, 11.0, 9.0])) + + def test_dataset_is_transposed_if_needed(self): + dim_name = "my_time_dim" + weather_data = xr.Dataset( + { + "TEMP": ( + (dim_name, "location"), + [ + [10.0, 11.0, 9.0, 12.0], + [11.0, 7.0, 8.0, 13.0], + [9.0, 3.0, 11.0, 11.0], + ], + ) + }, + ) + weather_data_iter = to_weather_data_iterator(weather_data, time_dim=dim_name) + assert isinstance(weather_data_iter, types.GeneratorType) + + weather_data_list = list(weather_data_iter) + assert len(weather_data_list) == 3 + first = weather_data_list[0] + assert torch.equal(first["TEMP"], torch.tensor([10.0, 11.0, 9.0, 12.0])) + + def test_raises_on_values_outside_valid_range(self): + weather_data_faulty = xr.Dataset( + {"TEMP": ("DAY", [10.0, 1000.0, 9.0, 12.0])}, # unrealistic temperature + coords={"DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"]}, + ) + with pytest.raises(ValueError, match="outside the range"): + to_weather_data_iterator(weather_data_faulty) + + def test_raises_on_wrong_date_intervals(self): + weather_data_faulty = xr.Dataset( + {"TEMP": ("DAY", [10.0, 11.0, 9.0, 12.0])}, + coords={ + "DAY": ["2020-04-01", "2020-04-03", "2020-04-04", "2020-04-05"] + }, # one day missing + ) + with pytest.raises(ValueError, match="consecutive daily dates"): + to_weather_data_iterator(weather_data_faulty) + + def test_raises_on_nan_values_if_skipna(self): + weather_data_with_nan = xr.Dataset( + {"TEMP": ("DAY", [10.0, 11.0, 9.0, np.nan])}, + coords={"DAY": ["2020-04-01", "2020-04-02", "2020-04-03", "2020-04-04"]}, + ) + with pytest.raises(ValueError, match="NaN"): + to_weather_data_iterator(weather_data_with_nan, skipna=False) + + def test_day_coordinate_is_optional(self): + weather_data = xr.Dataset({"TEMP": ("DAY", [10.0, 11.0, 9.0, 12.0])}) + weather_data_iter = to_weather_data_iterator(weather_data) + first = next(weather_data_iter) + assert len(first) == 1 + assert "TEMP" in first