-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy path05_create_labels.py
More file actions
104 lines (81 loc) · 3.68 KB
/
Copy path05_create_labels.py
File metadata and controls
104 lines (81 loc) · 3.68 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
"""
Step 5:建立標籤並合併資料(對應報告第 3.2.3 節末段,及 3.2.4 前置準備)
- 讀取前處理後的新聞 & 股價資料
- 用 shift(-1) 建立「明日漲幅 >= N%」的 label(N = 1~6)
- 合併新聞與標籤(以日期為 key)
- 輸出:data/merged_{company_id}_{N}pct.csv
對象:大盤(TWII)及四支個股(2330、3231、2368、3017)皆執行,
與報告 3.2.3 節以 TWII 為主體說明的分類邏輯一致。
"""
import pandas as pd
import os
from config import COMPANIES, THRESHOLDS, DATA_DIR, TRAIN_END, VAL_END, price_fname
def load_price(company_id: str) -> pd.DataFrame:
path = os.path.join(DATA_DIR, price_fname(company_id))
if not os.path.exists(path):
raise FileNotFoundError(f"找不到 {path},請先執行 01_download_stock_price.py")
df = pd.read_csv(path, parse_dates=["date"])
df["date"] = df["date"].dt.normalize()
return df
def load_news(company_id: str) -> pd.DataFrame:
path = os.path.join(DATA_DIR, f"news_{company_id}_preprocessed.csv")
if not os.path.exists(path):
raise FileNotFoundError(f"找不到 {path},請先執行 04_preprocess_text.py")
df = pd.read_csv(path, parse_dates=["date"])
df["date"] = pd.to_datetime(df["date"], errors="coerce").dt.normalize()
df = df.dropna(subset=["date", "clean_text"])
return df
def aggregate_daily_news(news_df: pd.DataFrame) -> pd.DataFrame:
"""同一天可能有多篇新聞,合併成一篇文字"""
agg = (
news_df.groupby("date")["clean_text"]
.apply(lambda x: " ".join(x))
.reset_index()
)
return agg
def create_merged(company_id: str, name: str):
print(f"\n── {name}({company_id})──")
# 檢查是否全部門檻都已完成
all_done = all(
os.path.exists(os.path.join(DATA_DIR, f"merged_{company_id}_{t}pct.csv"))
for t in THRESHOLDS
)
if all_done:
print(f" 所有門檻已完成,跳過。")
return
price_df = load_price(company_id)
news_df = load_news(company_id)
daily_news = aggregate_daily_news(news_df)
for t in THRESHOLDS:
out = os.path.join(DATA_DIR, f"merged_{company_id}_{t}pct.csv")
if os.path.exists(out):
print(f" {t}% 門檻已存在,跳過。")
continue
# 建立標籤:明天漲幅 >= t%
p = price_df.copy()
p["label"] = (p["漲幅百分比"].shift(-1) >= t).astype(int)
p = p.dropna(subset=["label"]).copy()
p["label"] = p["label"].astype(int)
# 合併(inner join:只保留有新聞也有股價的日期)
merged = pd.merge(daily_news, p[["date", "漲幅百分比", "close", "label"]],
on="date", how="inner")
merged = merged.sort_values("date").reset_index(drop=True)
if merged.empty:
print(f" ⚠️ {t}% 門檻合併後無資料!")
continue
# 加上分割欄位(方便後續模型直接使用)
merged["split"] = "train"
merged.loc[merged["date"] >= pd.Timestamp(TRAIN_END), "split"] = "val"
merged.loc[merged["date"] >= pd.Timestamp(VAL_END), "split"] = "test"
pos = (merged["label"] == 1).sum()
neg = (merged["label"] == 0).sum()
print(f" {t}% 門檻:total={len(merged)},正類={pos},負類={neg}")
merged.to_csv(out, index=False, encoding="utf-8-sig")
print(f" ✅ 儲存 → {out}")
def main():
print("=== Step 5:建立標籤並合併資料(報告 3.2.3 節) ===")
for cid, info in COMPANIES.items():
create_merged(cid, info["name"])
print("\n所有合併資料完成。")
if __name__ == "__main__":
main()