{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":81000,"databundleVersionId":8812083,"sourceType":"competition"}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# install pyro (PPL)\n!mamba install -y --quiet -c conda-forge pyro-ppl","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:14.392824Z","iopub.execute_input":"2024-07-30T12:39:14.393390Z","iopub.status.idle":"2024-07-30T12:39:36.900851Z","shell.execute_reply.started":"2024-07-30T12:39:14.393351Z","shell.execute_reply":"2024-07-30T12:39:36.899315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport torch\nimport pyro\nimport pyro.distributions as dist\nimport pyro.nn as nn\n\nfrom tqdm import tqdm\n\n# plotting\nimport matplotlib.pyplot as plt\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-30T12:39:36.902964Z","iopub.execute_input":"2024-07-30T12:39:36.903389Z","iopub.status.idle":"2024-07-30T12:39:36.915480Z","shell.execute_reply.started":"2024-07-30T12:39:36.903352Z","shell.execute_reply":"2024-07-30T12:39:36.913958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!git clone https://github.com/ataraxno/SIMPLE_crop_model","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:36.917269Z","iopub.execute_input":"2024-07-30T12:39:36.917642Z","iopub.status.idle":"2024-07-30T12:39:38.094752Z","shell.execute_reply.started":"2024-07-30T12:39:36.917611Z","shell.execute_reply":"2024-07-30T12:39:38.093271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fraction of plant available water-holding capacity (AWC; one number for entire soil profile, limited by potential root depth)\n# runoff curve number (RCN)\n# deep drainage coefficient (DDC)\n# root zone depth (RZD, a fixed maximum depth)\ncrop_params = pd.read_csv('SIMPLE_crop_model/SIMPLE/params/crop_params.csv', index_col='Crop').reset_index()\nsoil_params = pd.read_csv('SIMPLE_crop_model/SIMPLE/params/soil_params.csv', index_col='Crop').reset_index()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:38.097094Z","iopub.execute_input":"2024-07-30T12:39:38.097616Z","iopub.status.idle":"2024-07-30T12:39:38.139041Z","shell.execute_reply.started":"2024-07-30T12:39:38.097566Z","shell.execute_reply":"2024-07-30T12:39:38.137855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"crop_params.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:38.140477Z","iopub.execute_input":"2024-07-30T12:39:38.140843Z","iopub.status.idle":"2024-07-30T12:39:38.172961Z","shell.execute_reply.started":"2024-07-30T12:39:38.140812Z","shell.execute_reply":"2024-07-30T12:39:38.171807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"soil_params.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:38.174329Z","iopub.execute_input":"2024-07-30T12:39:38.174705Z","iopub.status.idle":"2024-07-30T12:39:38.194161Z","shell.execute_reply.started":"2024-07-30T12:39:38.174675Z","shell.execute_reply":"2024-07-30T12:39:38.192883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_data(crop: str, mode: str=\"train\"):\n    # note that years represent an offset from model spinup;\n    # soil co2 dataset has real year;\n    # 0-30 are days before sowing, 31-238 are days after sowing\n    tasmax = pd.read_parquet(f\"/kaggle/input/the-future-crop-challenge/tasmax_{crop}_{mode}.parquet\")\n    tasmin = pd.read_parquet(f\"/kaggle/input/the-future-crop-challenge/tasmin_{crop}_{mode}.parquet\")\n    pr = pd.read_parquet(f\"/kaggle/input/the-future-crop-challenge/pr_{crop}_{mode}.parquet\")\n    rsds = pd.read_parquet(f\"/kaggle/input/the-future-crop-challenge/rsds_{crop}_{mode}.parquet\")\n    soil_co2 = pd.read_parquet(f\"/kaggle/input/the-future-crop-challenge/soil_co2_{crop}_{mode}.parquet\")\n    target = pd.read_parquet(f\"/kaggle/input/the-future-crop-challenge/{mode}_solutions_{crop}.parquet\") if mode == \"train\" else None\n    return {\n        'tasmax': tasmax.drop([\"crop\",\"variable\"], axis=1).reset_index().set_index([\"lat\",\"lon\"]).sort_index(),\n        'tasmin': tasmin.drop([\"crop\",\"variable\"], axis=1).reset_index().set_index([\"lat\",\"lon\"]).sort_index(),\n        'pr': pr.drop([\"crop\",\"variable\"], axis=1).reset_index().set_index([\"lat\",\"lon\"]).sort_index(),\n        'rsds': rsds.drop([\"crop\",\"variable\"], axis=1).reset_index().set_index([\"lat\",\"lon\"]).sort_index(),\n        'soil_co2': soil_co2.drop([\"crop\"], axis=1).reset_index().set_index([\"lat\",\"lon\"]).sort_index(),\n        'target': target,\n    }","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:38.195761Z","iopub.execute_input":"2024-07-30T12:39:38.196949Z","iopub.status.idle":"2024-07-30T12:39:38.208756Z","shell.execute_reply.started":"2024-07-30T12:39:38.196905Z","shell.execute_reply":"2024-07-30T12:39:38.207691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This code is modified from the SIMPLE crop model:\n# https://github.com/ataraxno/SIMPLE_crop_model/blob/main/SIMPLE/crop.py\n# by T Moon, GHPF Lab of SNU.\n# Zhao C, Liu B, Xiao L, Hoogenboom G, Boote KJ, Kassie BT, Pavan W, Shelia V, Kim KS, Hernandez-Ochoa IM et al. (2019) A SIMPLE crop model. Eur J Agron 104:97-106\n#\n# which is licensed as free and open source software under the Apache v2 license:\n# https://www.apache.org/licenses/LICENSE-2.0\n#\n# All modifications are subject to copyright by Brian Groenke of UFZ Leipzig (c) 2024. All rights reserved.\n\nclass Forcings:\n    def __init__(self, tasmax, tasmin, precip, rad, co2, nitrogen, device=torch.get_default_device()):\n        \"\"\"Forcing data for the crop model.\n        \n        tasmax: array - Maximum daily temperature time series (deg C)\n        tasmin: array - Minimum daily temperature time series (deg C)\n        precip: array - Daily total precipitation (mm)\n        rad: array - Average daily incoming radiation (MJ m^-2 d^-1)\n        co2: array - CO2 concentration (ppm)\n        nitrogen: array - Nitrogen concentration (ppm); currently not used!\n        \n        \"\"\"\n        assert len(tasmax.shape) == len(tasmin.shape) == len(precip.shape) == len(rad.shape) == len(co2.shape) == 2, \\\n            \"forcing input tensors should be T x N where N is the batch size and T is the number of time steps\"\n        assert tasmax.shape == tasmin.shape == precip.shape == rad.shape == co2.shape, \\\n            \"all forcing input tensors must have matching shapes\"\n        self.num_timesteps = tasmax.shape[0]\n        self.tasmax = torch.Tensor(tasmax, device=device)\n        self.tasmin = torch.Tensor(tasmin, device=device)\n        self.precip = torch.Tensor(precip, device=device)\n        self.rad = torch.Tensor(rad, device=device)\n        self.co2 = torch.Tensor(co2, device=device)\n        # currently unused\n        self.nitrogen = torch.Tensor(nitrogen, device=device)\n\n# Probabilistic SIMPLE crop model implemented via pyro/pytorch.\nclass ProbCrop(nn.PyroModule):\n    def __init__(self, N: int, crop_params, soil_params, t0=0, co2_min=350.0, co2_max=700.0,\n                 fSolar_max=0.8, rwu_alpha=0.096, device=torch.get_default_device()):\n        \"\"\"Initializes the SIMPLE crop model with `N` independent states where `N` should match\n        the dimensionality of the forcing data at each timestep. `crop_params` and `soil_params`\n        are as documented in the original SIMPLE model implementation and are directly loaded from\n        the parameter files.\n        \"\"\"\n        super().__init__()\n        crop_params = torch.Tensor(np.array(crop_params), device=device)\n        soil_params = torch.Tensor(np.array(soil_params), device=device)\n        co2_min = torch.Tensor(np.array(co2_min), device=device)\n        co2_max = torch.Tensor(np.array(co2_max), device=device)\n        \n        # Crop parameters\n        ## Cumulative temperature requirement from sowing to maturity (degC d)\n        self.T_sum        = nn.PyroParam(crop_params[0], constraint=dist.constraints.positive)\n        ## Potential harvest index.\n        self.HI           = nn.PyroParam(crop_params[1], constraint=dist.constraints.unit_interval)\n        ## Cumulative temperature requirement for leaf area development to intercept 50% of radiation (degC d).\n        self.I_50A        = nn.PyroParam(crop_params[2], constraint=dist.constraints.positive)\n        ## Cumulative temperature till maturity to reach 50% radiation interception due to leaf senescence (degC d).\n        self.I_50B_0      = nn.PyroParam(crop_params[3], constraint=dist.constraints.positive)\n        self.T_base       = crop_params[4]\n        self.T_opt        = crop_params[5]\n        self.RUE          = crop_params[6]\n        self.I_50maxH     = crop_params[7]\n        self.I_50maxW     = crop_params[8]\n        self.T_heat       = crop_params[9]\n        self.T_ext        = crop_params[10]\n        self.S_CO2        = crop_params[11]\n        self.S_water      = crop_params[12]\n        self.fSolar_max   = torch.Tensor(np.array(fSolar_max))\n        self.dr           = torch.Tensor()\n        self.rwu_alpha    = torch.Tensor(np.array(rwu_alpha))\n        \n        # Soil parameters\n        self.AWC          = soil_params[0]\n        self.RCN          = soil_params[1]\n        self.DDC          = soil_params[2]\n        self.RZD          = soil_params[3]\n        \n        # Additional CO2 thresholds\n        self.co2_min      = co2_min\n        self.co2_max      = co2_max\n        \n        # Init state\n        self.initialize(N, t0)\n        \n    def initialize(self, N: int, t0=0):\n        # State variables\n        self.t            = t0 # days after sowing, planting, ... whatever.\n        self.ET_0         = torch.zeros(N)\n        self.PAW          = torch.ones(N)*self.RZD*self.AWC/1000 # plant-available water storage\n        self.I_50B        = torch.ones(N)*self.I_50B_0\n        self.TT           = torch.zeros(N) # cumulative mean temperature\n        self.biomass_cum  = torch.zeros(N) # cumulative biomass\n        self.yields       = torch.zeros(N) # yields\n        \n    def phenology(self, T_max, T_min):\n        T_mean = (T_max + T_min) / 2\n#         if T_mean > self.T_base:\n#             dTT = T_mean - self.T_base\n#         else:\n#             dTT = 0\n        dTT = torch.where(T_mean > self.T_base, T_mean - self.T_base, 0.0)\n        self.TT += dTT\n        return self.TT\n        \n        \n    def growth(self, T_max, T_min, rad, CO2):\n        T_mean = (T_max + T_min)/2\n        # fSolar calculation\n        fSolar = torch.min(\n            self.fSolar_max/(1 + torch.exp(-0.01*(self.TT - self.I_50A))),\n            self.fSolar_max/(1 + torch.exp(0.01*(self.TT - (self.T_sum - self.I_50B))))\n        )\n        torch._assert(~torch.isnan(fSolar).any(), f\"NaN detected in fSolar at t={self.t}\")\n        \n        # f(Temp) calculation\n#         if T_mean < self.T_base:\n#             fTemp = 0\n#         elif T_mean >= self.T_base and T_mean < self.T_opt:\n#             fTemp =\n#         else: # T_mean >= self.T_opt\n#             fTemp = 1\n        fTemp = torch.where(\n            T_mean < self.T_base,\n            0.0,\n            torch.where(\n                torch.logical_and(T_mean >= self.T_base, T_mean < self.T_opt),\n                (T_mean - self.T_base)/(self.T_opt - self.T_base),\n                1.0,\n            )\n        )\n        torch._assert(~torch.isnan(fTemp).any(), f\"NaN detected in fTemp at t={self.t}\")\n        \n        # f(heat) calculation\n#         if T_max <= self.T_heat:\n#             fHeat = 1\n#         elif T_max > self.T_heat and T_max <= self.T_ext:\n#             fHeat = 1 - (T_max - self.T_heat)/(self.T_ext - self.T_heat)\n#         else: # T_max > self.T_ext\n#             fHeat = 0\n        fHeat = torch.where(\n            T_max <= self.T_heat,\n            1.0,\n            torch.where(\n                torch.logical_and(T_max > self.T_heat, T_max <= self.T_ext),\n                1 - (T_max - self.T_heat)/(self.T_ext - self.T_heat),\n                0.0,\n            )\n        )\n        torch._assert(~torch.isnan(fHeat).any(), f\"NaN detected in fHeat at t={self.t}\")\n                \n        # f(CO2) calculation\n#         if CO2 >= 350 and CO2 < 700:\n#             fCO2 = 1 + self.S_CO2*(CO2 - 350)\n#         elif CO2 > 700:\n#             fCO2 = 1 + self.S_CO2*350\n        fCO2 = torch.where(\n            CO2 <= self.co2_min,\n            1.0, # this case is not handled in the original code but should be...\n            torch.where(\n                torch.logical_and(CO2 >= self.co2_min, CO2 < self.co2_max),\n                1 + self.S_CO2*(CO2 - self.co2_min),\n                1 + self.S_CO2*self.co2_min, # TODO: double check this, it doesn't make sense\n            )\n        )\n        torch._assert(~torch.isnan(fCO2).any(), f\"NaN detected in fCO2 at t={self.t}\")\n            \n        # f(Water) calculation\n        # note the zero ET_0 case is not handled in the original code; here we opt to just assume non-arid\n        # conditioins when ET_0 = 0 since ET cannot occur (?) under these circumstances.\n        ARID = 1 - torch.where(self.ET_0 > 0, torch.min(self.ET_0, self.rwu_alpha*self.PAW)/self.ET_0, 1.0)\n#         ARID = 0 # No information about ET_0 and PAW\n        fWater = 1 - self.S_water*ARID\n        torch._assert(~torch.isnan(fWater).any(), f\"NaN detected in fWater at t={self.t}\")\n        \n        # Updating I_50B with f(Heat) and f(Water)\n        self.I_50B += self.I_50maxH*(1 - fHeat)\n        self.I_50B += self.I_50maxW*(1 - fWater)\n        \n        biomass_rate = rad*fSolar*self.RUE*fCO2*fTemp*torch.min(fHeat, fWater)\n        self.biomass_cum += biomass_rate\n\n        return fSolar, self.biomass_cum\n    \n    def hydrology(self, T_max, T_min, prec, rad):\n        # Simplified hydrology scheme, neglecting irrigation\n        # see https://etcalc.hydrotools.tech/pageMain.php for ET calculations\n        # via the Priestley-Taylor method.\n        # first convert to meters\n        rzd = self.RZD/1000\n        prec = prec/1000\n        min_water = self.AWC*rzd # minimum water content\n        max_water = rzd # maximum water content\n        # compute drainage and runoff\n        drainage = self.DDC*rzd*(self.PAW/rzd - self.AWC)\n        runoff = prec**2 / (prec + (25400/self.RCN - 254)/1000)\n        # convert potential ET\n        T_air = (T_max + T_min)/2\n        vp_slope = torch.where(T_air > 0, 0.3221*torch.exp(0.0803*T_air**0.8876), 0.3405*torch.exp(0.0642*T_air))\n        self.ET_0 = 1.3 / (2260*1000)*vp_slope / (vp_slope + 4.95e-4)*rad\n        ET = torch.min(self.rwu_alpha*self.PAW, self.ET_0)\n        # update water balance\n        new_PAW = self.PAW + prec - ET - drainage - runoff\n#         print(f\"prec={prec.mean()} runoff={runoff.mean()} ET={ET.mean()} D={drainage.mean()}\")\n        self.PAW = torch.min(max_water, torch.max(min_water, new_PAW))\n        torch._assert(~torch.isnan(self.PAW).any(), f\"NaN detected in PAW at t={self.t}\")\n        torch._assert(~torch.isnan(self.ET_0).any(), f\"NaN detected in ET_0 at t={self.t}\")\n        return self.PAW, self.ET_0\n    \n    def step(self, forcings: Forcings):\n        # inputs\n        T_max = forcings.tasmax\n        T_min = forcings.tasmin\n        prec = forcings.precip\n        rad = forcings.rad\n        CO2 = forcings.co2\n        t = self.t\n        # update\n        TT_t = self.phenology(T_max[t,:], T_min[t,:])\n        PAW, ET_0 = self.hydrology(T_max[t,:], T_min[t,:], prec[t,:], rad[t,:])\n        fSolar, biomass_cum = self.growth(T_max[t,:], T_min[t,:], rad[t,:], CO2[t,:])\n        self.t += 1\n        return TT_t, fSolar, biomass_cum, PAW\n    \n    \n    def run(self, forcings: Forcings, t0=0):\n        # reset states\n        self.initialize(self.TT.shape[0], t0)\n        \n        # state storage\n        TTs = []\n        solar = []\n        biomass = []\n        PAW = []\n        \n        for _ in range(forcings.num_timesteps):\n            TT_t, fSolar, biomass_cum, PAW_t = self.step(forcings)\n            TTs.append(TT_t)\n            solar.append(fSolar)\n            biomass.append(biomass_cum)\n            PAW.append(PAW_t)\n        \n        # factor of 1/100 converts from g/m^2 to tonne/ha;\n        # g/m^2 x 1e4 m^2 / ha x 1 tonne / 1e6 g\n        self.yields = self.biomass_cum*self.HI/100\n        \n        return self.yields, TTs, biomass, solar, PAW","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:38.210644Z","iopub.execute_input":"2024-07-30T12:39:38.211088Z","iopub.status.idle":"2024-07-30T12:39:38.264194Z","shell.execute_reply.started":"2024-07-30T12:39:38.211057Z","shell.execute_reply":"2024-07-30T12:39:38.262805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_params_for(crop: str, crop_params, soil_params):\n    # for now we'll just arbitrarily select the first available option for the crop type;\n    # ideally this should be done more intelligently, or perhaps built into a mixture prior\n    crop_params_sel = crop_params.loc[crop_params.Crop == crop].iloc[0,2:].astype(np.float32)\n    soil_params_sel = soil_params.loc[soil_params.Crop == crop].iloc[0,5:].astype(np.float32)\n    return crop_params_sel, soil_params_sel\n\ndef set_up_model(data, crop_params, soil_params, timestep_offset=31):\n    model = ProbCrop(data[\"soil_co2\"].shape[0], crop_params_wheat.values, soil_params_wheat.values)\n    num_timesteps = data[\"tasmax\"].shape[1] - timestep_offset - 2\n    forcings = Forcings(\n        data[\"tasmax\"].drop(columns=[\"ID\",\"year\"]).values.T[timestep_offset:,:],\n        data[\"tasmin\"].drop(columns=[\"ID\",\"year\"]).values.T[timestep_offset:,:],\n        # convert precipitation to mm\n        data[\"pr\"].drop(columns=[\"ID\",\"year\"]).values.T[timestep_offset:,:]*24*3600,\n        # convert radiation from W/m^2 to MJ/m^2/day\n        data[\"rsds\"].drop(columns=[\"ID\",\"year\"]).values.T[timestep_offset:,:]*24*3600/1e6,\n        data[\"soil_co2\"].co2.values.reshape((1,-1))*np.ones((num_timesteps,1)),\n        data[\"soil_co2\"].nitrogen.values.reshape((1,-1))*np.ones((num_timesteps,1)),\n    )\n    return model, forcings","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:38.268612Z","iopub.execute_input":"2024-07-30T12:39:38.269077Z","iopub.status.idle":"2024-07-30T12:39:38.283528Z","shell.execute_reply.started":"2024-07-30T12:39:38.269040Z","shell.execute_reply":"2024-07-30T12:39:38.281954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_train = load_data(\"wheat\", \"train\")","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:39:38.285171Z","iopub.execute_input":"2024-07-30T12:39:38.285660Z","iopub.status.idle":"2024-07-30T12:40:05.200208Z","shell.execute_reply.started":"2024-07-30T12:39:38.285623Z","shell.execute_reply":"2024-07-30T12:40:05.198863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"crop_params_wheat, soil_params_wheat = extract_params_for(\"Wheat\", crop_params, soil_params)\ncrop_params_wheat","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:40:05.201808Z","iopub.execute_input":"2024-07-30T12:40:05.202265Z","iopub.status.idle":"2024-07-30T12:40:05.217215Z","shell.execute_reply.started":"2024-07-30T12:40:05.202228Z","shell.execute_reply":"2024-07-30T12:40:05.215840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"soil_params_wheat","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:40:05.218946Z","iopub.execute_input":"2024-07-30T12:40:05.219383Z","iopub.status.idle":"2024-07-30T12:40:05.232253Z","shell.execute_reply.started":"2024-07-30T12:40:05.219336Z","shell.execute_reply":"2024-07-30T12:40:05.230803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_train_co2 = wheat_train[\"soil_co2\"]\nwheat_train_co2.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:40:05.233742Z","iopub.execute_input":"2024-07-30T12:40:05.234124Z","iopub.status.idle":"2024-07-30T12:40:05.257048Z","shell.execute_reply.started":"2024-07-30T12:40:05.234093Z","shell.execute_reply":"2024-07-30T12:40:05.255582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_model, wheat_forcings = set_up_model(wheat_train, crop_params_wheat, soil_params_wheat)","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:47:56.229493Z","iopub.execute_input":"2024-07-30T12:47:56.229976Z","iopub.status.idle":"2024-07-30T12:47:59.770044Z","shell.execute_reply.started":"2024-07-30T12:47:56.229922Z","shell.execute_reply":"2024-07-30T12:47:59.769094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_model.step(wheat_forcings)\n# yields, TTs, biomass, solar, PAW = wheat_model.run(wheat_forcings)","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:40:09.921095Z","iopub.execute_input":"2024-07-30T12:40:09.921550Z","iopub.status.idle":"2024-07-30T12:40:10.022118Z","shell.execute_reply.started":"2024-07-30T12:40:09.921510Z","shell.execute_reply":"2024-07-30T12:40:10.020985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define loss and optimize\nparams = list(wheat_model.parameters())\nloss_fn = torch.nn.MSELoss(reduction='sum')\noptim = torch.optim.Adam(wheat_model.parameters(), lr=1e-4)\ntrue_yields = torch.Tensor(wheat_train[\"target\"][\"yield\"].values)\nnum_iter = 10\nparams[0].retain_grad()\nfor i in tqdm(range(num_iter)):\n    yields, TTs, biomass, solar, PAW = wheat_model.run(wheat_forcings)\n    loss = loss_fn(yields, true_yields)\n    # initialize gradients to zero\n    optim.zero_grad()\n    # backpropagate\n    loss.backward()\n    # seems to be an issue with gradient overflow for some parameters\n    print([p.grad for p in params])\n    # take a gradient step\n    optim.step()\n    print(list(wheat_model.parameters()))","metadata":{"execution":{"iopub.status.busy":"2024-07-30T12:47:59.771827Z","iopub.execute_input":"2024-07-30T12:47:59.772970Z","iopub.status.idle":"2024-07-30T12:48:09.255783Z","shell.execute_reply.started":"2024-07-30T12:47:59.772927Z","shell.execute_reply":"2024-07-30T12:48:09.254112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maize_train = load_data(\"maize\", \"train\")\ncrop_params_maize, soil_params_maize = extract_params_for(\"Maize\", crop_params, soil_params)","metadata":{},"execution_count":null,"outputs":[]}]}