{"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":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# install additional dependencies for probabilistic modeling and geostats\n!mamba install -y -c conda-forge pymc bambi arviz geodatasets","metadata":{"execution":{"iopub.status.busy":"2024-07-29T14:55:12.332717Z","iopub.execute_input":"2024-07-29T14:55:12.333460Z","iopub.status.idle":"2024-07-29T14:57:39.275473Z","shell.execute_reply.started":"2024-07-29T14:55:12.333425Z","shell.execute_reply":"2024-07-29T14:57:39.274238Z"},"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 xarray as xr\n\nimport cartopy\nimport geopandas as gpd\nimport matplotlib.pyplot as plt\nfrom geodatasets import get_path\n\nworldmap = gpd.read_file(get_path(\"naturalearth.land\"))\n\nimport itertools\n\nimport re, gc\nimport warnings\nwarnings.filterwarnings(\"ignore\")\npd.set_option('display.max_columns', None)\ngc.enable()\n\n# modeling\nimport pymc as pm\nimport bambi as bmb\n\nfrom tqdm import tqdm\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-29T14:57:39.277591Z","iopub.execute_input":"2024-07-29T14:57:39.277936Z","iopub.status.idle":"2024-07-29T14:57:45.005214Z","shell.execute_reply.started":"2024-07-29T14:57:39.277905Z","shell.execute_reply":"2024-07-29T14:57:45.004103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import geopandas as geo\n\ndef 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    }\n\ndef to_xarray(df):\n    melted_df = df.melt(id_vars=[\"year\",\"lon\",\"lat\"], var_name=\"timestep\")\n    melted_df[\"timestep\"] = melted_df[\"timestep\"].astype(float)\n    melted_df['loc'] = melted_df[['lat','lon']].agg(tuple, axis=1)\n    return xr.Dataset.from_dataframe(melted_df.drop(columns=['lat','lon']).set_index(['year','loc','timestep']))\n\ndef standardize(x, axis=0, shift=True):\n    shift = x.mean(axis=0) if shift else 0.0\n    scale = x.std(axis=0)\n    return (x - shift) / scale, shift, scale","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:09:49.422773Z","iopub.execute_input":"2024-07-29T12:09:49.423430Z","iopub.status.idle":"2024-07-29T12:09:49.440327Z","shell.execute_reply.started":"2024-07-29T12:09:49.423396Z","shell.execute_reply":"2024-07-29T12:09:49.439196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_train = load_data(\"wheat\", \"train\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:09:49.443461Z","iopub.execute_input":"2024-07-29T12:09:49.443834Z","iopub.status.idle":"2024-07-29T12:11:00.842668Z","shell.execute_reply.started":"2024-07-29T12:09:49.443805Z","shell.execute_reply":"2024-07-29T12:11:00.840655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ax = worldmap.plot(color=\"white\", edgecolor=\"black\")\ndf = wheat_train[\"soil_co2\"].reset_index()\ngdf = geo.GeoDataFrame(df, geometry=geo.points_from_xy(df.lon, df.lat), crs=\"EPSG:4326\")\np1 = gdf.plot(ax=ax, column=\"texture_class\", markersize=0.1, alpha=0.5, legend=True, cmap='gist_earth', legend_kwds={'orientation': 'horizontal'})\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:11:00.844900Z","iopub.execute_input":"2024-07-29T12:11:00.845443Z","iopub.status.idle":"2024-07-29T12:11:54.365057Z","shell.execute_reply.started":"2024-07-29T12:11:00.845394Z","shell.execute_reply":"2024-07-29T12:11:54.363715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def temperature_features(var, name):\n    \"\"\"Reduces temeprature time series to pre and post sowing feature set.\"\"\"\n    # drop ID and year\n    var = var.drop([\"ID\",\"year\"], axis=1).astype(np.float32)\n    # get remaining column names (time indices)\n    cols = var.columns\n    # pre- and post-sowing periods\n    var_pre = var.iloc[:,:30]\n    var_post = var.iloc[:,30:]\n    # mean/min/max\n    var_pre_mean = var_pre.mean(axis=1)\n    var_pre_max = var_pre.max(axis=1)\n    var_pre_min = var_pre.min(axis=1)\n    # degree-days: freezing for tasmin and thawing for tasmax\n    var_pre_dd = var_pre.where(var_pre > 0).sum(axis=1) if name is \"tasmax\" else var_pre.where(var_pre <= 0).sum(axis=1)\n    # again for post-sowing period\n    var_post_mean = var_post.mean(axis=1)\n    var_post_max = var_post.max(axis=1)\n    var_post_min = var_post.min(axis=1)\n    var_post_dd = var_post.where(var_post > 0).sum(axis=1) if name is \"tasmax\" else var_post.where(var_post <= 0).sum(axis=1)\n    features = var.drop(cols, axis=1)\n    features[f\"{name}_pre_mean\"] = var_pre_mean\n    features[f\"{name}_pre_min\"] = var_pre_min\n    features[f\"{name}_pre_max\"] = var_pre_max\n    features[f\"{name}_pre_dd\"] = var_pre_dd\n    features[f\"{name}_post_mean\"] = var_post_mean\n    features[f\"{name}_post_min\"] = var_post_min\n    features[f\"{name}_post_max\"] = var_post_max\n    features[f\"{name}_post_dd\"] = var_post_dd\n    return features\n\ndef precip_features(var, name):\n    \"\"\"Reduces precipitation time series to pre and post sowing features.\"\"\"\n    # drop ID and year and rescale to meters\n    var = var.drop(columns=[\"ID\",\"year\"]).astype(np.float32)\n    # get remaining column names (time indices)\n    cols = var.columns\n    # compute features\n    var_pre = var.iloc[:,:30]\n    var_post = var.iloc[:,30:]\n    var_pre_sum = var_pre.sum(axis=1)\n    var_pre_sdii = var_pre.where(var_pre > 0).mean(axis=1).fillna(0.0)\n    var_post_sum = var_post.sum(axis=1).fillna(0.0)\n    var_post_sdii = var_post.where(var_post > 0).mean(axis=1).fillna(0.0)\n    features = var.drop(cols, axis=1)\n    features[f\"{name}_pre_sum\"] = var_pre_sum\n    features[f\"{name}_pre_sdii\"] = var_pre_sdii\n    features[f\"{name}_post_sum\"] = var_post_sum\n    features[f\"{name}_post_sdii\"] = var_post_sdii\n    return features\n\ndef rad_features(var, name):\n    \"\"\"Reduces radiation time series to pre and post sowing feature set.\"\"\"\n    # drop ID and year\n    var = var.drop([\"ID\",\"year\"], axis=1).astype(np.float32)\n    # get remaining column names (time indices)\n    cols = var.columns\n    # compute features\n    var_pre = var.iloc[:,:30]\n    var_post = var.iloc[:,30:]\n    var_pre_mean = var_pre.mean(axis=1)\n    var_pre_max = var_pre.max(axis=1)\n    var_pre_min = var_pre.min(axis=1)\n    var_post_mean = var_post.mean(axis=1)\n    var_post_max = var_post.max(axis=1)\n    var_post_min = var_post.min(axis=1)\n    features = var.drop(cols, axis=1)\n    features[f\"{name}_pre_mean\"] = var_pre_mean\n    features[f\"{name}_pre_min\"] = var_pre_min\n    features[f\"{name}_pre_max\"] = var_pre_max\n    features[f\"{name}_post_mean\"] = var_post_mean\n    features[f\"{name}_post_min\"] = var_post_min\n    features[f\"{name}_post_max\"] = var_post_max\n    return features\n\ndef lat_lon_features(data):\n    lat_lon = data.reset_index()[[\"lat\",\"lon\"]]\n    # convert lat/lon to cartesian coordinates which should theoretically have better geometric properties\n    lat_lon[\"x\"] = np.cos(lat_lon.lat*np.pi/180)*np.cos(lat_lon.lon*np.pi/180)\n    lat_lon[\"y\"] = np.cos(lat_lon.lat*np.pi/180)*np.sin(lat_lon.lon*np.pi/180)\n    lat_lon[\"z\"] = np.sin(lat_lon.lat*np.pi/180)\n    return lat_lon.set_index([\"lat\",\"lon\"])\n\ndef get_features(data):\n    # temperature\n    tasmax_features = temperature_features(data[\"tasmax\"], \"tasmax\")\n    tasmin_features = temperature_features(data[\"tasmin\"], \"tasmin\")\n    # precipitation\n    pr_features = precip_features(data[\"pr\"], \"pr\")\n    # solar radiation\n    rsds_features = rad_features(data[\"rsds\"], \"rsds\")\n    # co2 and nitrogen\n    co2_features = data[\"soil_co2\"][[\"co2\",\"nitrogen\"]]\n    # spatial coordinates\n    coord_features = lat_lon_features(data[\"soil_co2\"])\n    return pd.concat([co2_features, coord_features, tasmax_features, tasmin_features, pr_features, rsds_features], axis=1)\n\ndef batch_by_loc(df, batch_size=200):\n    \"\"\"Batches the given data frame by lat/lon coordinates. Returns an array of arrays where each element is a set of indices\n    corresponding to a btach of locations of size `batch_size`.\n    \"\"\"\n    grouped_df = df.groupby(['lat','lon'])\n    loc2idx = grouped_df.indices\n    batches = []\n    indices = list(loc2idx.values())\n    for st in range(0, len(indices), batch_size):\n        en = st + batch_size\n        batches.append(np.concatenate(indices[st:en]))\n    return batches","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:11:54.366976Z","iopub.execute_input":"2024-07-29T12:11:54.367438Z","iopub.status.idle":"2024-07-29T12:11:54.402692Z","shell.execute_reply.started":"2024-07-29T12:11:54.367398Z","shell.execute_reply":"2024-07-29T12:11:54.401241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# join with soil co2 on ID to match order of predictors\ny_train = wheat_train[\"soil_co2\"].reset_index().merge(wheat_train[\"target\"].reset_index(), on=\"ID\")[[\"lat\",\"lon\",\"real_year\",\"yield\"]]\ny_train.describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:11:54.403991Z","iopub.execute_input":"2024-07-29T12:11:54.404407Z","iopub.status.idle":"2024-07-29T12:11:54.603713Z","shell.execute_reply.started":"2024-07-29T12:11:54.404365Z","shell.execute_reply":"2024-07-29T12:11:54.602431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.hist(y_train[\"yield\"], bins=100);","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:11:54.605410Z","iopub.execute_input":"2024-07-29T12:11:54.605879Z","iopub.status.idle":"2024-07-29T12:11:55.067393Z","shell.execute_reply.started":"2024-07-29T12:11:54.605840Z","shell.execute_reply":"2024-07-29T12:11:55.065903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pymc_var_datasets(data, container_type=pm.MutableData, **kwargs):\n    tasmax = data[\"tasmax\"].drop(columns=[\"ID\",\"year\"])\n    tasmax_data = container_type(\"tasmax\", tasmax.values, dims=('obs','day'), **kwargs)\n    tasmin = data[\"tasmin\"].drop(columns=[\"ID\",\"year\"])\n    tasmin_data = container_type(\"tasmin\", tasmin.values, dims=('obs','day'), **kwargs)\n    prec = data[\"pr\"].drop(columns=[\"ID\",\"year\"])\n    prec_data = container_type(\"prec\", prec.values, dims=('obs','day'), **kwargs)\n    rsds = data[\"rsds\"].drop(columns=[\"ID\",\"year\"])\n    rsds_data = container_type(\"rsds\", rsds.values, dims=('obs','day'), **kwargs)\n    co2_data = container_type(\"co2\", data[\"soil_co2\"].co2.values, dims=('obs',), **kwargs)\n    nitrogen_data = container_type(\"nitrogen\", data[\"soil_co2\"].nitrogen.values, dims=('obs',), **kwargs)\n    texture_data = container_type(\"texture\", data[\"soil_co2\"].texture_class.values, dims=('obs',), **kwargs)\n    return {'tasmax': tasmax_data, 'tasmin': tasmin_data, 'prec': prec_data, 'rsds': rsds_data, 'co2': co2_data, 'nitrogen': nitrogen_data, 'texture': texture_data}\n\ndef mlr_formula(X: pd.DataFrame, y: pd.Series):\n    covariates = X.columns.values\n    return f\"{y.name} ~ 1 + \" + \" + \".join(covariates), covariates\n\ndef glm_model(X: pd.DataFrame, y: pd.Series, family=\"guassian\", **kwargs):\n    X = X.drop(columns=[\"lat\",\"lon\",\"year\"])\n    X_std, shift, scale = standardize(X)\n    # drop columns with nan/inf values (i.e. zero variance); otherwise bambi will complain\n    X_std = X_std.replace((np.inf,-np.inf), np.nan).dropna(axis=1, how=\"any\")\n    # combine X and y into single dataframe\n    train_data = pd.concat([X_std, y], axis=1)\n    # generate simple multiple linear regression formula from term names\n    formula, varnames = mlr_formula(X_std, y)\n    # build model\n    model = bmb.Model(formula, train_data, family=family, **kwargs)\n    # set mildly informative priors that should promote reasonable-ish predictions;\n    # half-normal on the intercept enforces strictly positive values and the effect prior scale of 10/k should constrain the effect sizes\n    # based on the number of features.\n    intercept_prior = bmb.Prior(\"HalfNormal\", sigma=10.0) if family is \"gaussian\" else bmb.Prior(\"Normal\", sigma=2.0)\n    model.set_priors(priors={'Intercept': intercept_prior}, common=bmb.Prior(\"Normal\",mu=0,sigma=10/len(varnames)))\n    return model, X_std, shift, scale","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:53:33.682702Z","iopub.execute_input":"2024-07-29T12:53:33.683168Z","iopub.status.idle":"2024-07-29T12:53:33.703431Z","shell.execute_reply.started":"2024-07-29T12:53:33.683131Z","shell.execute_reply":"2024-07-29T12:53:33.702174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train = get_features(wheat_train)\nX_train[\"year\"] = wheat_train[\"soil_co2\"].real_year\nX_train = X_train.reset_index()\nX_train.describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:11:55.093052Z","iopub.execute_input":"2024-07-29T12:11:55.094561Z","iopub.status.idle":"2024-07-29T12:12:08.461015Z","shell.execute_reply.started":"2024-07-29T12:11:55.094520Z","shell.execute_reply":"2024-07-29T12:12:08.459734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_glm_pre, _, _, _ = glm_model(X_train.iloc[:1000], y_train[\"yield\"].iloc[:1000], family=\"gaussian\")\nwheat_glm_pre.build()\nwheat_glm_pre","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:28:47.899766Z","iopub.execute_input":"2024-07-29T12:28:47.900289Z","iopub.status.idle":"2024-07-29T12:28:48.075214Z","shell.execute_reply.started":"2024-07-29T12:28:47.900250Z","shell.execute_reply":"2024-07-29T12:28:48.073776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here we do a \"prior predictive check\" to assess whether the model *a priori* produces reasonable-ish predictions.\n\nThis will take a couple of minutes to run due to the size of the dataset.\n\nNote that the model produces negative yield values. This is due to the unconstrained Gaussian likelihood. It's not ideal, but it's a reasonable first order approximation, and it is much easier to fit than nonnegative likelihoods such as the hurdle Gamma.\n\nYou can change to a (theoretically) more appropriate GLM family above by setting `family = \"hurdle_gamma\"` or `family = \"hurdle_lognormal\"`.","metadata":{}},{"cell_type":"code","source":"idata_prior_wheat = wheat_glm_pre.prior_predictive(draws=100)\n# check prior predictive samples to see if predicted yields fall in a reasonable range\nidata_prior_wheat.prior_predictive[\"yield\"].to_dataframe().reset_index().drop(columns=[\"chain\",\"draw\",\"__obs__\"]).describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:28:53.434592Z","iopub.execute_input":"2024-07-29T12:28:53.435030Z","iopub.status.idle":"2024-07-29T12:28:55.987915Z","shell.execute_reply.started":"2024-07-29T12:28:53.434999Z","shell.execute_reply":"2024-07-29T12:28:55.986560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_glm, X_train_std, wheat_scale, wheat_shift = glm_model(X_train, y_train[\"yield\"], family=\"gaussian\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:53:54.850979Z","iopub.execute_input":"2024-07-29T12:53:54.851450Z","iopub.status.idle":"2024-07-29T12:53:55.438006Z","shell.execute_reply.started":"2024-07-29T12:53:54.851413Z","shell.execute_reply":"2024-07-29T12:53:55.436869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" # due to the size of the dataset, we will fit the model using a fast variational inference (VI) approximation\napprox_posterior_wheat = wheat_glm.fit(method=\"vi\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:37:55.092377Z","iopub.execute_input":"2024-07-29T12:37:55.092866Z","iopub.status.idle":"2024-07-29T12:40:39.931763Z","shell.execute_reply.started":"2024-07-29T12:37:55.092829Z","shell.execute_reply":"2024-07-29T12:40:39.930533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## alternatively, we can use a batched training strategy where train a series of models on different batches of sites.\n# num_iter = 10_000\n# trained_models_wheat = []\n# batches = batch_by_loc(X_train, batch_size=100)\n# for i, idx in enumerate(batches):\n#     X_i = X_train.iloc[idx].drop(columns=['x','y','z'])\n#     y_i = y_train.iloc[idx]['yield']\n#     glm, shift, scale = glm_model(X_i, y_i, family=\"gaussian\")\n#     print(f\"Training GLM on batch {i+1}/{len(batches)} with {X_i.shape[0]} samples\")\n#     res = glm.fit(num_iter, inference_method=\"vi\", method=\"advi\")\n#     trained_models_wheat.append((idx, glm, res, shift, scale))","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:43:36.139481Z","iopub.execute_input":"2024-07-29T12:43:36.140019Z","iopub.status.idle":"2024-07-29T12:43:36.150269Z","shell.execute_reply.started":"2024-07-29T12:43:36.139981Z","shell.execute_reply":"2024-07-29T12:43:36.148953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"approx_posterior_samples = approx_posterior_wheat.sample()\napprox_posterior_samples","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:43:36.463648Z","iopub.execute_input":"2024-07-29T12:43:36.464096Z","iopub.status.idle":"2024-07-29T12:43:41.990184Z","shell.execute_reply.started":"2024-07-29T12:43:36.464059Z","shell.execute_reply":"2024-07-29T12:43:41.988836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_glm.predict(approx_posterior_samples, kind=\"response\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:44:15.453092Z","iopub.execute_input":"2024-07-29T12:44:15.453545Z","iopub.status.idle":"2024-07-29T12:44:25.952210Z","shell.execute_reply.started":"2024-07-29T12:44:15.453510Z","shell.execute_reply":"2024-07-29T12:44:25.950911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"posterior_yield_mean = approx_posterior_samples.posterior_predictive[\"yield\"].mean(dim=['chain','draw']).values\nposterior_yield_std = approx_posterior_samples.posterior_predictive[\"yield\"].std(dim=['chain','draw']).values","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:44:35.651500Z","iopub.execute_input":"2024-07-29T12:44:35.654146Z","iopub.status.idle":"2024-07-29T12:44:44.241935Z","shell.execute_reply.started":"2024-07-29T12:44:35.654093Z","shell.execute_reply":"2024-07-29T12:44:44.240248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred = y_train.reset_index().copy()\ny_pred[\"year\"] = X_train.reset_index().year\ny_pred[\"pred_mean\"] = posterior_yield_mean\ny_pred[\"pred_std\"] = posterior_yield_std\ny_pred","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:44:44.244894Z","iopub.execute_input":"2024-07-29T12:44:44.245471Z","iopub.status.idle":"2024-07-29T12:44:44.315686Z","shell.execute_reply.started":"2024-07-29T12:44:44.245416Z","shell.execute_reply":"2024-07-29T12:44:44.314233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred.describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:45:50.467038Z","iopub.execute_input":"2024-07-29T12:45:50.467702Z","iopub.status.idle":"2024-07-29T12:45:50.616936Z","shell.execute_reply.started":"2024-07-29T12:45:50.467666Z","shell.execute_reply":"2024-07-29T12:45:50.615578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred.groupby(\"year\")[\"yield\"].mean().plot(label=\"true yield\")\ny_pred.groupby(\"year\")[\"pred_mean\"].mean().plot(label=\"predicted\")\nplt.legend()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:44:47.543760Z","iopub.execute_input":"2024-07-29T12:44:47.544216Z","iopub.status.idle":"2024-07-29T12:44:47.976953Z","shell.execute_reply.started":"2024-07-29T12:44:47.544180Z","shell.execute_reply":"2024-07-29T12:44:47.975714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred.groupby(['lat','lon']).apply(lambda x: (x.pred_mean - x[\"yield\"]).abs().mean()).reset_index().rename(columns={0: \"error\"}).plot.scatter(x=\"lon\",y=\"lat\",c=\"error\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:44:59.097870Z","iopub.execute_input":"2024-07-29T12:44:59.098278Z","iopub.status.idle":"2024-07-29T12:45:02.860068Z","shell.execute_reply.started":"2024-07-29T12:44:59.098249Z","shell.execute_reply":"2024-07-29T12:45:02.858649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred[:300].groupby(['lat','lon'])[[\"pred_mean\",\"yield\"]].plot()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:45:09.848992Z","iopub.execute_input":"2024-07-29T12:45:09.849430Z","iopub.status.idle":"2024-07-29T12:45:12.444858Z","shell.execute_reply.started":"2024-07-29T12:45:09.849397Z","shell.execute_reply":"2024-07-29T12:45:12.443631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_test = load_data(\"wheat\", \"test\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:47:38.121513Z","iopub.execute_input":"2024-07-29T12:47:38.121989Z","iopub.status.idle":"2024-07-29T12:48:48.461747Z","shell.execute_reply.started":"2024-07-29T12:47:38.121957Z","shell.execute_reply":"2024-07-29T12:48:48.460467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_test = get_features(wheat_test)\nX_test_std = ((X_test - wheat_shift) / wheat_scale)[X_train_std.columns]\nX_test_std.describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T12:54:16.594631Z","iopub.execute_input":"2024-07-29T12:54:16.595079Z","iopub.status.idle":"2024-07-29T12:54:42.619112Z","shell.execute_reply.started":"2024-07-29T12:54:16.595044Z","shell.execute_reply":"2024-07-29T12:54:42.617934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_glm.predict(approx_posterior_samples, kind=\"response\", data=X_test_std)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wheat_pred = wheat_test[\"soil_co2\"][[\"ID\"]].reset_index().drop(columns=[\"lat\",\"lon\"])\nwheat_pred[\"yield\"] = approx_posterior_samples.posterior_predictive[\"yield\"].mean(dim=[\"chain\",\"draw\"]).values","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:00:03.025492Z","iopub.execute_input":"2024-07-29T13:00:03.026570Z","iopub.status.idle":"2024-07-29T13:00:04.996569Z","shell.execute_reply.started":"2024-07-29T13:00:03.026522Z","shell.execute_reply":"2024-07-29T13:00:04.994793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# explicitly delete variables to avoid running out of memory on Kaggle\ndel wheat_glm, wheat_train, wheat_test, X_train, X_train_std, y_train, y_pred, X_test, X_test_std, approx_posterior_samples","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:01:21.015785Z","iopub.execute_input":"2024-07-29T13:01:21.016824Z","iopub.status.idle":"2024-07-29T13:01:21.025785Z","shell.execute_reply.started":"2024-07-29T13:01:21.016773Z","shell.execute_reply":"2024-07-29T13:01:21.024458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect(2)","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:07:19.993385Z","iopub.execute_input":"2024-07-29T13:07:19.994626Z","iopub.status.idle":"2024-07-29T13:07:21.463974Z","shell.execute_reply.started":"2024-07-29T13:07:19.994582Z","shell.execute_reply":"2024-07-29T13:07:21.462465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maize_train = load_data(\"maize\", \"train\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:03:54.204298Z","iopub.execute_input":"2024-07-29T13:03:54.205495Z","iopub.status.idle":"2024-07-29T13:04:26.328556Z","shell.execute_reply.started":"2024-07-29T13:03:54.205448Z","shell.execute_reply":"2024-07-29T13:04:26.327076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train = get_features(maize_train)\nX_train[\"year\"] = maize_train[\"soil_co2\"].real_year\nX_train = X_train.reset_index()\nX_train.describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:04:35.126568Z","iopub.execute_input":"2024-07-29T13:04:35.127053Z","iopub.status.idle":"2024-07-29T13:04:49.074034Z","shell.execute_reply.started":"2024-07-29T13:04:35.127018Z","shell.execute_reply":"2024-07-29T13:04:49.072695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# join with soil co2 on ID to match order of predictors\ny_train = maize_train[\"soil_co2\"].reset_index().merge(maize_train[\"target\"].reset_index(), on=\"ID\")[[\"lat\",\"lon\",\"real_year\",\"yield\"]]\ny_train.describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:05:22.792126Z","iopub.execute_input":"2024-07-29T13:05:22.792654Z","iopub.status.idle":"2024-07-29T13:05:22.974015Z","shell.execute_reply.started":"2024-07-29T13:05:22.792596Z","shell.execute_reply":"2024-07-29T13:05:22.972349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maize_glm, X_train_std, maize_shift, maize_scale = glm_model(X_train, y_train[\"yield\"], family=\"gaussian\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:11:04.595209Z","iopub.execute_input":"2024-07-29T13:11:04.595711Z","iopub.status.idle":"2024-07-29T13:11:05.288651Z","shell.execute_reply.started":"2024-07-29T13:11:04.595664Z","shell.execute_reply":"2024-07-29T13:11:05.287273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"approx_posterior_maize = maize_glm.fit(method=\"vi\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:07:33.338523Z","iopub.execute_input":"2024-07-29T13:07:33.338971Z","iopub.status.idle":"2024-07-29T13:11:04.494783Z","shell.execute_reply.started":"2024-07-29T13:07:33.338936Z","shell.execute_reply":"2024-07-29T13:11:04.493368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"approx_posterior_samples = approx_posterior_maize.sample()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:11:23.167044Z","iopub.execute_input":"2024-07-29T13:11:23.167562Z","iopub.status.idle":"2024-07-29T13:11:29.913198Z","shell.execute_reply.started":"2024-07-29T13:11:23.167522Z","shell.execute_reply":"2024-07-29T13:11:29.911909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maize_test = load_data(\"maize\", \"test\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:11:29.915871Z","iopub.execute_input":"2024-07-29T13:11:29.916348Z","iopub.status.idle":"2024-07-29T13:12:50.568218Z","shell.execute_reply.started":"2024-07-29T13:11:29.916292Z","shell.execute_reply":"2024-07-29T13:12:50.566829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_test = get_features(maize_test)\nX_test_std = ((X_test - maize_shift) / maize_scale)[X_train_std.columns]\nX_test_std.describe()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:12:50.569864Z","iopub.execute_input":"2024-07-29T13:12:50.570228Z","iopub.status.idle":"2024-07-29T13:13:20.485636Z","shell.execute_reply.started":"2024-07-29T13:12:50.570197Z","shell.execute_reply":"2024-07-29T13:13:20.484424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del maize_train, X_train, X_train_std\ngc.collect(2)","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:13:43.729896Z","iopub.execute_input":"2024-07-29T13:13:43.730259Z","iopub.status.idle":"2024-07-29T13:13:45.649644Z","shell.execute_reply.started":"2024-07-29T13:13:43.730230Z","shell.execute_reply":"2024-07-29T13:13:45.648395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maize_glm.predict(approx_posterior_samples, kind=\"response\", data=X_test_std)","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:13:20.487974Z","iopub.execute_input":"2024-07-29T13:13:20.488383Z","iopub.status.idle":"2024-07-29T13:13:43.728467Z","shell.execute_reply.started":"2024-07-29T13:13:20.488321Z","shell.execute_reply":"2024-07-29T13:13:43.727231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maize_pred = maize_test[\"soil_co2\"][[\"ID\"]].reset_index().drop(columns=[\"lat\",\"lon\"])\nmaize_pred[\"yield\"] = approx_posterior_samples.posterior_predictive[\"yield\"].mean(dim=[\"chain\",\"draw\"]).values","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:13:56.075737Z","iopub.execute_input":"2024-07-29T13:13:56.076191Z","iopub.status.idle":"2024-07-29T13:14:04.188646Z","shell.execute_reply.started":"2024-07-29T13:13:56.076137Z","shell.execute_reply":"2024-07-29T13:14:04.186497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_pred = pd.concat([maize_pred, wheat_pred], axis=0)\nall_pred.set_index(\"ID\").to_csv(\"/kaggle/working/submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T13:14:04.192496Z","iopub.execute_input":"2024-07-29T13:14:04.193589Z","iopub.status.idle":"2024-07-29T13:14:09.266934Z","shell.execute_reply.started":"2024-07-29T13:14:04.193537Z","shell.execute_reply":"2024-07-29T13:14:09.265164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}