{"metadata":{"kernelspec":{"display_name":"dl_env","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.9.20"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84896,"databundleVersionId":10305135,"sourceType":"competition"}],"dockerImageVersionId":30822,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<a id='c1'></a>\n# <div style=\"text-align:center; border-radius:15px 15px; padding:15px; color:#333333; margin:0; ; padding:15px; font-size:100%; font:'Verdana'; background-color:#F5F5F5;border: 1px; overflow:hidden\"><b>1. 导入模块，数据以及函数</b></div>","metadata":{}},{"cell_type":"code","source":"13139+7793","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n#显示所有的行\npd.set_option('display.max_rows', None)\npd.set_option('display.max_columns', None)\nimport plotly.express as px\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.preprocessing import StandardScaler, OneHotEncoder, OrdinalEncoder\nimport torch\nfrom torch import nn, optim\nfrom torch.utils.data import DataLoader, TensorDataset\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import mean_squared_error\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nfrom tqdm import tqdm\nimport imageio\nimport os\n","metadata":{"ExecuteTime":{"end_time":"2024-12-20T11:49:47.811846Z","start_time":"2024-12-20T11:49:47.743045Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/playground-series-s4e12/train.csv')\ntest_df = pd.read_csv('/kaggle/input/playground-series-s4e12/test.csv')\ninsurance_df = pd.concat([train_df, test_df], ignore_index=True)\nidx = insurance_df['Premium Amount'].isna()","metadata":{"ExecuteTime":{"end_time":"2024-12-20T11:49:50.412509Z","start_time":"2024-12-20T11:49:47.812833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PlotlyVisvalizer:\n    def __init__(self, dataframe):\n        self.df = dataframe\n        \n    def histogram(self, column, title = 'Histogram', color = 'blue', hue = None):\n        fig = px.histogram(self.df, x = column, title = title, color = hue)\n        fig.update_traces(marker_line_color = 'black', marker_line_width = 1.5, opacity = 0.8)\n        fig.update_xaxes(title_text = column, title_font = dict(size = 15, color = 'black'), tickangle = 45)\n        fig.update_yaxes(title_text = 'Count', title_font = dict(size = 15, color = 'black'))\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n    \n    def scatter(self, x_column, y_column, title = 'Scatter Plot', color = 'blue', hue = None):\n        fig = px.scatter(self.df, x = x_column, y = y_column, title = title, color = hue)\n        fig.update_traces(marker = dict(size = 10, opacity = 0.8, line = dict(width = 1.5, color = 'black')))\n        fig.update_xaxes(title_text = x_column, title_font = dict(size = 15, color = 'black'))\n        fig.update_yaxes(title_text = y_column, title_font = dict(size = 15, color = 'black'))\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n    \n    def boxplot(self, x_column, y_column, title = 'Box Plot', color = 'blue', hue = None):\n        fig = px.box(self.df, x = x_column, y = y_column, title = title, color = hue)\n        fig.update_traces(marker_line_color = 'black', marker_line_width = 1.5, opacity = 0.8)\n        fig.update_xaxes(title_text = x_column, title_font = dict(size = 15, color = 'black'))\n        fig.update_yaxes(title_text = y_column, title_font = dict(size = 15, color = 'black'))\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n        \n    def line_plot(self, x_column, y_column, title = 'Line Plot', color = 'blue', hue = None):\n        fig = px.line(self.df, x = x_column, y = y_column, title = title, color = hue)\n        fig.update_traces(line = dict(width = 2), marker = dict(size = 5))\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5,xaxis_title = x_column, yaxis_title = y_column, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n        \n    def bar_chart(self, x_column, y_column, title = 'Bar Chart', color = 'blue', hue = None):\n        fig = px.bar(self.df, x = x_column, y = y_column, title = title, color = hue)\n        fig.update_traces(marker_color = color)\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, xaxis_title = x_column, yaxis_title = y_column, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n        \n    def pie_chart(self, column, title = 'Pie Chart', color = 'blue', hue = None):\n        value_counts = self.df[column].value_counts()\n        total = value_counts.sum()\n        percentage = value_counts / total * 100\n        \n        filtered_percentage = percentage[percentage > 1]\n        \n        other_sum = value_counts[percentage < 1].sum()\n        if other_sum > 0:\n            filtered_percentage['Other'] = other_sum\n        \n        \n        fig = px.pie(filtered_percentage, values = percentage, title = title, color = hue)\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n        \n    def heatmap(self, x_column, y_column, z_column, title = 'Heatmap'):\n        fig = px.density_heatmap(self.df, x = x_column, y = y_column, z = z_column, title = title)\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n    \n    def area_chart(self, x_column, y_column, title = 'Area Chart', color = 'blue', hue = None):\n        fig = px.area(self.df, x = x_column, y = y_column, title = title, color = hue)\n        fig.update_traces(marker = dict(size = 10, opacity = 0.8, line = dict(width = 1.5, color = 'black')))\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, xaxis_title = x_column, yaxis_title = y_column, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n        \n    def violin_plot(self, x_column, y_column, title = 'Violin Plot', color = 'blue', hue = None):\n        fig = px.violin(self.df, x = x_column, y = y_column, title = title, color = hue)\n        fig.update_traces(marker = dict(size = 10, opacity = 0.8, line = dict(width = 1.5, color = 'black')))\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, xaxis_title = x_column, yaxis_title = y_column, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n    \n    def sunburst_chart(self, x_column, y_column, title = 'Sunburst Chart', color = 'blue', hue = None):\n        fig = px.sunburst(self.df, x = x_column, y = y_column, title = title, color = hue)\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n    \n    def parallel_coordinates(self, x_column, y_column, title = 'Parallel Coordinates', color = 'blue', hue = None):\n        fig = px.parallel_coordinates(self.df, x = x_column, y = y_column, title = title, color = hue)\n        fig.update_layout(title_font = dict(size = 20, color = 'darkblue'), title_x = 0.5, plot_bgcolor = 'white', paper_bgcolor = 'white')\n        fig.show()\n    \n    def missmap(self, title_text='Missing Map'):  # 改变参数名以避免混淆\n        # 计算每个列的缺失值数量\n        missing_matrix = self.df.isna().sum().astype(int)\n\n        # 将一维数据转换为 DataFrame，并进行转置\n        missing_matrix = missing_matrix.to_frame(name='Missing').T\n\n        # 使用 px.imshow 创建热力图\n        fig = px.imshow(\n            missing_matrix,  # 此时是一个二维 DataFrame\n            color_continuous_scale='Viridis',\n            labels={\"color\": \"Missing\"},\n            aspect='auto'  # 自动调整宽高比\n        )\n\n        # 更新布局\n        fig.update_layout(\n            title=dict(\n                text=title_text,  # 确保 title 是字符串\n                font=dict(size=20, color='darkblue'),\n                x=0.5\n            ),\n            plot_bgcolor='white',\n            paper_bgcolor='white',\n            yaxis_title=\"Features\",  # 特征名称\n            xaxis_title=\"Samples\"    # 样本索引\n        )\n        fig.show()\n\npv = PlotlyVisvalizer(insurance_df)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id='c2'></a>\n# <div style=\"text-align:center; border-radius:15px 15px; padding:15px; color:#333333; margin:0; ; padding:15px; font-size:100%; font:'Verdana'; background-color:#F5F5F5;border: 1px; overflow:hidden\"><b>2. EDA</b></div>","metadata":{}},{"cell_type":"code","source":"train_df.info()","metadata":{"ExecuteTime":{"end_time":"2024-12-20T11:49:50.436663Z","start_time":"2024-12-20T11:49:50.427205Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 数据集描述\n\n该数据集包含关于个人人口统计、健康状况和财务信息的详细数据，共计1,200,000条记录，21个字段，涵盖收入、教育水平及保险相关属性等内容。\n\n#### 字段说明\n1. **id**：每条记录的唯一标识符（int）。\n2. **Age（年龄）**：个人的年龄（float，98.4% 非空）。\n3. **Gender（性别）**：个人的性别（string，100% 非空）。\n4. **Annual Income（年收入）**：个人的年收入（float，96.2% 非空）。\n5. **Marital Status（婚姻状况）**：个人的婚姻状况（string，98.4% 非空）。\n6. **Number of Dependents（抚养人数）**：家庭中抚养的人数（float，90.9% 非空）。\n7. **Education Level（教育水平）**：个人的教育水平（string，100% 非空）。\n8. **Occupation（职业）**：个人的职业（string，70.2% 非空）。\n9. **Health Score（健康评分）**：个人的健康评分（float，93.8% 非空）。\n10. **Location（地区）**：个人的居住地区（string，100% 非空）。\n11. **Policy Type（保单类型）**：保险保单的类型（string，100% 非空）。\n12. **Previous Claims（历史理赔次数）**：历史理赔记录次数（float，69.7% 非空）。\n13. **Vehicle Age（车辆年龄）**：所拥有车辆的年龄（float，100% 非空）。\n14. **Credit Score（信用评分）**：个人的信用评分（float，88.5% 非空）。\n15. **Insurance Duration（保险期限）**：保险的持续时间（float，100% 非空）。\n16. **Policy Start Date（保单开始日期）**：保险保单的开始日期（string，100% 非空）。\n17. **Customer Feedback（客户反馈）**：客户反馈信息（string，93.5% 非空）。\n18. **Smoking Status（吸烟状态）**：个人是否吸烟（string，100% 非空）。\n19. **Exercise Frequency（锻炼频率）**：个人的锻炼频率（string，100% 非空）。\n20. **Property Type（财产类型）**：个人财产的类型（string，100% 非空）。\n21. **Premium Amount（保费金额）**：个人支付的保费金额（float，100% 非空）。\n\n#### 数据总结\n- **总记录数**：1,200,000  \n- **总字段数**：21  \n- **内存占用**：约192.3 MB  \n- **字段类型**：\n  - 浮点型（float64）：9个  \n  - 整型（int64）：1个  \n  - 字符串型（object）：11个","metadata":{}},{"cell_type":"markdown","source":"### 将Policy Start Date转换为日期类型\n\n","metadata":{}},{"cell_type":"code","source":"insurance_df['Policy Start Date'] = pd.to_datetime(insurance_df['Policy Start Date'])\n# insurance_df['year'] = insurance_df['Policy Start Date'].dt.year\n# insurance_df['month'] = insurance_df['Policy Start Date'].dt.month\n# insurance_df['day'] = insurance_df['Policy Start Date'].dt.day\n# insurance_df['day_of_week'] = insurance_df['Policy Start Date'].dt.dayofweek\n# insurance_df['quarter'] = insurance_df['Policy Start Date'].dt.quarter\n# insurance_df.drop(columns=['Policy Start Date'], inplace=True)\n# 将日期转换为时间戳数值\ninsurance_df['timestamp'] = insurance_df['Policy Start Date'].astype(np.int64) // 10**9\n\n# 删除原始日期列\ninsurance_df.drop(columns=['Policy Start Date'], inplace=True)\n\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"insurance_df.info()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target = 'Premium Amount'\nnumerical_columns = insurance_df.select_dtypes(include=['float64', 'int64','int32']).columns.drop(['id',target]).tolist()\none_hot_columns = ['Gender', 'Marital Status', 'Policy Type', 'Smoking Status',  'Property Type','Occupation']\nlabel_columns = ['Location', 'Customer Feedback','Education Level', 'Exercise Frequency']\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\n# 定义绘制子图的函数\ndef plot_scatterplots_optimized(data, target, hue_column, columns, sample_size=None):\n    n_cols = 3  # 每行放置的子图数量\n    n_rows = (len(columns) + n_cols - 1) // n_cols  # 根据列数计算行数\n\n    # 下采样数据（可选）\n    if sample_size and len(data) > sample_size:\n        data = data.sample(sample_size, random_state=42)\n\n    # 创建多子图\n    fig, axes = plt.subplots(n_rows, n_cols, figsize=(n_cols * 3, n_rows * 3))\n    axes = axes.flatten()  # 将axes展平以便逐一访问\n\n    for i, column in enumerate(columns):\n        sns.scatterplot(ax=axes[i], data=data, x=column, y=target, hue=hue_column, alpha=0.7)\n        axes[i].set_title(f'{column} vs {target}', fontsize=12)\n        axes[i].set_xlabel(column, fontsize=10)\n        axes[i].set_ylabel(target, fontsize=10)\n        axes[i].legend(title=hue_column, loc='best')\n        axes[i].grid(True)\n\n    # 隐藏多余的子图\n    for j in range(len(columns), len(axes)):\n        fig.delaxes(axes[j])\n\n    plt.tight_layout()\n    plt.show()\n\n# 获取列名\ncolumns_to_plot = insurance_df.loc[:, 'Age':].columns\n\n# 调用函数绘制，限制采样到1000条数据\nplot_scatterplots_optimized(data=insurance_df, target='Premium Amount', hue_column='Age', columns=columns_to_plot, sample_size=1000)","metadata":{"ExecuteTime":{"end_time":"2024-12-20T11:49:55.094815Z","start_time":"2024-12-20T11:49:50.437351Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id='c3'></a>\n# <div style=\"text-align:center; border-radius:15px 15px; padding:15px; color:#333333; margin:0; ; padding:15px; font-size:100%; font:'Verdana'; background-color:#F5F5F5;border: 1px; overflow:hidden\"><b>3. 数据预处理</b></div>","metadata":{}},{"cell_type":"code","source":"class DataPreprocessor:\n    \"\"\"\n    A class for preprocessing data with methods to fit and transform datasets.\n    Allows applying the same encoders across multiple datasets.\n\n    Params\n    ----------\n    numerical_columns : list of str\n    one_hot_columns : list of str\n    label_columns : list of str\n\n    Attributes\n    ----------\n    scaler : StandardScaler\n        Scaler object for numerical features.\n    one_hot_encoder : OneHotEncoder\n        Encoder object for one-hot encoding.\n    label_encoders : dict of OrdinalEncoder\n        Dictionary of OrdinalEncoders for label encoding.\n\n    Methods\n    -------\n    fit(df)\n        Fits the encoders on the provided DataFrame.\n    transform(df)\n        Transforms the DataFrame using the fitted encoders.\n    fit_transform(df)\n        Fits the encoders and transforms the DataFrame.\n    \"\"\"\n\n    def __init__(self, numerical_columns, one_hot_columns, label_columns):\n        self.numerical_columns = numerical_columns\n        self.one_hot_columns = one_hot_columns\n        self.label_columns = label_columns\n\n        self.scaler = StandardScaler()\n        self.one_hot_encoder = OneHotEncoder(drop='first', sparse_output=False, handle_unknown='ignore')\n        \n        self.label_encoders = {\n            col: OrdinalEncoder(handle_unknown='use_encoded_value', unknown_value=-1) for col in self.label_columns\n        }\n\n        self.one_hot_feature_names = None\n\n    def fit(self, df):\n        self.scaler.fit(df[self.numerical_columns])\n\n        self.one_hot_encoder.fit(df[self.one_hot_columns])\n        self.one_hot_feature_names = self.one_hot_encoder.get_feature_names_out(self.one_hot_columns)\n\n        for col in self.label_columns:\n            self.label_encoders[col].fit(df[[col]])\n\n    def transform(self, df):\n        df_scaled = df.copy()\n\n        df_scaled[self.numerical_columns] = self.scaler.transform(df[self.numerical_columns])\n\n        encoded_columns = self.one_hot_encoder.transform(df[self.one_hot_columns])\n        encoded_df = pd.DataFrame(encoded_columns, columns=self.one_hot_feature_names, index=df.index)\n\n        for col in self.label_columns:\n            df_scaled[col] = self.label_encoders[col].transform(df[[col]])\n\n        df_final = pd.concat([df_scaled.drop(self.one_hot_columns, axis=1), encoded_df], axis=1)\n\n        return df_final\n\n    def fit_transform(self, df):\n        self.fit(df)\n        return self.transform(df)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(one_hot_columns + label_columns)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(numerical_columns)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"insurance_df = insurance_df.drop(columns=['id'])\ninsurance_df[one_hot_columns + label_columns] = insurance_df[one_hot_columns + label_columns].fillna('Unknown')\n\nfor col in ['Number of Dependents', 'Previous Claims', 'Insurance Duration']:\n    insurance_df[col] = insurance_df[col].fillna(0)\n\nmedian_values = insurance_df[numerical_columns].median()\ninsurance_df[numerical_columns] = insurance_df[numerical_columns].fillna(median_values)\n\npreprocessor = DataPreprocessor(numerical_columns, one_hot_columns, label_columns)\ninsurance_df = preprocessor.fit_transform(insurance_df)\n\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n# 确保只选择数值型列\nnumerical_data = insurance_df[~idx][numerical_columns + label_columns+ [target]].select_dtypes(include=['number'])\n\ncorr = numerical_data.corr()\nmask = np.triu(np.ones_like(corr, dtype=bool))\nf, ax = plt.subplots(figsize=(11, 9))\ncamp = sns.diverging_palette(230, 20, as_cmap=True)\nsns.heatmap(corr, mask=mask, cmap=camp, vmax=.3, center=0, square=True, linewidths=.5, cbar_kws={\"shrink\": .5})\nplt.show()\n\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"numerical_data.head()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 1000\nlearning_rate = 0.001\nsize_hidden1 = 128\nsize_hidden2 = 64\nsize_hidden3 = 32\nsize_hidden4 = 32\nsize_hidden5 = 1\nbatch_size = 1024","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PredictModel(nn.Module):\n    def __init__(self, input_size):\n        super(PredictModel, self).__init__()\n        self.fc1 = nn.Linear(input_size, size_hidden1)\n        self.batch_norm1 = nn.BatchNorm1d(size_hidden1)\n        self.relu1 = nn.ReLU()\n        self.dropout1 = nn.Dropout(0.1)\n        self.fc2 = nn.Linear(size_hidden1, size_hidden2)\n        self.batch_norm2 = nn.BatchNorm1d(size_hidden2)\n        self.silu2 = nn.SiLU()\n        self.dropout2 = nn.Dropout(0.1)\n        self.fc3 = nn.Linear(size_hidden2, size_hidden3)\n        self.batch_norm3 = nn.BatchNorm1d(size_hidden3)\n        self.silu3 = nn.SiLU()\n        self.dropout3 = nn.Dropout(0.1)\n        self.fc4 = nn.Linear(size_hidden3, size_hidden4)\n        self.batch_norm4 = nn.BatchNorm1d(size_hidden4)\n        self.relu4 = nn.ReLU()\n        self.dropout4 = nn.Dropout(0.1)\n        self.fc5 = nn.Linear(size_hidden4, size_hidden5)\n        # self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.batch_norm1(x)\n        x = self.relu1(x)\n        x = self.dropout1(x)\n        \n        x = self.fc2(x)\n        x = self.batch_norm2(x)\n        x = self.silu2(x)\n        x = self.dropout2(x)\n        \n        x = self.fc3(x)\n        x = self.batch_norm3(x)\n        x = self.silu3(x)\n        x = self.dropout3(x)\n        \n        x = self.fc4(x)\n        x = self.batch_norm4(x)\n        x = self.relu4(x)\n        x = self.dropout4(x)\n        \n        x = self.fc5(x)\n        return x","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"X = insurance_df[~idx].drop(target, axis = 1).apply(pd.to_numeric, errors='coerce').values\ny = insurance_df[~idx][target].values\ny = np.log1p(y)\n\nX = torch.tensor(X, dtype=torch.float32).to(device)\ny = torch.tensor(y, dtype=torch.long).to(device)\n\nX_test = torch.tensor(insurance_df[idx].drop(target, axis = 1).apply(pd.to_numeric, errors='coerce').values, dtype=torch.float32).to(device)\n\n\nmodel = PredictModel(input_size=X.shape[1]).to(device)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class KFoldTrainer:\n    def __init__(\n        self,\n        X,\n        y,\n        model_class,\n        num_epochs,\n        learning_rate,\n        k_folds = 5,\n        batch_size = 32,\n        device = device\n    ):\n        self.X = X\n        self.y = y\n        self.model_class = model_class\n        self.num_epochs = num_epochs\n        self.learning_rate = learning_rate\n        self.k_folds = k_folds\n        self.device = device\n        self.batch_size = batch_size\n        self.input_size = X.shape[1]\n        \n        self.fold_rmsle = []\n        \n    def train(self):\n        kfold = KFold(n_splits=self.k_folds, shuffle=True)\n        \n        for fold, (train_idx, val_idx) in enumerate(kfold.split(self.X)):\n            print(f'Fold {fold + 1}')\n            try:\n                rmsle = self._train_fold(fold, train_idx, val_idx)\n                if rmsle is not None:\n                    self.fold_rmsle.append(rmsle)\n                else:\n                    print(f'RMSLE is None for fold {fold + 1}')\n            except Exception as e:\n                print(f'Error in fold {fold + 1}: {e}')\n                self.fold_rmsle.append(float('inf'))\n\n        # 找到最好的模型\n        if self.fold_rmsle:\n            best_fold = np.argmin(self.fold_rmsle)\n            print(f'Best fold {best_fold+1} with RMSLE {self.fold_rmsle[best_fold]}')\n            self.best_fold = best_fold\n        else:\n            print('No valid RMSLE found')\n            self.best_fold = None\n            \n    def _train_fold(self, fold, train_idx, val_idx):\n        image_filenames = []\n\n        # 准备训练和验证数据\n        X_train = self.X[train_idx]\n        y_train = self.y[train_idx]\n\n        X_val = self.X[val_idx]\n        y_val = self.y[val_idx]\n        \n        # 建立tensor dataset和data loader\n        train_dataset = TensorDataset(X_train, y_train)\n        val_dataset = TensorDataset(X_val, y_val)\n\n        train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\n        val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n\n        # 定义模型，损失函数，优化器和调度器\n        model = self.model_class(self.input_size).to(self.device)\n        \n        # 定义RMSLE损失函数\n        def rmsle_loss(pred, target):\n            # 将预测值和目标值从log scale转换回原始值\n            pred = torch.exp(pred) - 1\n            target = torch.exp(target) - 1\n            \n            # 添加一个小的常数避免log(0)\n            epsilon = 1e-6\n            pred = torch.clamp(pred, min=epsilon)\n            target = torch.clamp(target, min=epsilon)\n            \n            log_pred = torch.log(pred + 1)\n            log_target = torch.log(target + 1)\n            squared_error = torch.pow(log_pred - log_target, 2)\n            return torch.sqrt(torch.mean(squared_error))\n            \n        criterion = rmsle_loss\n        optimizer = optim.Adam(model.parameters(), lr=self.learning_rate)\n        scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=1000, gamma=0.1)\n\n        train_losses = []\n        val_losses = []\n        # 训练模型\n        for epoch in tqdm(range(self.num_epochs),desc=f'Training Fold {fold + 1}'):\n            model.train()\n            optimizer.zero_grad()\n            \n            outputs = model(X_train).squeeze()\n            loss = criterion(outputs, y_train.float())\n            loss.backward()\n            optimizer.step()\n        \n            train_losses.append(loss.item())\n            scheduler.step()\n            \n            if (epoch + 1) % 10 == 0:\n                model.eval()\n                with torch.no_grad():\n                    val_outputs = model(X_val).squeeze()\n                    val_loss = criterion(val_outputs, y_val.float())\n                    val_losses.append(val_loss.item())\n                    \n                # 保存图片\n                filename = self._plot_and_save(train_losses, val_losses, fold, epoch)\n                image_filenames.append(filename)\n\n        # 在验证集上评估模型\n        model.eval()\n        with torch.no_grad():\n            val_outputs = model(X_val).squeeze()\n            final_rmsle = criterion(val_outputs, y_val.float()).item()\n            print(f'Final RMSLE on fold {fold + 1}: {final_rmsle:.4f}\\n')\n\n        # 保存该fold的模型参数\n        torch.save(model.state_dict(), f'model_{fold+1}.pth')\n        \n        # 保存最终图片\n        final_plot_filename = self._plot_final(train_losses, val_losses, fold)\n        image_filenames.append(final_plot_filename)\n        \n        # 创建GIF动画\n        self._create_gif(image_filenames, fold)\n        \n        # 删除临时图片文件\n        for filename in set(image_filenames):\n            os.remove(filename)\n        \n        return final_rmsle\n        \n    def _plot_and_save(self, train_losses, val_losses, fold, epoch):\n        plt.figure(figsize=(12, 5))\n\n        # Plot training loss\n        plt.subplot(1, 2, 1)\n        plt.plot(range(1, len(train_losses) + 1), train_losses, label=\"Training Loss\")\n        plt.xlabel(\"Epoch\")\n        plt.ylabel(\"RMSLE Loss\")\n        plt.title(\"Training Loss Over Epochs\")\n        plt.legend()\n        plt.grid(True)\n\n        # Plot validation loss\n        plt.subplot(1, 2, 2)\n        epochs_for_val = [(i+1)*10 for i in range(len(val_losses))]\n        plt.plot(\n            epochs_for_val, \n            val_losses, \n            label=\"Validation Loss (every 10 epochs)\", \n            marker='o'\n        )\n        plt.xlabel(\"Epoch\")\n        plt.ylabel(\"RMSLE Loss\")\n        plt.title(\"Validation Loss Over Epochs\")\n        plt.legend()\n        plt.grid(True)\n\n        plt.tight_layout()\n        filename = f'fold_{fold+1}_epoch_{epoch+1}.png'\n        plt.savefig(filename)\n        plt.close()\n        return filename\n    \n    def _plot_final(self, train_losses, val_losses, fold):\n        plt.figure(figsize=(12, 5))\n        \n        # Final training loss plot\n        plt.subplot(1, 2, 1)\n        plt.plot(range(1, len(train_losses) + 1), train_losses, marker='o')\n        plt.title(f'Training Loss Over Epochs - Fold {fold + 1}')\n        plt.xlabel('Epoch')\n        plt.ylabel('RMSLE Loss')\n        plt.grid(True)\n        \n        # Final validation loss plot\n        plt.subplot(1, 2, 2)\n        epochs_for_val = [(i+1)*10 for i in range(len(val_losses))]\n        plt.plot(epochs_for_val, val_losses, marker='o')\n        plt.title(f'Validation Loss Over Epochs - Fold {fold + 1}')\n        plt.xlabel('Epoch')\n        plt.ylabel('RMSLE Loss')\n        plt.grid(True)\n        \n        plt.tight_layout()\n        final_plot_filename = f'fold_{fold+1}_final.png'\n        plt.savefig(final_plot_filename)\n        plt.close()\n        return final_plot_filename\n    \n    def _create_gif(self, image_filenames, fold):\n        # Create a GIF animation from the saved plots\n        with imageio.get_writer(f'animation_fold_{fold+1}.gif', mode='I', duration=0.5) as writer:\n            for filename in image_filenames:\n                image = imageio.imread(filename)\n                writer.append_data(image)\n                ","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer = KFoldTrainer(X, y, PredictModel, num_epochs, learning_rate, 5,batch_size, device=device)\ntrainer.train()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# 假设trainer是已经训练完成的KFoldTrainer实例\nbest_fold = trainer.best_fold\nbest_model_path = f'model_{best_fold+1}.pth'\n\n# 用最佳fold的参数加载到model\nmodel = PredictModel(input_size=X.shape[1]).to(device)\nmodel.load_state_dict(torch.load(best_model_path))\nmodel.eval()\n\nwith torch.no_grad():\n    y_pred_log = model(X_test).squeeze()  # 使用加载的模型进行预测\n    y_pred = np.expm1(y_pred_log.cpu().numpy())  # 将预测值从log scale转换回来\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.DataFrame({'id': test_df['id'], 'Premium Amount': y_pred})\nsubmission.to_csv('submission.csv', index=False)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.head()","metadata":{},"outputs":[],"execution_count":null}]}