diff --git a/dacon/lstm_base.ipynb b/dacon/lstm_base.ipynb
new file mode 100644
index 0000000..08cd2bc
--- /dev/null
+++ b/dacon/lstm_base.ipynb
@@ -0,0 +1,1158 @@
+{
+ "cells": [
+ {
+ "cell_type": "code",
+ "execution_count": 5,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "import pandas as pd\n",
+ "import numpy as np\n",
+ "import random\n",
+ "import os\n",
+ "\n",
+ "from tqdm import tqdm\n",
+ "from statsmodels.tsa.arima.model import ARIMA\n",
+ "\n",
+ "import warnings\n",
+ "warnings.filterwarnings(\"ignore\")"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 6,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "def seed_everything(seed):\n",
+ " random.seed(seed)\n",
+ " os.environ['PYTHONHASHSEED'] = str(seed)\n",
+ " np.random.seed(seed)\n",
+ "\n",
+ "seed_everything(42) # Seed 고정"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 7,
+ "metadata": {},
+ "outputs": [
+ {
+ "data": {
+ "text/html": [
+ "
\n",
+ "\n",
+ "
\n",
+ " \n",
+ " \n",
+ " | \n",
+ " 일자 | \n",
+ " 종목코드 | \n",
+ " 종목명 | \n",
+ " 거래량 | \n",
+ " 시가 | \n",
+ " 고가 | \n",
+ " 저가 | \n",
+ " 종가 | \n",
+ "
\n",
+ " \n",
+ " \n",
+ " \n",
+ " | 0 | \n",
+ " 20210601 | \n",
+ " A060310 | \n",
+ " 3S | \n",
+ " 166690 | \n",
+ " 2890 | \n",
+ " 2970 | \n",
+ " 2885 | \n",
+ " 2920 | \n",
+ "
\n",
+ " \n",
+ " | 1 | \n",
+ " 20210601 | \n",
+ " A095570 | \n",
+ " AJ네트웍스 | \n",
+ " 63836 | \n",
+ " 5860 | \n",
+ " 5940 | \n",
+ " 5750 | \n",
+ " 5780 | \n",
+ "
\n",
+ " \n",
+ " | 2 | \n",
+ " 20210601 | \n",
+ " A006840 | \n",
+ " AK홀딩스 | \n",
+ " 103691 | \n",
+ " 35500 | \n",
+ " 35600 | \n",
+ " 34150 | \n",
+ " 34400 | \n",
+ "
\n",
+ " \n",
+ " | 3 | \n",
+ " 20210601 | \n",
+ " A054620 | \n",
+ " APS | \n",
+ " 462544 | \n",
+ " 14600 | \n",
+ " 14950 | \n",
+ " 13800 | \n",
+ " 14950 | \n",
+ "
\n",
+ " \n",
+ " | 4 | \n",
+ " 20210601 | \n",
+ " A265520 | \n",
+ " AP시스템 | \n",
+ " 131987 | \n",
+ " 29150 | \n",
+ " 29150 | \n",
+ " 28800 | \n",
+ " 29050 | \n",
+ "
\n",
+ " \n",
+ "
\n",
+ "
"
+ ],
+ "text/plain": [
+ " 일자 종목코드 종목명 거래량 시가 고가 저가 종가\n",
+ "0 20210601 A060310 3S 166690 2890 2970 2885 2920\n",
+ "1 20210601 A095570 AJ네트웍스 63836 5860 5940 5750 5780\n",
+ "2 20210601 A006840 AK홀딩스 103691 35500 35600 34150 34400\n",
+ "3 20210601 A054620 APS 462544 14600 14950 13800 14950\n",
+ "4 20210601 A265520 AP시스템 131987 29150 29150 28800 29050"
+ ]
+ },
+ "execution_count": 7,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "train = pd.read_csv('./train.csv')\n",
+ "train.head()"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": 8,
+ "metadata": {},
+ "outputs": [
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 0/2000 [00:00, ?it/s]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "WARNING:tensorflow:6 out of the last 6 calls to .predict_function at 0x0000023EBD534EE0> triggered tf.function retracing. Tracing is expensive and the excessive number of tracings could be due to (1) creating @tf.function repeatedly in a loop, (2) passing tensors with different shapes, (3) passing Python objects instead of tensors. For (1), please define your @tf.function outside of the loop. For (2), @tf.function has reduce_retracing=True option that can avoid unnecessary retracing. For (3), please refer to https://www.tensorflow.org/guide/function#controlling_retracing and https://www.tensorflow.org/api_docs/python/tf/function for more details.\n",
+ "1/1 [==============================] - 1s 993ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 1/2000 [00:08<4:42:21, 8.48s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 460ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 2/2000 [00:16<4:32:19, 8.18s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 480ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 3/2000 [00:24<4:26:21, 8.00s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 478ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 4/2000 [00:33<4:47:30, 8.64s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 488ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 5/2000 [00:42<4:47:49, 8.66s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 481ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 6/2000 [00:51<4:54:23, 8.86s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 637ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 7/2000 [01:00<4:54:09, 8.86s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 469ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 8/2000 [01:09<4:58:49, 9.00s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 897ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 9/2000 [01:18<4:53:43, 8.85s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 489ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 0%| | 10/2000 [01:27<4:53:19, 8.84s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 1s/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 11/2000 [01:36<4:58:09, 8.99s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 478ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 12/2000 [01:45<4:56:37, 8.95s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 497ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 13/2000 [01:54<4:52:02, 8.82s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 485ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 14/2000 [02:03<4:57:09, 8.98s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 487ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 15/2000 [02:11<4:47:54, 8.70s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 482ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 16/2000 [02:20<4:49:31, 8.76s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 540ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 17/2000 [02:28<4:48:13, 8.72s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 465ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 18/2000 [02:37<4:49:40, 8.77s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 887ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 19/2000 [02:46<4:49:00, 8.75s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 463ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 20/2000 [02:54<4:43:16, 8.58s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 876ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 21/2000 [03:03<4:44:48, 8.64s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 456ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 22/2000 [03:11<4:41:30, 8.54s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 462ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 23/2000 [03:20<4:39:24, 8.48s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 463ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%| | 24/2000 [03:29<4:44:48, 8.65s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 535ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%|▏ | 25/2000 [03:37<4:41:28, 8.55s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 471ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%|▏ | 26/2000 [03:47<4:52:01, 8.88s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 514ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%|▏ | 27/2000 [03:55<4:44:06, 8.64s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 496ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%|▏ | 28/2000 [04:03<4:44:39, 8.66s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 465ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 1%|▏ | 29/2000 [04:11<4:37:32, 8.45s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 508ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 30/2000 [04:20<4:40:27, 8.54s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 454ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 31/2000 [04:28<4:36:24, 8.42s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 506ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 32/2000 [04:37<4:43:05, 8.63s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 516ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 33/2000 [04:46<4:37:41, 8.47s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 532ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 34/2000 [04:54<4:42:28, 8.62s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 472ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 35/2000 [05:04<4:52:05, 8.92s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 504ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 36/2000 [05:13<4:55:41, 9.03s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 463ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 37/2000 [05:22<4:54:40, 9.01s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 477ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 38/2000 [05:32<4:59:13, 9.15s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 929ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 39/2000 [05:42<5:05:31, 9.35s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 506ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 40/2000 [05:50<4:59:41, 9.17s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 445ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 41/2000 [05:59<4:58:03, 9.13s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 501ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 42/2000 [06:09<5:02:33, 9.27s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 0s 463ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 43/2000 [06:18<4:56:43, 9.10s/it]"
+ ]
+ },
+ {
+ "name": "stdout",
+ "output_type": "stream",
+ "text": [
+ "1/1 [==============================] - 1s 588ms/step\n"
+ ]
+ },
+ {
+ "name": "stderr",
+ "output_type": "stream",
+ "text": [
+ " 2%|▏ | 44/2000 [06:28<5:03:23, 9.31s/it]"
+ ]
+ }
+ ],
+ "source": [
+ "import pandas as pd\n",
+ "import numpy as np\n",
+ "from sklearn.preprocessing import MinMaxScaler\n",
+ "from tensorflow.keras.models import Sequential\n",
+ "from tensorflow.keras.layers import LSTM, Dense\n",
+ "\n",
+ "# 추론 결과를 저장하기 위한 dataframe 생성\n",
+ "results_df = pd.DataFrame(columns=['종목코드', 'final_return'])\n",
+ "\n",
+ "# train 데이터에 존재하는 독립적인 종목코드 추출\n",
+ "unique_codes = train['종목코드'].unique()\n",
+ "\n",
+ "features = ['거래량','시가', '고가', '저가', '종가']\n",
+ "# 각 종목코드에 대해서 모델 학습 및 추론 반복\n",
+ "for code in tqdm(unique_codes):\n",
+ " \n",
+ " # 학습 데이터 생성\n",
+ " train_close = train[train['종목코드'] == code][features]\n",
+ " # train_close['일자'] = pd.to_datetime(train_close['일자'], format='%Y%m%d')\n",
+ " # train_close.set_index('일자', inplace=True)\n",
+ " tc = train_close['종가']\n",
+ " \n",
+ " \n",
+ " # 데이터 스케일링\n",
+ " tc_scaled = train_close\n",
+ "\n",
+ " scaler = MinMaxScaler(feature_range=(0, 1))\n",
+ " tc_scaled['거래량'] = scaler.fit_transform(train_close['거래량'].values.reshape(-1, 1))\n",
+ " scaler = MinMaxScaler(feature_range=(0, 1))\n",
+ " tc_scaled['시가'] = scaler.fit_transform(train_close['시가'].values.reshape(-1, 1))\n",
+ " scaler = MinMaxScaler(feature_range=(0, 1))\n",
+ " tc_scaled['고가'] = scaler.fit_transform(train_close['고가'].values.reshape(-1, 1))\n",
+ " scaler = MinMaxScaler(feature_range=(0, 1))\n",
+ " tc_scaled['저가'] = scaler.fit_transform(train_close['저가'].values.reshape(-1, 1))\n",
+ " scaler = MinMaxScaler(feature_range=(0, 1))\n",
+ " tc_scaled['종가'] = scaler.fit_transform(train_close['종가'].values.reshape(-1, 1))\n",
+ "\n",
+ "\n",
+ " # print(tc_scaled)\n",
+ " # tc_scaled = scaler.fit_transform(train_close)\n",
+ "\n",
+ " # 데이터셋 생성\n",
+ " def create_dataset(dataset, time_steps=1):\n",
+ " X, y = [], []\n",
+ " for i in range(len(dataset)-time_steps):\n",
+ " X.append(dataset.iloc[i:(i+time_steps), :].values)\n",
+ " y.append(dataset['종가'].iloc[i+time_steps])\n",
+ " \n",
+ " return np.array(X), np.array(y)\n",
+ "\n",
+ " time_steps = 10 # 시퀀스 길이 설정\n",
+ " X, y = create_dataset(tc_scaled, time_steps)\n",
+ "\n",
+ " # 데이터셋 분할: 학습 데이터와 테스트 데이터\n",
+ " train_size = int(len(X) * 0.8)\n",
+ " X_train, X_test = X[:train_size], X[train_size:]\n",
+ " y_train, y_test = y[:train_size], y[train_size:]\n",
+ "\n",
+ " # LSTM 모델 구축\n",
+ " model = Sequential()\n",
+ " model.add(LSTM(50, return_sequences=True, input_shape=(time_steps, len(features))))\n",
+ " model.add(LSTM(50))\n",
+ " model.add(Dense(1))\n",
+ " model.compile(optimizer='adam', loss='mean_squared_error')\n",
+ "\n",
+ " # 모델 학습\n",
+ " model.fit(X_train, y_train, epochs=50, batch_size=32, verbose=0)\n",
+ " \n",
+ " # 향후 15개의 거래일에 대한 예측\n",
+ " predictions = model.predict(X_test[-15:])\n",
+ " predictions = scaler.inverse_transform(predictions)\n",
+ "\n",
+ " # 최종 수익률 계산\n",
+ " final_return = (predictions[-1] - predictions[0]) / predictions[0]\n",
+ "\n",
+ " # 결과 저장\n",
+ " results_df = results_df.append({'종목코드': code, 'final_return': final_return}, ignore_index=True)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "results_df"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [
+ {
+ "data": {
+ "text/html": [
+ "\n",
+ "\n",
+ "
\n",
+ " \n",
+ " \n",
+ " | \n",
+ " 종목코드 | \n",
+ " final_return | \n",
+ " 순위 | \n",
+ "
\n",
+ " \n",
+ " \n",
+ " \n",
+ " | 5 | \n",
+ " A211270 | \n",
+ " [-0.06308102] | \n",
+ " 1 | \n",
+ "
\n",
+ " \n",
+ " | 1 | \n",
+ " A095570 | \n",
+ " [-0.04776682] | \n",
+ " 2 | \n",
+ "
\n",
+ " \n",
+ " | 8 | \n",
+ " A126600 | \n",
+ " [-0.038834322] | \n",
+ " 3 | \n",
+ "
\n",
+ " \n",
+ " | 7 | \n",
+ " A282330 | \n",
+ " [-0.01267713] | \n",
+ " 4 | \n",
+ "
\n",
+ " \n",
+ " | 6 | \n",
+ " A027410 | \n",
+ " [-0.009586] | \n",
+ " 5 | \n",
+ "
\n",
+ " \n",
+ " | 2 | \n",
+ " A006840 | \n",
+ " [0.016609492] | \n",
+ " 6 | \n",
+ "
\n",
+ " \n",
+ " | 9 | \n",
+ " A138930 | \n",
+ " [0.01706481] | \n",
+ " 7 | \n",
+ "
\n",
+ " \n",
+ " | 3 | \n",
+ " A054620 | \n",
+ " [0.01961614] | \n",
+ " 8 | \n",
+ "
\n",
+ " \n",
+ " | 4 | \n",
+ " A265520 | \n",
+ " [0.05338428] | \n",
+ " 9 | \n",
+ "
\n",
+ " \n",
+ " | 0 | \n",
+ " A060310 | \n",
+ " [0.073663026] | \n",
+ " 10 | \n",
+ "
\n",
+ " \n",
+ "
\n",
+ "
"
+ ],
+ "text/plain": [
+ " 종목코드 final_return 순위\n",
+ "5 A211270 [-0.06308102] 1\n",
+ "1 A095570 [-0.04776682] 2\n",
+ "8 A126600 [-0.038834322] 3\n",
+ "7 A282330 [-0.01267713] 4\n",
+ "6 A027410 [-0.009586] 5\n",
+ "2 A006840 [0.016609492] 6\n",
+ "9 A138930 [0.01706481] 7\n",
+ "3 A054620 [0.01961614] 8\n",
+ "4 A265520 [0.05338428] 9\n",
+ "0 A060310 [0.073663026] 10"
+ ]
+ },
+ "execution_count": 123,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "results_df['순위'] = results_df['final_return'].rank(method='first').astype('int') # 각 순위를 중복없이 생성\n",
+ "results_df.sort_values('순위')"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [
+ {
+ "data": {
+ "text/html": [
+ "\n",
+ "\n",
+ "
\n",
+ " \n",
+ " \n",
+ " | \n",
+ " 종목코드 | \n",
+ " final_return | \n",
+ " 순위 | \n",
+ "
\n",
+ " \n",
+ " \n",
+ " \n",
+ " | 0 | \n",
+ " A060310 | \n",
+ " [-8.4089585e-08] | \n",
+ " 1 | \n",
+ "
\n",
+ " \n",
+ " | 1 | \n",
+ " A095570 | \n",
+ " [0.0] | \n",
+ " 2 | \n",
+ "
\n",
+ " \n",
+ " | 2 | \n",
+ " A006840 | \n",
+ " [8.503587e-08] | \n",
+ " 9 | \n",
+ "
\n",
+ " \n",
+ " | 3 | \n",
+ " A054620 | \n",
+ " [0.0] | \n",
+ " 3 | \n",
+ "
\n",
+ " \n",
+ " | 4 | \n",
+ " A265520 | \n",
+ " [0.0] | \n",
+ " 4 | \n",
+ "
\n",
+ " \n",
+ " | 5 | \n",
+ " A211270 | \n",
+ " [0.0] | \n",
+ " 5 | \n",
+ "
\n",
+ " \n",
+ " | 6 | \n",
+ " A027410 | \n",
+ " [0.0] | \n",
+ " 6 | \n",
+ "
\n",
+ " \n",
+ " | 7 | \n",
+ " A282330 | \n",
+ " [0.0] | \n",
+ " 7 | \n",
+ "
\n",
+ " \n",
+ " | 8 | \n",
+ " A126600 | \n",
+ " [0.0] | \n",
+ " 8 | \n",
+ "
\n",
+ " \n",
+ " | 9 | \n",
+ " A138930 | \n",
+ " [8.691532e-08] | \n",
+ " 10 | \n",
+ "
\n",
+ " \n",
+ "
\n",
+ "
"
+ ],
+ "text/plain": [
+ " 종목코드 final_return 순위\n",
+ "0 A060310 [-8.4089585e-08] 1\n",
+ "1 A095570 [0.0] 2\n",
+ "2 A006840 [8.503587e-08] 9\n",
+ "3 A054620 [0.0] 3\n",
+ "4 A265520 [0.0] 4\n",
+ "5 A211270 [0.0] 5\n",
+ "6 A027410 [0.0] 6\n",
+ "7 A282330 [0.0] 7\n",
+ "8 A126600 [0.0] 8\n",
+ "9 A138930 [8.691532e-08] 10"
+ ]
+ },
+ "execution_count": 118,
+ "metadata": {},
+ "output_type": "execute_result"
+ }
+ ],
+ "source": [
+ "sample_submission = pd.read_csv('./sample_submission.csv')\n",
+ "sample_submission"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "baseline_submission = sample_submission[['종목코드']].merge(results_df[['종목코드', '순위']], on='종목코드', how='left')\n",
+ "baseline_submission"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": [
+ "baseline_submission.to_csv('baseline_submission.csv', index=False)"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "execution_count": null,
+ "metadata": {},
+ "outputs": [],
+ "source": []
+ }
+ ],
+ "metadata": {
+ "kernelspec": {
+ "display_name": "Python 3",
+ "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.12"
+ },
+ "orig_nbformat": 4
+ },
+ "nbformat": 4,
+ "nbformat_minor": 2
+}
diff --git a/robo/finance_predict.ipynb b/robo/finance_predict.ipynb
index cc06f36..7e58067 100644
--- a/robo/finance_predict.ipynb
+++ b/robo/finance_predict.ipynb
@@ -332,28 +332,6 @@
"df_list[0].head()"
]
},
- {
- "cell_type": "code",
- "execution_count": 127,
- "metadata": {},
- "outputs": [
- {
- "ename": "AttributeError",
- "evalue": "'Series' object has no attribute 'to_datetime'",
- "output_type": "error",
- "traceback": [
- "\u001b[1;31m---------------------------------------------------------------------------\u001b[0m",
- "\u001b[1;31mAttributeError\u001b[0m Traceback (most recent call last)",
- "Cell \u001b[1;32mIn[127], line 1\u001b[0m\n\u001b[1;32m----> 1\u001b[0m a \u001b[39m=\u001b[39m df_list[\u001b[39m0\u001b[39;49m][\u001b[39m'\u001b[39;49m\u001b[39mdate\u001b[39;49m\u001b[39m'\u001b[39;49m]\u001b[39m.\u001b[39;49mto_datetime()\n\u001b[0;32m 2\u001b[0m a\n",
- "File \u001b[1;32mc:\\Users\\chocs\\AppData\\Local\\Programs\\Python\\Python39\\lib\\site-packages\\pandas\\core\\generic.py:5902\u001b[0m, in \u001b[0;36mNDFrame.__getattr__\u001b[1;34m(self, name)\u001b[0m\n\u001b[0;32m 5895\u001b[0m \u001b[39mif\u001b[39;00m (\n\u001b[0;32m 5896\u001b[0m name \u001b[39mnot\u001b[39;00m \u001b[39min\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_internal_names_set\n\u001b[0;32m 5897\u001b[0m \u001b[39mand\u001b[39;00m name \u001b[39mnot\u001b[39;00m \u001b[39min\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_metadata\n\u001b[0;32m 5898\u001b[0m \u001b[39mand\u001b[39;00m name \u001b[39mnot\u001b[39;00m \u001b[39min\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_accessors\n\u001b[0;32m 5899\u001b[0m \u001b[39mand\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_info_axis\u001b[39m.\u001b[39m_can_hold_identifiers_and_holds_name(name)\n\u001b[0;32m 5900\u001b[0m ):\n\u001b[0;32m 5901\u001b[0m \u001b[39mreturn\u001b[39;00m \u001b[39mself\u001b[39m[name]\n\u001b[1;32m-> 5902\u001b[0m \u001b[39mreturn\u001b[39;00m \u001b[39mobject\u001b[39;49m\u001b[39m.\u001b[39;49m\u001b[39m__getattribute__\u001b[39;49m(\u001b[39mself\u001b[39;49m, name)\n",
- "\u001b[1;31mAttributeError\u001b[0m: 'Series' object has no attribute 'to_datetime'"
- ]
- }
- ],
- "source": [
- "knn 분류 선이 그어지면 "
- ]
- },
{
"cell_type": "code",
"execution_count": 79,