{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":81000,"databundleVersionId":8812083,"sourceType":"competition"}],"dockerImageVersionId":30747,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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)\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-08-07T07:01:37.904022Z","iopub.execute_input":"2024-08-07T07:01:37.904427Z","iopub.status.idle":"2024-08-07T07:01:37.914168Z","shell.execute_reply.started":"2024-08-07T07:01:37.904367Z","shell.execute_reply":"2024-08-07T07:01:37.913101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load package and fix random seed","metadata":{}},{"cell_type":"code","source":"import os\nimport copy\nimport random\nimport torch\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.autograd import Variable\nimport torch.nn.functional as F\nfrom sklearn.preprocessing import StandardScaler\ndevice = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nDATA_DIR = r'/kaggle/input/the-future-crop-challenge'\nprint(f'Running on {device}')\n\ndef setup_seed(seed):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.backends.cudnn.deterministic = True\n    return\n\nsetup_seed(0)","metadata":{"execution":{"iopub.status.busy":"2024-08-07T07:01:37.916090Z","iopub.execute_input":"2024-08-07T07:01:37.916427Z","iopub.status.idle":"2024-08-07T07:01:37.927424Z","shell.execute_reply.started":"2024-08-07T07:01:37.916392Z","shell.execute_reply":"2024-08-07T07:01:37.926432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define dataset\nThese code are based on [Example LSTM+FCN in PyTorch](https://www.kaggle.com/code/lobsterjesus/example-lstm-fcn-in-pytorch)\n\nInput feature includes **tasmax, tasmin, pr, rsds, cumulative rsds，lon, lat, co2, nitrogen and texture_class**\n\nThe **texture_class** is converted to a one-hot vector with **13 classes**\n\nTime series data from the planting date (the 30th day) are used for input.\n\nThe raw climate data is also returned for the next step of temperature sum","metadata":{}},{"cell_type":"code","source":"class ClimateDataset(Dataset):\n    \"\"\"\n    The ClimateDataset class provides a convenient way for acessing, merging and scaling parquet data. \n    This is used in the main training loop to access features and target variables.\n    \"\"\"\n    def __init__(self, crop: str, mode: str, data_dir: str, scalers: list = None):\n        self.tasmax = pd.read_parquet(os.path.join(data_dir, f\"tasmax_{crop}_{mode}.parquet\"))\n        self.tasmin = pd.read_parquet(os.path.join(data_dir, f\"tasmin_{crop}_{mode}.parquet\"))\n        # self.tas = pd.read_parquet(os.path.join(data_dir, f\"tas_{crop}_{mode}.parquet\"))\n        self.pr = pd.read_parquet(os.path.join(data_dir, f\"pr_{crop}_{mode}.parquet\"))\n        self.rsds = pd.read_parquet(os.path.join(data_dir, f\"rsds_{crop}_{mode}.parquet\"))\n        self.soil_co2 = pd.read_parquet(os.path.join(data_dir, f\"soil_co2_{crop}_{mode}.parquet\"))\n        \n        if mode == 'train':\n            self.yield_ = pd.read_parquet(os.path.join(data_dir, f\"{mode}_solutions_{crop}.parquet\"))\n        else:\n            self.yield_ = None\n\n        if scalers is None:\n            self._init_scalers()\n        else:\n            self.scaler_climate, self.scaler_soil, self.scaler_yield = scalers\n        self._check_data([self.tasmax, self.tasmin, self.pr, self.rsds, self.soil_co2], self.yield_)\n\n    def __getitem__(self, index):\n        # 240x4 climate matrix per location/year (features in last dimension by convention) + cumulative rsds\n        climate = np.vstack([\n            self.tasmax.iloc[index, 35:].astype(np.float32), \n            self.tasmin.iloc[index, 35:].astype(np.float32),\n            # self.tas.iloc[index, 35:].astype(np.float32),\n            self.pr.iloc[index, 35:].astype(np.float32),\n            self.rsds.iloc[index, 35:].astype(np.float32),\n            np.cumsum(self.rsds.iloc[index, 35:]).astype(np.float32),\n        ]).T\n\n        # Fixed soil properties per location/year\n        soil = self.soil_co2.iloc[index][['lon','lat','co2', 'nitrogen']].astype(np.float32)\n        texture = self.soil_co2.iloc[index][['texture_class']].values.astype(np.int64)\n        id = soil.name\n        soil = soil.values\n\n        # Yield estimated by process model\n        if self.yield_ is not None:\n            yield_ = self.yield_.iloc[index].astype(np.float32).values\n        else:\n            yield_ = []\n\n        climate_scaler = self.scaler_climate.transform(climate)\n        soil = self.scaler_soil.transform(soil.reshape(1, -1)).reshape(-1)\n        texture = F.one_hot(torch.from_numpy((texture-1).astype(np.int64)),13)[0].to(torch.float32)\n        soil_texture = torch.concat([torch.tensor(soil),texture])\n        # print(yield_)\n        return torch.tensor(climate), torch.tensor(climate_scaler), soil_texture, torch.tensor(yield_), id\n\n    def __len__(self):\n        return self.tasmax.shape[0]\n\n    def _init_scalers(self):\n        # Draw random sample from climate data to estimate distribution moments for scaler.\n        climate_sample = np.vstack([\n            self.tasmax.sample(1000).iloc[:, 5:].values.flatten(), \n            self.tasmin.sample(1000).iloc[:, 5:].values.flatten(),\n            # self.tas.sample(1000).iloc[:, 5:].values.flatten(),\n            self.pr.sample(1000).iloc[:, 5:].values.flatten(),\n            self.rsds.sample(1000).iloc[:, 5:].values.flatten(),\n            np.cumsum(self.rsds.sample(1000).iloc[:, 5:],1).values.flatten(),\n        ]).T\n        self.scaler_climate = StandardScaler()\n        self.scaler_climate.fit(climate_sample)\n        # Scaler for fixed soil properties\n        self.scaler_soil = StandardScaler()\n        self.scaler_soil.fit(self.soil_co2[['lon','lat','co2', 'nitrogen']].values)\n\n        # Scaler for yield\n        self.scaler_yield = StandardScaler()\n        if self.yield_ is not None:\n            self.scaler_yield.fit(self.yield_.values)\n\n    def _check_data(self, climate: list, yield_: pd.DataFrame) -> bool:\n        # Check for matching year, lon, lat columns\n        for i in range(1, len(climate)):\n            assert np.all(climate[0][['year', 'lon', 'lat']] == climate[i][['year', 'lon', 'lat']])\n        # Check label for matching length\n        assert yield_ is None or climate[0].shape[0] == yield_.shape[0]","metadata":{"execution":{"iopub.status.busy":"2024-08-07T07:01:37.928967Z","iopub.execute_input":"2024-08-07T07:01:37.929364Z","iopub.status.idle":"2024-08-07T07:01:37.954924Z","shell.execute_reply.started":"2024-08-07T07:01:37.929329Z","shell.execute_reply":"2024-08-07T07:01:37.954017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define model\nThe LSTM is used as backbone\n\nThe input features consist of time series data (tasmax, tasmin, pr, rsds, cumulative rsds) and static attributes (lon, lat, co2, nitrogen and texture_class)\n\nMature date is calculated using temperature sum according to [WOFOST](https://wofost.readthedocs.io/en/7.2/)\n\nThe average value of last ten days is used as prediction to increase robustness","metadata":{}},{"cell_type":"code","source":"class MODEL(nn.Module):\n    def __init__(self, crop,input_fea, hidden_fea, layer_num, output_fea):\n        super().__init__()\n        self.hidden_fea = hidden_fea\n        self.layer_num = layer_num\n        self.lstm = nn.LSTM(input_fea, hidden_fea, layer_num, batch_first=True, dropout = 0.0)\n        self.fc1 = nn.Linear(hidden_fea, 256)\n        self.fc2 = nn.Linear(256, output_fea)\n        state_size = (1, 1, self.hidden_fea)\n        self.h0 = Variable(torch.randn(state_size)).to(device)\n        self.c0 = Variable(torch.randn(state_size)).to(device)\n        \n        if crop == \"wheat\": # https://github.com/ajwdewit/WOFOST_crop_parameters/blob/master/wheat.yaml\n            Tsum, Tbase, TEFFMX = 2000, 0, 30\n        elif crop == \"maize\": # https://github.com/ajwdewit/WOFOST_crop_parameters/blob/master/maize.yaml\n            Tsum, Tbase, TEFFMX = 1600, 8, 22\n        self.Tsum = Tsum\n        self.Tbase = Tbase\n        self.TEFFMX = TEFFMX\n    \n    def forward(self, climate_inverse,climate,soil):\n        batch_size = climate.shape[0]\n        time_steps = climate.shape[1]\n        soil = soil[:,None,:].repeat([1,210,1])\n        X = torch.concat([climate,soil],-1)\n                        \n        h0 = self.h0.repeat(self.layer_num,batch_size,1)\n        c0 = self.c0.repeat(self.layer_num,batch_size,1)\n        hn_all, (hn, cn) = self.lstm(X, (h0,c0))\n        Y = F.relu(self.fc1(hn_all))\n        Y = self.fc2(Y)\n        \n        # Adjusting the mature_indices to gather a sequence of last 10 days\n        mature_indices = self.get_growth_mask(climate_inverse[:,:,0],climate_inverse[:,:,1])\n        mature_indices = mature_indices.unsqueeze(-1) -10 + torch.arange(10).to(mature_indices.device)\n        mature_indices = mature_indices.clamp(0, time_steps - 1)  # Clamping to valid indices\n        mature_indices = mature_indices.unsqueeze(-1)\n        gathered_Y = torch.gather(Y, 1, mature_indices)\n        average_Y = torch.mean(torch.abs(gathered_Y), dim=1)\n        # return torch.abs(gathered_Y[:,-1,:])\n        return average_Y\n    \n    def get_growth_mask(self,tmin,tmax):\n        T = (tmin+tmax)/2\n        tt_rate = self.TEFFMX - F.relu(self.TEFFMX -  F.relu(T-self.Tbase))\n        tt = torch.cumsum(tt_rate,-1)\n        dvs = tt/self.Tsum\n        diff = torch.abs(dvs - 1)\n        _, indices = torch.min(diff, dim=1)\n        return indices","metadata":{"execution":{"iopub.status.busy":"2024-08-07T07:01:37.956976Z","iopub.execute_input":"2024-08-07T07:01:37.957277Z","iopub.status.idle":"2024-08-07T07:01:37.973585Z","shell.execute_reply.started":"2024-08-07T07:01:37.957252Z","shell.execute_reply":"2024-08-07T07:01:37.972647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load data, train and predict","metadata":{}},{"cell_type":"code","source":"    ids = []\n    yields_pred = []\n    for crop in [\"maize\",\"wheat\"]:\n        # Intialize model, data loader, loss and optimizer\n        ds_train = ClimateDataset(crop, 'train', data_dir=DATA_DIR)\n        ds_test = ClimateDataset(crop, 'test', data_dir=DATA_DIR, scalers=(ds_train.scaler_climate, ds_train.scaler_soil, ds_train.scaler_yield))\n        train_loader = DataLoader(ds_train, batch_size=100, shuffle=True)\n        test_loader = DataLoader(ds_test, batch_size=1024, shuffle=False)\n\n        model = MODEL(crop,5+4+13, 32, 1, 1)\n        model.to(device)\n        cost_mse = torch.nn.MSELoss()\n        optimizer = torch.optim.Adam(model.parameters(), lr=0.01)\n        scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda epoch: 0.9**(int(epoch)))\n        # Main training loop   \n        # %%train\n        num_epochs = 5\n        loss_min = 100\n        for epoch in range(num_epochs):\n            loss_epoch = 0\n            for i, (climate_inverse, climate, soil, yield_true, _), in enumerate(train_loader):\n                optimizer.zero_grad()\n                yield_pred = model(climate_inverse.to(device),climate.to(device), soil.to(device))\n                data_loss = cost_mse(yield_pred, yield_true.to(device))\n                data_loss.backward()\n                optimizer.step()\n                loss_epoch += data_loss.item()\n                print(\"\\r\",f\"e-i: {epoch:02d}-{i:05d}/{len(train_loader)} Loss: {data_loss.item():5.10f}\", end='')\n            scheduler.step()\n            loss_epoch = loss_epoch/len(train_loader)\n            if loss_epoch<=loss_min:\n                model_best = copy.deepcopy(model.state_dict())\n            print(f\"epoch: {epoch:02d} Loss: {loss_epoch:5.10f}\")\n        # predict\n\n        model.load_state_dict(model_best,strict=True)  \n        for i,(climate_inverse, climate, soil, _, id) in enumerate(test_loader):\n            ids.append(id.detach().numpy())\n            yields_pred.append(model(climate_inverse.to(device),climate.to(device), soil.to(device)).detach().cpu().numpy())\n            if i%100==0: print(f'{i}/{len(test_loader)}')\n    predictions = pd.Series(np.concatenate(yields_pred).reshape(-1), index=np.concatenate(ids))\n    predictions.index.name = 'ID'\n    predictions.name = 'yield'\n    predictions.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-08-07T07:01:38.003934Z","iopub.execute_input":"2024-08-07T07:01:38.004209Z","iopub.status.idle":"2024-08-07T12:50:25.430755Z","shell.execute_reply.started":"2024-08-07T07:01:38.004184Z","shell.execute_reply":"2024-08-07T12:50:25.429859Z"},"trusted":true},"execution_count":null,"outputs":[]}]}