From 4e0f5b8b91259cc078c4de42e10d7b1ed73c25a7 Mon Sep 17 00:00:00 2001 From: chan Date: Tue, 11 Jul 2023 11:57:54 +0900 Subject: [PATCH] chan --- dacon/lstm_base.ipynb | 1158 ++++++++++++++++++++++++++++++++++++ robo/finance_predict.ipynb | 22 - 2 files changed, 1158 insertions(+), 22 deletions(-) create mode 100644 dacon/lstm_base.ipynb 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", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
일자종목코드종목명거래량시가고가저가종가
020210601A0603103S1666902890297028852920
120210601A095570AJ네트웍스638365860594057505780
220210601A006840AK홀딩스10369135500356003415034400
320210601A054620APS46254414600149501380014950
420210601A265520AP시스템13198729150291502880029050
\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.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", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
종목코드final_return순위
5A211270[-0.06308102]1
1A095570[-0.04776682]2
8A126600[-0.038834322]3
7A282330[-0.01267713]4
6A027410[-0.009586]5
2A006840[0.016609492]6
9A138930[0.01706481]7
3A054620[0.01961614]8
4A265520[0.05338428]9
0A060310[0.073663026]10
\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", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
종목코드final_return순위
0A060310[-8.4089585e-08]1
1A095570[0.0]2
2A006840[8.503587e-08]9
3A054620[0.0]3
4A265520[0.0]4
5A211270[0.0]5
6A027410[0.0]6
7A282330[0.0]7
8A126600[0.0]8
9A138930[8.691532e-08]10
\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,