{ "cells": [ { "cell_type": "markdown", "id": "a31d3c40-3593-4ca3-980e-578ee51e171a", "metadata": {}, "source": [ "# 用神经网络进行回归预测" ] }, { "cell_type": "code", "execution_count": 1, "id": "738c03b4-3ca0-4a87-9143-53ddea5179be", "metadata": {}, "outputs": [], "source": [ "import numpy as np\n", "import pandas as pd\n", "import torch\n", "import torch.nn as nn\n", "import torch.optim as optim\n", "from torch.utils.data import Dataset, DataLoader\n", "from sklearn.preprocessing import StandardScaler\n", "from sklearn.model_selection import train_test_split" ] }, { "cell_type": "markdown", "id": "637063c9-5924-416c-8201-adaae8514ffd", "metadata": {}, "source": [ "设置神经网络超参数,批大小为500,学习率为0.01,一共分别进行100次正向传播和反向传播" ] }, { "cell_type": "code", "execution_count": 2, "id": "b909615c-783d-4bfb-acf6-1bc3bacfbb18", "metadata": {}, "outputs": [], "source": [ "# torch参数\n", "batch_size = 500\n", "lr = 0.01\n", "max_epochs = 100\n", "num_workers = 0\n", "device = torch.device(\"cuda:1\" if torch.cuda.is_available() else \"cpu\")" ] }, { "cell_type": "markdown", "id": "904cfade-5667-440a-a13c-34fdf57f7b0f", "metadata": {}, "source": [ "定义所需数据集类" ] }, { "cell_type": "code", "execution_count": 3, "id": "f6e65156-8535-49c2-8d9d-ecb8639855c9", "metadata": {}, "outputs": [], "source": [ "variables = ['number_of_reviews', 'price', 'accommodates',\n", " 'host_response_rate', 'host_acceptance_rate', 'review_scores_rating']\n", "# 数据集类\n", "\n", "\n", "class USDataset(Dataset):\n", " def __init__(self, df):\n", " '''\n", " 初始化\n", " df: 处理后的数据集\n", " '''\n", " self.df = df\n", " self.info = df[['number_of_reviews', 'price', 'accommodates',\n", " 'host_response_rate', 'host_acceptance_rate']].values\n", " self.target = df['review_scores_rating'].values\n", "\n", " def __getitem__(self, index):\n", " '''\n", " 根据编号返回信息\n", " index: 样本编号\n", " '''\n", " info = self.info[index]\n", " target = self.target[index]\n", " return info, target\n", "\n", " def __len__(self):\n", " '''\n", " 返回数据集样本个数\n", " '''\n", " return len(self.df)" ] }, { "cell_type": "markdown", "id": "fc574f2e-6443-48af-9fee-dcf5fb29af58", "metadata": {}, "source": [ "读入数据,由于样本量很大,直接删除有缺失值的样本。" ] }, { "cell_type": "code", "execution_count": 4, "id": "17f8aedb-abf1-4622-98ec-d525a779274b", "metadata": {}, "outputs": [ { "data": { "text/html": [ "
| \n", " | host_response_rate | \n", "host_acceptance_rate | \n", "accommodates | \n", "price | \n", "number_of_reviews | \n", "review_scores_rating | \n", "
|---|---|---|---|---|---|---|
| 0 | \n", "1.00 | \n", "0.33 | \n", "2.0 | \n", "120.0 | \n", "90.0 | \n", "4.50 | \n", "
| 1 | \n", "1.00 | \n", "0.98 | \n", "2.0 | \n", "90.0 | \n", "351.0 | \n", "4.58 | \n", "
| 2 | \n", "1.00 | \n", "0.98 | \n", "2.0 | \n", "66.0 | \n", "67.0 | \n", "4.52 | \n", "
| 3 | \n", "1.00 | \n", "0.98 | \n", "1.0 | \n", "33.0 | \n", "297.0 | \n", "4.70 | \n", "
| 5 | \n", "1.00 | \n", "1.00 | \n", "2.0 | \n", "45.0 | \n", "42.0 | \n", "4.98 | \n", "
| ... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "
| 203252 | \n", "1.00 | \n", "0.93 | \n", "4.0 | \n", "152.0 | \n", "1.0 | \n", "4.00 | \n", "
| 203253 | \n", "1.00 | \n", "0.97 | \n", "2.0 | \n", "45.0 | \n", "1.0 | \n", "3.00 | \n", "
| 203254 | \n", "1.00 | \n", "0.97 | \n", "2.0 | \n", "40.0 | \n", "1.0 | \n", "1.00 | \n", "
| 203276 | \n", "0.99 | \n", "0.99 | \n", "2.0 | \n", "43.0 | \n", "1.0 | \n", "5.00 | \n", "
| 203308 | \n", "1.00 | \n", "1.00 | \n", "3.0 | \n", "110.0 | \n", "1.0 | \n", "5.00 | \n", "
134835 rows × 6 columns
\n", "