|
72 | 72 | "source": [ |
73 | 73 | "import logging\n", |
74 | 74 | "from typing import List, Optional, Union\n", |
| 75 | + "\n", |
75 | 76 | "import numpy as np\n", |
76 | 77 | "import pandas as pd\n", |
77 | 78 | "import tensorflow as tf\n", |
78 | | - "from tfts import AutoModel, AutoConfig, KerasTrainer" |
| 79 | + "\n", |
| 80 | + "from tfts import AutoConfig, AutoModel, KerasTrainer" |
79 | 81 | ] |
80 | 82 | }, |
81 | 83 | { |
|
103 | 105 | "class CFG:\n", |
104 | 106 | " input_dir = \"/kaggle/input/china-vehicle-sales-data/china_vehicle_sales_data.csv\"\n", |
105 | 107 | " train_sequence_length = 12\n", |
106 | | - " predict_sequence_length = 3\n" |
| 108 | + " predict_sequence_length = 3" |
107 | 109 | ] |
108 | 110 | }, |
109 | 111 | { |
|
316 | 318 | }, |
317 | 319 | { |
318 | 320 | "cell_type": "code", |
319 | | - "execution_count": 5, |
| 321 | + "execution_count": null, |
320 | 322 | "id": "ad6327cc", |
321 | 323 | "metadata": { |
322 | 324 | "execution": { |
|
340 | 342 | "\n", |
341 | 343 | "logger = logging.getLogger(__name__)\n", |
342 | 344 | "\n", |
| 345 | + "\n", |
343 | 346 | "def add_lagging_feature(\n", |
344 | 347 | " data: pd.DataFrame,\n", |
345 | 348 | " groupby_column: Union[str, List[str]],\n", |
|
364 | 367 | " for lag in lags:\n", |
365 | 368 | " feature_col_name = f\"{column}_lag{lag}\"\n", |
366 | 369 | " feature_columns.append(feature_col_name)\n", |
367 | | - " logger.debug(\n", |
368 | | - " f\"Creating lagging feature: {feature_col_name} for column '{column}' with lag {lag} and groupby '{groupby_column}'.\"\n", |
369 | | - " )\n", |
370 | 370 | " data[feature_col_name] = data.groupby(groupby_column)[column].shift(lag)\n", |
371 | 371 | " return data" |
372 | 372 | ] |
|
759 | 759 | "source": [ |
760 | 760 | "feature_columns = []\n", |
761 | 761 | "\n", |
762 | | - "data = add_lagging_feature(data, groupby_column=[\"provinceId\", \"model\"], value_columns=[\"salesVolume\"], lags=list(range(1, 12)), feature_columns=feature_columns)\n", |
| 762 | + "data = add_lagging_feature(\n", |
| 763 | + " data,\n", |
| 764 | + " groupby_column=[\"provinceId\", \"model\"],\n", |
| 765 | + " value_columns=[\"salesVolume\"],\n", |
| 766 | + " lags=list(range(1, 12)),\n", |
| 767 | + " feature_columns=feature_columns,\n", |
| 768 | + ")\n", |
763 | 769 | "\n", |
764 | 770 | "data" |
765 | 771 | ] |
|
854 | 860 | ], |
855 | 861 | "source": [ |
856 | 862 | "grouped_sequence = data.groupby([\"provinceId\", \"model\"]).apply(\n", |
857 | | - " lambda x: x.sort_values('Date')[[\"salesVolume\", \"salesVolume_lag1\", \"salesVolume_lag2\", \"salesVolume_lag3\"]].to_numpy()\n", |
| 863 | + " lambda x: x.sort_values(\"Date\")[\n", |
| 864 | + " [\"salesVolume\", \"salesVolume_lag1\", \"salesVolume_lag2\", \"salesVolume_lag3\"]\n", |
| 865 | + " ].to_numpy()\n", |
858 | 866 | ")\n", |
859 | 867 | "\n", |
860 | 868 | "data_3d = np.stack(grouped_sequence.values)\n", |
|
902 | 910 | " self.total_samples = self.num_ids * self.samples_per_id\n", |
903 | 911 | "\n", |
904 | 912 | " # Precompute all valid (id, start_idx) pairs\n", |
905 | | - " self.indices = [\n", |
906 | | - " (i, j)\n", |
907 | | - " for i in range(self.num_ids)\n", |
908 | | - " for j in range(self.samples_per_id)\n", |
909 | | - " ]\n", |
910 | | - " \n", |
| 913 | + " self.indices = [(i, j) for i in range(self.num_ids) for j in range(self.samples_per_id)]\n", |
| 914 | + "\n", |
911 | 915 | " def __getitem__(self, index):\n", |
912 | | - " # batch-wise item \n", |
913 | | - " batch_indices = self.indices[index * self.batch_size:(index + 1) * self.batch_size]\n", |
914 | | - " \n", |
| 916 | + " # batch-wise item\n", |
| 917 | + " batch_indices = self.indices[index * self.batch_size : (index + 1) * self.batch_size]\n", |
| 918 | + "\n", |
915 | 919 | " x_batch = []\n", |
916 | 920 | " y_batch = []\n", |
917 | 921 | "\n", |
918 | 922 | " for id_idx, start_idx in batch_indices:\n", |
919 | | - " x = self.data[id_idx, start_idx:start_idx + self.train_seq_len, 1:]\n", |
920 | | - " y = self.data[id_idx, start_idx + self.train_seq_len:start_idx + self.train_seq_len + self.pred_seq_len, 0]\n", |
| 923 | + " x = self.data[id_idx, start_idx : start_idx + self.train_seq_len, 1:]\n", |
| 924 | + " y = self.data[\n", |
| 925 | + " id_idx, start_idx + self.train_seq_len : start_idx + self.train_seq_len + self.pred_seq_len, 0\n", |
| 926 | + " ]\n", |
921 | 927 | " x_batch.append(x)\n", |
922 | 928 | " y_batch.append(y)\n", |
923 | 929 | "\n", |
924 | 930 | " return np.nan_to_num(np.array(x_batch)), np.nan_to_num(np.array(y_batch))\n", |
925 | | - " \n", |
| 931 | + "\n", |
926 | 932 | " def __len__(self):\n", |
927 | 933 | " # depends on how many samples you want to extract from 1 ID\n", |
928 | 934 | " return int(np.ceil(len(self.indices) / self.batch_size))" |
|
1086 | 1092 | "source": [ |
1087 | 1093 | "def build_model():\n", |
1088 | 1094 | " inputs = tf.keras.Input(shape=(CFG.train_sequence_length, 3))\n", |
1089 | | - " \n", |
| 1095 | + "\n", |
1090 | 1096 | " config = AutoConfig()(\"rnn\")\n", |
1091 | 1097 | " config.rnn_type = \"lstm\"\n", |
1092 | 1098 | " backbone = AutoModel.from_config(config=config)\n", |
1093 | | - " \n", |
| 1099 | + "\n", |
1094 | 1100 | " outputs = backbone(inputs)\n", |
1095 | 1101 | " model = tf.keras.Model(inputs=inputs, outputs=outputs)\n", |
1096 | | - " model.compile(loss=tf.keras.losses.MeanAbsoluteError(), optimizer=tf.keras.optimizers.Adam(), metrics = ['mae'])\n", |
| 1102 | + " model.compile(loss=tf.keras.losses.MeanAbsoluteError(), optimizer=tf.keras.optimizers.Adam(), metrics=[\"mae\"])\n", |
1097 | 1103 | " return model\n", |
1098 | 1104 | "\n", |
1099 | 1105 | "\n", |
|
1165 | 1171 | } |
1166 | 1172 | ], |
1167 | 1173 | "source": [ |
1168 | | - "history = model.fit(train_dataset, validation_data=valid_dataset, epochs=10) \n", |
1169 | | - "model.save_weights('./sales_model.weights.h5')" |
| 1174 | + "history = model.fit(train_dataset, validation_data=valid_dataset, epochs=10)\n", |
| 1175 | + "model.save_weights(\"./sales_model.weights.h5\")" |
1170 | 1176 | ] |
1171 | 1177 | } |
1172 | 1178 | ], |
|
0 commit comments