mirror of
https://github.com/shiyu-coder/Kronos.git
synced 2026-10-06 15:04:09 +08:00
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,629 @@
|
||||
import pandas as pd
|
||||
import requests
|
||||
import json
|
||||
from datetime import datetime, timedelta
|
||||
import os
|
||||
import time
|
||||
import random
|
||||
|
||||
|
||||
def get_stock_market(stock_code):
|
||||
"""
|
||||
根据股票代码判断市场类型
|
||||
返回: 市场前缀 '0'-深交所, '1'-上交所
|
||||
"""
|
||||
if stock_code.startswith(('0', '2', '3')):
|
||||
return '0' # 深交所
|
||||
elif stock_code.startswith(('6', '9')):
|
||||
return '1' # 上交所
|
||||
else:
|
||||
return '1' # 默认上交所
|
||||
|
||||
|
||||
def get_stock_data_eastmoney(stock_code="002354", start_year=2024, end_year=2025):
|
||||
"""
|
||||
使用东方财富网API获取指定年份范围的股票数据 - 修复版
|
||||
"""
|
||||
try:
|
||||
print(f"正在从东方财富网获取股票 {stock_code} 的 {start_year}-{end_year} 年数据...")
|
||||
|
||||
# 计算日期范围
|
||||
start_date = f"{start_year}0101"
|
||||
current_date = datetime.now()
|
||||
|
||||
if current_date.year > end_year:
|
||||
end_date = f"{end_year}1231"
|
||||
else:
|
||||
end_date = current_date.strftime('%Y%m%d')
|
||||
|
||||
print(f"时间范围: {start_date} 到 {end_date}")
|
||||
|
||||
# 获取市场类型
|
||||
market = get_stock_market(stock_code)
|
||||
secid = f"{market}.{stock_code}"
|
||||
|
||||
# 使用更简单的东方财富API
|
||||
url = "http://push2his.eastmoney.com/api/qt/stock/kline/get"
|
||||
|
||||
params = {
|
||||
'secid': secid,
|
||||
'fields1': 'f1,f2,f3,f4,f5,f6',
|
||||
'fields2': 'f51,f52,f53,f54,f55,f56,f57,f58,f59,f60,f61',
|
||||
'klt': '101', # 日线
|
||||
'fqt': '1', # 前复权
|
||||
'beg': start_date,
|
||||
'end': end_date,
|
||||
'lmt': '10000',
|
||||
'ut': 'fa5fd1943c7b386f172d6893dbfba10b',
|
||||
'cb': f'jQuery{random.randint(1000000, 9999999)}_{int(time.time()*1000)}'
|
||||
}
|
||||
|
||||
headers = {
|
||||
'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/117.0.0.0 Safari/537.36',
|
||||
'Referer': 'https://quote.eastmoney.com/',
|
||||
'Accept': '*/*',
|
||||
}
|
||||
|
||||
time.sleep(random.uniform(1, 2))
|
||||
|
||||
response = requests.get(url, params=params, headers=headers, timeout=10)
|
||||
|
||||
print(f"API响应状态码: {response.status_code}")
|
||||
|
||||
if response.status_code == 200:
|
||||
# 处理JSONP响应
|
||||
response_text = response.text
|
||||
|
||||
# 提取JSON数据(处理JSONP格式)
|
||||
if response_text.startswith('/**/'):
|
||||
response_text = response_text[4:]
|
||||
|
||||
# 查找JSON数据的开始和结束位置
|
||||
start_idx = response_text.find('(')
|
||||
end_idx = response_text.rfind(')')
|
||||
|
||||
if start_idx != -1 and end_idx != -1:
|
||||
json_str = response_text[start_idx + 1:end_idx]
|
||||
try:
|
||||
data = json.loads(json_str)
|
||||
except json.JSONDecodeError:
|
||||
print("❌ JSON解析失败,尝试直接解析...")
|
||||
# 如果JSON解析失败,尝试直接提取数据
|
||||
return parse_kline_data_directly(response_text, stock_code, start_year, end_year)
|
||||
else:
|
||||
print("❌ 无法找到JSON数据边界")
|
||||
return None
|
||||
|
||||
print(f"API返回数据状态: {data.get('rc', 'N/A')}")
|
||||
|
||||
if data and data.get('data') is not None:
|
||||
klines = data['data'].get('klines', [])
|
||||
print(f"获取到 {len(klines)} 条K线数据")
|
||||
|
||||
if not klines:
|
||||
print("⚠️ K线数据为空")
|
||||
return None
|
||||
|
||||
# 解析数据
|
||||
stock_data = []
|
||||
for kline in klines:
|
||||
try:
|
||||
items = kline.split(',')
|
||||
if len(items) >= 6:
|
||||
stock_data.append({
|
||||
'日期': items[0],
|
||||
'股票代码': stock_code,
|
||||
'开盘价': float(items[1]),
|
||||
'收盘价': float(items[2]),
|
||||
'最高价': float(items[3]),
|
||||
'最低价': float(items[4]),
|
||||
'成交量': float(items[5]),
|
||||
'成交额': float(items[6]) if len(items) > 6 else 0,
|
||||
'振幅': float(items[7]) if len(items) > 7 else 0,
|
||||
'涨跌幅': float(items[8]) if len(items) > 8 else 0,
|
||||
'涨跌额': float(items[9]) if len(items) > 9 else 0,
|
||||
'换手率': float(items[10]) if len(items) > 10 else 0
|
||||
})
|
||||
except (ValueError, IndexError) as e:
|
||||
continue
|
||||
|
||||
if not stock_data:
|
||||
print("❌ 解析后无有效数据")
|
||||
return None
|
||||
|
||||
df = pd.DataFrame(stock_data)
|
||||
df['日期'] = pd.to_datetime(df['日期'])
|
||||
df.set_index('日期', inplace=True)
|
||||
df = df.sort_index()
|
||||
|
||||
# 筛选指定年份的数据
|
||||
df = df[(df.index.year >= start_year) & (df.index.year <= end_year)]
|
||||
|
||||
print(f"✅ 成功获取 {len(df)} 条有效数据")
|
||||
print(f"实际时间范围: {df.index.min().strftime('%Y-%m-%d')} 到 {df.index.max().strftime('%Y-%m-%d')}")
|
||||
return df
|
||||
else:
|
||||
print("❌ API返回数据为空")
|
||||
return None
|
||||
else:
|
||||
print(f"❌ 请求失败,状态码: {response.status_code}")
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ 获取数据时出错: {str(e)}")
|
||||
return None
|
||||
|
||||
|
||||
def parse_kline_data_directly(response_text, stock_code, start_year, end_year):
|
||||
"""
|
||||
直接解析K线数据(当JSON解析失败时使用)
|
||||
"""
|
||||
try:
|
||||
# 尝试直接从响应文本中提取K线数据
|
||||
if '"klines":[' in response_text:
|
||||
start_idx = response_text.find('"klines":[') + 10
|
||||
end_idx = response_text.find(']', start_idx)
|
||||
klines_str = response_text[start_idx:end_idx]
|
||||
|
||||
# 清理字符串并分割
|
||||
klines = klines_str.replace('"', '').split(',')
|
||||
|
||||
stock_data = []
|
||||
for kline in klines:
|
||||
if kline.strip():
|
||||
items = kline.split(',')
|
||||
if len(items) >= 6:
|
||||
stock_data.append({
|
||||
'日期': items[0],
|
||||
'股票代码': stock_code,
|
||||
'开盘价': float(items[1]),
|
||||
'收盘价': float(items[2]),
|
||||
'最高价': float(items[3]),
|
||||
'最低价': float(items[4]),
|
||||
'成交量': float(items[5]),
|
||||
'成交额': float(items[6]) if len(items) > 6 else 0,
|
||||
})
|
||||
|
||||
if stock_data:
|
||||
df = pd.DataFrame(stock_data)
|
||||
df['日期'] = pd.to_datetime(df['日期'])
|
||||
df.set_index('日期', inplace=True)
|
||||
df = df.sort_index()
|
||||
df = df[(df.index.year >= start_year) & (df.index.year <= end_year)]
|
||||
print(f"✅ 直接解析获取 {len(df)} 条数据")
|
||||
return df
|
||||
except Exception as e:
|
||||
print(f"❌ 直接解析也失败: {e}")
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def get_stock_data_akshare(stock_code="002354", start_year=2024, end_year=2025):
|
||||
"""
|
||||
使用AKShare作为备用数据源 - 修复版
|
||||
"""
|
||||
try:
|
||||
print(f"尝试使用AKShare获取股票 {stock_code} 数据...")
|
||||
import akshare as ak
|
||||
|
||||
# 计算日期范围
|
||||
start_date = f"{start_year}0101"
|
||||
end_date = datetime.now().strftime('%Y%m%d')
|
||||
|
||||
# 获取数据
|
||||
df = ak.stock_zh_a_hist(symbol=stock_code, period="daily",
|
||||
start_date=start_date, end_date=end_date,
|
||||
adjust="qfq")
|
||||
|
||||
if df is not None and not df.empty:
|
||||
# 重命名列以匹配我们的格式
|
||||
column_mapping = {
|
||||
'日期': '日期',
|
||||
'开盘': '开盘价',
|
||||
'收盘': '收盘价',
|
||||
'最高': '最高价',
|
||||
'最低': '最低价',
|
||||
'成交量': '成交量',
|
||||
'成交额': '成交额',
|
||||
'振幅': '振幅',
|
||||
'涨跌幅': '涨跌幅',
|
||||
'涨跌额': '涨跌额',
|
||||
'换手率': '换手率'
|
||||
}
|
||||
|
||||
# 只映射存在的列
|
||||
actual_mapping = {k: v for k, v in column_mapping.items() if k in df.columns}
|
||||
df = df.rename(columns=actual_mapping)
|
||||
|
||||
# 添加股票代码列
|
||||
df['股票代码'] = stock_code
|
||||
df['日期'] = pd.to_datetime(df['日期'])
|
||||
df.set_index('日期', inplace=True)
|
||||
df = df.sort_index()
|
||||
|
||||
# 筛选指定年份
|
||||
df = df[(df.index.year >= start_year) & (df.index.year <= end_year)]
|
||||
|
||||
print(f"✅ AKShare成功获取 {len(df)} 条数据")
|
||||
return df
|
||||
else:
|
||||
print("❌ AKShare未返回数据")
|
||||
return None
|
||||
|
||||
except ImportError:
|
||||
print("⚠️ AKShare未安装,使用 pip install akshare 安装")
|
||||
return None
|
||||
except Exception as e:
|
||||
print(f"❌ AKShare获取数据失败: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def get_stock_data_baostock(stock_code="002354", start_year=2024, end_year=2025):
|
||||
"""
|
||||
使用Baostock作为第三个数据源
|
||||
"""
|
||||
try:
|
||||
print(f"尝试使用Baostock获取股票 {stock_code} 数据...")
|
||||
import baostock as bs
|
||||
import pandas as pd
|
||||
|
||||
# 登录系统
|
||||
lg = bs.login()
|
||||
|
||||
# 计算日期范围
|
||||
start_date = f"{start_year}-01-01"
|
||||
end_date = datetime.now().strftime('%Y-%m-%d')
|
||||
|
||||
# 根据市场添加前缀
|
||||
market = get_stock_market(stock_code)
|
||||
if market == '0':
|
||||
full_code = f"sz.{stock_code}"
|
||||
else:
|
||||
full_code = f"sh.{stock_code}"
|
||||
|
||||
# 获取数据
|
||||
rs = bs.query_history_k_data_plus(
|
||||
full_code,
|
||||
"date,open,high,low,close,volume,amount,turn,pctChg",
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
frequency="d",
|
||||
adjustflag="2" # 前复权
|
||||
)
|
||||
|
||||
data_list = []
|
||||
while (rs.error_code == '0') & rs.next():
|
||||
data_list.append(rs.get_row_data())
|
||||
|
||||
# 退出系统
|
||||
bs.logout()
|
||||
|
||||
if data_list:
|
||||
df = pd.DataFrame(data_list, columns=rs.fields)
|
||||
|
||||
# 数据类型转换
|
||||
df['date'] = pd.to_datetime(df['date'])
|
||||
df['open'] = pd.to_numeric(df['open'])
|
||||
df['high'] = pd.to_numeric(df['high'])
|
||||
df['low'] = pd.to_numeric(df['low'])
|
||||
df['close'] = pd.to_numeric(df['close'])
|
||||
df['volume'] = pd.to_numeric(df['volume'])
|
||||
df['amount'] = pd.to_numeric(df['amount'])
|
||||
df['turn'] = pd.to_numeric(df['turn'])
|
||||
df['pctChg'] = pd.to_numeric(df['pctChg'])
|
||||
|
||||
# 重命名列
|
||||
df = df.rename(columns={
|
||||
'date': '日期',
|
||||
'open': '开盘价',
|
||||
'high': '最高价',
|
||||
'low': '最低价',
|
||||
'close': '收盘价',
|
||||
'volume': '成交量',
|
||||
'amount': '成交额',
|
||||
'turn': '换手率',
|
||||
'pctChg': '涨跌幅'
|
||||
})
|
||||
|
||||
# 添加股票代码列
|
||||
df['股票代码'] = stock_code
|
||||
df.set_index('日期', inplace=True)
|
||||
df = df.sort_index()
|
||||
|
||||
# 筛选指定年份
|
||||
df = df[(df.index.year >= start_year) & (df.index.year <= end_year)]
|
||||
|
||||
# 计算涨跌额
|
||||
df['涨跌额'] = df['收盘价'].diff()
|
||||
|
||||
print(f"✅ Baostock成功获取 {len(df)} 条数据")
|
||||
return df
|
||||
else:
|
||||
print("❌ Baostock未返回数据")
|
||||
return None
|
||||
|
||||
except ImportError:
|
||||
print("⚠️ Baostock未安装,使用 pip install baostock 安装")
|
||||
return None
|
||||
except Exception as e:
|
||||
print(f"❌ Baostock获取数据失败: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def get_stock_data_with_retry(stock_code="002354", start_year=2024, end_year=2025, retry_count=2):
|
||||
"""
|
||||
带重试机制的数据获取 - 多数据源版本
|
||||
"""
|
||||
data_sources = [
|
||||
("AKShare", get_stock_data_akshare),
|
||||
("Baostock", get_stock_data_baostock),
|
||||
("东方财富", get_stock_data_eastmoney)
|
||||
]
|
||||
|
||||
for source_name, data_func in data_sources:
|
||||
print(f"\n🔍 尝试从 {source_name} 获取数据...")
|
||||
data = data_func(stock_code, start_year, end_year)
|
||||
|
||||
if data is not None and not data.empty:
|
||||
# 检查数据是否包含目标年份
|
||||
available_years = data.index.year.unique()
|
||||
print(f"获取到的数据年份: {sorted(available_years)}")
|
||||
|
||||
if any(year in available_years for year in range(start_year, end_year + 1)):
|
||||
print(f"✅ {source_name} 数据获取成功!")
|
||||
# 标记数据来源
|
||||
data.attrs['data_source'] = source_name
|
||||
return data
|
||||
else:
|
||||
print(f"⚠️ 数据未包含目标年份数据")
|
||||
|
||||
print("❌ 所有真实数据源都失败,使用示例数据...")
|
||||
return create_sample_data(stock_code, start_year, end_year)
|
||||
|
||||
|
||||
def create_sample_data(stock_code="002354", start_year=2024, end_year=2025):
|
||||
"""
|
||||
创建更真实的示例数据
|
||||
"""
|
||||
print(f"📊 创建 {start_year}-{end_year} 年的示例数据...")
|
||||
|
||||
# 生成交易日(排除周末)
|
||||
start_date = datetime(start_year, 1, 1)
|
||||
end_date = datetime.now()
|
||||
all_dates = pd.bdate_range(start=start_date, end=end_date, freq='B')
|
||||
|
||||
# 只保留目标年份的数据
|
||||
trading_dates = [date for date in all_dates if start_year <= date.year <= end_year]
|
||||
|
||||
# 生成更真实的股价数据
|
||||
import numpy as np
|
||||
np.random.seed(42)
|
||||
|
||||
# 设置合理的基准价格
|
||||
base_prices = {
|
||||
'600580': 12.0, # 卧龙电驱 - 更合理的价格
|
||||
'002354': 5.0, # 天娱数科
|
||||
'300207': 15.0, # 欣旺达
|
||||
}
|
||||
base_price = base_prices.get(stock_code, 10.0)
|
||||
|
||||
stock_data = []
|
||||
current_price = base_price
|
||||
|
||||
for i, date in enumerate(trading_dates):
|
||||
# 更真实的股价波动
|
||||
volatility = 0.015 # 1.5%的日波动率
|
||||
|
||||
if i > 0:
|
||||
# 使用更真实的随机游走
|
||||
daily_return = np.random.normal(0, volatility)
|
||||
# 添加一些趋势
|
||||
if i < len(trading_dates) * 0.3: # 前30%的时间
|
||||
trend_bias = 0.0005 # 轻微上涨趋势
|
||||
elif i < len(trading_dates) * 0.7: # 中间40%的时间
|
||||
trend_bias = -0.0003 # 轻微下跌趋势
|
||||
else: # 后30%的时间
|
||||
trend_bias = 0.0002 # 轻微上涨趋势
|
||||
|
||||
daily_return += trend_bias
|
||||
current_price = current_price * (1 + daily_return)
|
||||
|
||||
# 价格边界限制 - 更合理
|
||||
current_price = max(base_price * 0.5, min(base_price * 2.0, current_price))
|
||||
else:
|
||||
current_price = base_price
|
||||
|
||||
# 生成OHLC数据
|
||||
open_variation = np.random.normal(0, volatility * 0.2)
|
||||
open_price = current_price * (1 + open_variation)
|
||||
|
||||
daily_range = abs(np.random.normal(volatility * 0.8, volatility * 0.3))
|
||||
high_price = max(open_price, current_price) * (1 + daily_range)
|
||||
low_price = min(open_price, current_price) * (1 - daily_range)
|
||||
close_price = current_price
|
||||
|
||||
# 确保价格合理性
|
||||
high_price = max(open_price, close_price, low_price, high_price)
|
||||
low_price = min(open_price, close_price, high_price, low_price)
|
||||
|
||||
# 生成成交量(更合理)
|
||||
base_volume = 500000 # 基础成交量
|
||||
volume_variation = abs(daily_return) * 3000000 if i > 0 else 0
|
||||
volume = int(base_volume + volume_variation + np.random.randint(-100000, 200000))
|
||||
volume = max(100000, volume)
|
||||
|
||||
# 计算成交额(万元)
|
||||
amount = volume * close_price / 10000
|
||||
|
||||
# 计算涨跌幅和涨跌额
|
||||
if i > 0:
|
||||
prev_close = stock_data[-1]['收盘价']
|
||||
price_change = close_price - prev_close
|
||||
pct_change = (price_change / prev_close) * 100
|
||||
else:
|
||||
price_change = 0
|
||||
pct_change = 0
|
||||
|
||||
# 计算振幅
|
||||
amplitude = ((high_price - low_price) / open_price) * 100
|
||||
|
||||
# 生成换手率(0.5%-8%之间)
|
||||
turnover_rate = np.random.uniform(0.5, 8.0)
|
||||
|
||||
stock_data.append({
|
||||
'日期': date,
|
||||
'股票代码': stock_code,
|
||||
'开盘价': round(open_price, 2),
|
||||
'收盘价': round(close_price, 2),
|
||||
'最高价': round(high_price, 2),
|
||||
'最低价': round(low_price, 2),
|
||||
'成交量': volume,
|
||||
'成交额': round(amount, 2),
|
||||
'振幅': round(amplitude, 2),
|
||||
'涨跌幅': round(pct_change, 2),
|
||||
'涨跌额': round(price_change, 2),
|
||||
'换手率': round(turnover_rate, 2)
|
||||
})
|
||||
|
||||
df = pd.DataFrame(stock_data)
|
||||
df.set_index('日期', inplace=True)
|
||||
|
||||
print(f"✅ 已创建 {len(df)} 条 {start_year}-{end_year} 年的模拟数据")
|
||||
print(f"时间范围: {df.index.min().strftime('%Y-%m-%d')} 到 {df.index.max().strftime('%Y-%m-%d')}")
|
||||
|
||||
# 标记为模拟数据
|
||||
df.attrs['data_source'] = '模拟数据'
|
||||
|
||||
return df
|
||||
|
||||
|
||||
def display_data_info(df, stock_code, start_year, end_year):
|
||||
"""显示数据信息"""
|
||||
if df is None or df.empty:
|
||||
print("没有数据可显示")
|
||||
return
|
||||
|
||||
# 获取数据来源
|
||||
data_source = df.attrs.get('data_source', '未知来源')
|
||||
|
||||
print(f"\n{'=' * 60}")
|
||||
print(f"股票 {stock_code} {start_year}-{end_year} 年数据摘要")
|
||||
print(f"{'=' * 60}")
|
||||
|
||||
print(f"数据时间范围: {df.index.min().strftime('%Y-%m-%d')} 到 {df.index.max().strftime('%Y-%m-%d')}")
|
||||
print(f"总交易天数: {len(df)}")
|
||||
print(f"数据来源: {data_source}")
|
||||
|
||||
# 按年份显示统计
|
||||
for year in sorted(df.index.year.unique()):
|
||||
year_data = df[df.index.year == year]
|
||||
print(f"\n{year}年统计:")
|
||||
print(f" 交易天数: {len(year_data)}")
|
||||
print(f" 平均收盘价: {year_data['收盘价'].mean():.2f} 元")
|
||||
print(f" 最高价: {year_data['最高价'].max():.2f} 元")
|
||||
print(f" 最低价: {year_data['最低价'].min():.2f} 元")
|
||||
if len(year_data) > 1:
|
||||
year_return = (year_data['收盘价'].iloc[-1] / year_data['收盘价'].iloc[0] - 1) * 100
|
||||
print(f" 年度涨跌幅: {year_return:+.2f}%")
|
||||
|
||||
# 显示最新交易日数据
|
||||
latest_date = df.index.max()
|
||||
print(f"\n最新交易日 ({latest_date.strftime('%Y-%m-%d')}) 数据:")
|
||||
latest_data = df.loc[latest_date]
|
||||
for col, value in latest_data.items():
|
||||
if col != '股票代码':
|
||||
if col in ['成交量']:
|
||||
print(f" {col}: {value:,.0f}")
|
||||
elif col in ['成交额']:
|
||||
print(f" {col}: {value:,.2f} 万元")
|
||||
else:
|
||||
print(f" {col}: {value}")
|
||||
|
||||
|
||||
def save_stock_data(df, stock_code, save_dir="D:/lianghuajiaoyi/Kronos/examples/data"):
|
||||
"""
|
||||
保存股票数据到指定目录
|
||||
"""
|
||||
if df is not None and not df.empty:
|
||||
# 确保保存目录存在
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
# 保存CSV文件
|
||||
csv_file = os.path.join(save_dir, f"{stock_code}_stock_data.csv")
|
||||
|
||||
# 重置索引以便保存日期列
|
||||
df_reset = df.reset_index()
|
||||
df_reset.to_csv(csv_file, encoding='utf-8-sig', index=False)
|
||||
|
||||
print(f"\n📁 股票数据已保存: {csv_file}")
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def main(stock_code="002354", start_year=2024, end_year=2025):
|
||||
"""
|
||||
主函数:获取并保存股票数据 - 最终版
|
||||
"""
|
||||
# 设置保存目录
|
||||
save_directory = "D:/lianghuajiaoyi/Kronos/examples/data"
|
||||
|
||||
print("=" * 60)
|
||||
print(f"开始获取股票 {stock_code} 的 {start_year}-{end_year} 年数据")
|
||||
print("=" * 60)
|
||||
print(f"数据将保存到: {save_directory}")
|
||||
|
||||
# 检查必要库
|
||||
try:
|
||||
import requests
|
||||
import numpy as np
|
||||
except ImportError:
|
||||
print("正在安装必要库...")
|
||||
import subprocess
|
||||
subprocess.check_call(["pip", "install", "requests", "numpy", "pandas"])
|
||||
import requests
|
||||
import numpy as np
|
||||
|
||||
# 获取数据(多数据源)
|
||||
stock_data = get_stock_data_with_retry(stock_code, start_year, end_year)
|
||||
|
||||
if stock_data is not None:
|
||||
# 显示数据信息
|
||||
display_data_info(stock_data, stock_code, start_year, end_year)
|
||||
|
||||
# 保存数据到指定目录
|
||||
save_stock_data(stock_data, stock_code, save_directory)
|
||||
|
||||
print(f"\n🎉 股票 {stock_code} 数据处理完成!")
|
||||
print(f"最新数据日期: {stock_data.index.max().strftime('%Y-%m-%d')}")
|
||||
|
||||
# 显示保存的文件
|
||||
csv_file = os.path.join(save_directory, f"{stock_code}_stock_data.csv")
|
||||
if os.path.exists(csv_file):
|
||||
file_size = os.path.getsize(csv_file) / 1024 # KB
|
||||
print(f"📄 生成的文件: {csv_file} ({file_size:.1f} KB)")
|
||||
else:
|
||||
print("❌ 未能获取股票数据")
|
||||
|
||||
|
||||
# 使用方法说明
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
使用方法:
|
||||
修改下面的参数来获取不同股票的数据
|
||||
"""
|
||||
|
||||
# ==================== 在这里修改参数 ====================
|
||||
TARGET_STOCK_CODE = "300418" # 股票代码
|
||||
START_YEAR = 2024 # 开始年份
|
||||
END_YEAR = 2025 # 结束年份
|
||||
# =====================================================
|
||||
|
||||
print("股票数据获取工具 - 终极优化版")
|
||||
print("说明:修改代码中的 TARGET_STOCK_CODE 来获取不同股票的数据")
|
||||
print(f"当前设置: 股票代码={TARGET_STOCK_CODE}, 年份范围={START_YEAR}-{END_YEAR}")
|
||||
print()
|
||||
|
||||
# 运行主程序
|
||||
main(stock_code=TARGET_STOCK_CODE, start_year=START_YEAR, end_year=END_YEAR)
|
||||
|
||||
print(f"\n💡 提示:要获取其他股票数据,请修改代码中的 TARGET_STOCK_CODE 变量")
|
||||
@@ -0,0 +1,661 @@
|
||||
import pandas as pd
|
||||
import requests
|
||||
import json
|
||||
from datetime import datetime, timedelta
|
||||
import os
|
||||
import time
|
||||
import random
|
||||
|
||||
|
||||
def get_stock_market(stock_code):
|
||||
"""
|
||||
根据股票代码判断市场类型
|
||||
返回: 市场前缀 '0'-深交所, '1'-上交所
|
||||
"""
|
||||
if stock_code.startswith(('0', '2', '3')):
|
||||
return '0' # 深交所
|
||||
elif stock_code.startswith(('6', '9')):
|
||||
return '1' # 上交所
|
||||
else:
|
||||
return '1' # 默认上交所
|
||||
|
||||
|
||||
def get_stock_data_eastmoney_all_history(stock_code="002354"):
|
||||
"""
|
||||
使用东方财富网API获取股票所有历史数据
|
||||
"""
|
||||
try:
|
||||
print(f"正在从东方财富网获取股票 {stock_code} 的全部历史数据...")
|
||||
|
||||
# 获取市场类型
|
||||
market = get_stock_market(stock_code)
|
||||
secid = f"{market}.{stock_code}"
|
||||
|
||||
# 使用东方财富API获取所有历史数据
|
||||
url = "http://push2his.eastmoney.com/api/qt/stock/kline/get"
|
||||
|
||||
# 设置足够早的起始日期(中国股市从1990年开始)
|
||||
start_date = "19900101"
|
||||
end_date = datetime.now().strftime('%Y%m%d')
|
||||
|
||||
params = {
|
||||
'secid': secid,
|
||||
'fields1': 'f1,f2,f3,f4,f5,f6',
|
||||
'fields2': 'f51,f52,f53,f54,f55,f56,f57,f58,f59,f60,f61',
|
||||
'klt': '101', # 日线
|
||||
'fqt': '1', # 前复权
|
||||
'beg': start_date,
|
||||
'end': end_date,
|
||||
'lmt': '50000', # 增加限制数量以获取更多历史数据
|
||||
'ut': 'fa5fd1943c7b386f172d6893dbfba10b',
|
||||
'cb': f'jQuery{random.randint(1000000, 9999999)}_{int(time.time() * 1000)}'
|
||||
}
|
||||
|
||||
headers = {
|
||||
'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/117.0.0.0 Safari/537.36',
|
||||
'Referer': 'https://quote.eastmoney.com/',
|
||||
'Accept': '*/*',
|
||||
}
|
||||
|
||||
time.sleep(random.uniform(1, 2))
|
||||
|
||||
response = requests.get(url, params=params, headers=headers, timeout=15)
|
||||
|
||||
print(f"API响应状态码: {response.status_code}")
|
||||
|
||||
if response.status_code == 200:
|
||||
# 处理JSONP响应
|
||||
response_text = response.text
|
||||
|
||||
# 提取JSON数据(处理JSONP格式)
|
||||
if response_text.startswith('/**/'):
|
||||
response_text = response_text[4:]
|
||||
|
||||
# 查找JSON数据的开始和结束位置
|
||||
start_idx = response_text.find('(')
|
||||
end_idx = response_text.rfind(')')
|
||||
|
||||
if start_idx != -1 and end_idx != -1:
|
||||
json_str = response_text[start_idx + 1:end_idx]
|
||||
try:
|
||||
data = json.loads(json_str)
|
||||
except json.JSONDecodeError:
|
||||
print("❌ JSON解析失败,尝试直接解析...")
|
||||
return parse_kline_data_directly_all_history(response_text, stock_code)
|
||||
else:
|
||||
print("❌ 无法找到JSON数据边界")
|
||||
return None
|
||||
|
||||
print(f"API返回数据状态: {data.get('rc', 'N/A')}")
|
||||
|
||||
if data and data.get('data') is not None:
|
||||
klines = data['data'].get('klines', [])
|
||||
print(f"获取到 {len(klines)} 条历史K线数据")
|
||||
|
||||
if not klines:
|
||||
print("⚠️ K线数据为空")
|
||||
return None
|
||||
|
||||
# 解析数据
|
||||
stock_data = []
|
||||
for kline in klines:
|
||||
try:
|
||||
items = kline.split(',')
|
||||
if len(items) >= 6:
|
||||
stock_data.append({
|
||||
'日期': items[0],
|
||||
'股票代码': stock_code,
|
||||
'开盘价': float(items[1]),
|
||||
'收盘价': float(items[2]),
|
||||
'最高价': float(items[3]),
|
||||
'最低价': float(items[4]),
|
||||
'成交量': float(items[5]),
|
||||
'成交额': float(items[6]) if len(items) > 6 else 0,
|
||||
'振幅': float(items[7]) if len(items) > 7 else 0,
|
||||
'涨跌幅': float(items[8]) if len(items) > 8 else 0,
|
||||
'涨跌额': float(items[9]) if len(items) > 9 else 0,
|
||||
'换手率': float(items[10]) if len(items) > 10 else 0
|
||||
})
|
||||
except (ValueError, IndexError) as e:
|
||||
continue
|
||||
|
||||
if not stock_data:
|
||||
print("❌ 解析后无有效数据")
|
||||
return None
|
||||
|
||||
df = pd.DataFrame(stock_data)
|
||||
df['日期'] = pd.to_datetime(df['日期'])
|
||||
df.set_index('日期', inplace=True)
|
||||
df = df.sort_index()
|
||||
|
||||
print(f"✅ 成功获取 {len(df)} 条历史数据")
|
||||
print(
|
||||
f"历史数据时间范围: {df.index.min().strftime('%Y-%m-%d')} 到 {df.index.max().strftime('%Y-%m-%d')}")
|
||||
return df
|
||||
else:
|
||||
print("❌ API返回数据为空")
|
||||
return None
|
||||
else:
|
||||
print(f"❌ 请求失败,状态码: {response.status_code}")
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ 获取历史数据时出错: {str(e)}")
|
||||
return None
|
||||
|
||||
|
||||
def parse_kline_data_directly_all_history(response_text, stock_code):
|
||||
"""
|
||||
直接解析K线数据(当JSON解析失败时使用)- 全历史版本
|
||||
"""
|
||||
try:
|
||||
# 尝试直接从响应文本中提取K线数据
|
||||
if '"klines":[' in response_text:
|
||||
start_idx = response_text.find('"klines":[') + 10
|
||||
end_idx = response_text.find(']', start_idx)
|
||||
klines_str = response_text[start_idx:end_idx]
|
||||
|
||||
# 清理字符串并分割
|
||||
klines = [k.strip().strip('"') for k in klines_str.split('","') if k.strip()]
|
||||
|
||||
stock_data = []
|
||||
for kline in klines:
|
||||
if kline.strip():
|
||||
items = kline.split(',')
|
||||
if len(items) >= 6:
|
||||
stock_data.append({
|
||||
'日期': items[0],
|
||||
'股票代码': stock_code,
|
||||
'开盘价': float(items[1]),
|
||||
'收盘价': float(items[2]),
|
||||
'最高价': float(items[3]),
|
||||
'最低价': float(items[4]),
|
||||
'成交量': float(items[5]),
|
||||
'成交额': float(items[6]) if len(items) > 6 else 0,
|
||||
})
|
||||
|
||||
if stock_data:
|
||||
df = pd.DataFrame(stock_data)
|
||||
df['日期'] = pd.to_datetime(df['日期'])
|
||||
df.set_index('日期', inplace=True)
|
||||
df = df.sort_index()
|
||||
print(f"✅ 直接解析获取 {len(df)} 条历史数据")
|
||||
return df
|
||||
except Exception as e:
|
||||
print(f"❌ 直接解析也失败: {e}")
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def get_stock_data_akshare_all_history(stock_code="002354"):
|
||||
"""
|
||||
使用AKShare作为备用数据源 - 全历史版本
|
||||
"""
|
||||
try:
|
||||
print(f"尝试使用AKShare获取股票 {stock_code} 全部历史数据...")
|
||||
import akshare as ak
|
||||
|
||||
# 获取所有历史数据
|
||||
df = ak.stock_zh_a_hist(symbol=stock_code, period="daily",
|
||||
adjust="qfq")
|
||||
|
||||
if df is not None and not df.empty:
|
||||
# 重命名列以匹配我们的格式
|
||||
column_mapping = {
|
||||
'日期': '日期',
|
||||
'开盘': '开盘价',
|
||||
'收盘': '收盘价',
|
||||
'最高': '最高价',
|
||||
'最低': '最低价',
|
||||
'成交量': '成交量',
|
||||
'成交额': '成交额',
|
||||
'振幅': '振幅',
|
||||
'涨跌幅': '涨跌幅',
|
||||
'涨跌额': '涨跌额',
|
||||
'换手率': '换手率'
|
||||
}
|
||||
|
||||
# 只映射存在的列
|
||||
actual_mapping = {k: v for k, v in column_mapping.items() if k in df.columns}
|
||||
df = df.rename(columns=actual_mapping)
|
||||
|
||||
# 添加股票代码列
|
||||
df['股票代码'] = stock_code
|
||||
df['日期'] = pd.to_datetime(df['日期'])
|
||||
df.set_index('日期', inplace=True)
|
||||
df = df.sort_index()
|
||||
|
||||
print(f"✅ AKShare成功获取 {len(df)} 条历史数据")
|
||||
print(f"时间范围: {df.index.min().strftime('%Y-%m-%d')} 到 {df.index.max().strftime('%Y-%m-%d')}")
|
||||
return df
|
||||
else:
|
||||
print("❌ AKShare未返回数据")
|
||||
return None
|
||||
|
||||
except ImportError:
|
||||
print("⚠️ AKShare未安装,使用 pip install akshare 安装")
|
||||
return None
|
||||
except Exception as e:
|
||||
print(f"❌ AKShare获取历史数据失败: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def get_stock_data_baostock_all_history(stock_code="002354"):
|
||||
"""
|
||||
使用Baostock作为第三个数据源 - 全历史版本
|
||||
"""
|
||||
try:
|
||||
print(f"尝试使用Baostock获取股票 {stock_code} 全部历史数据...")
|
||||
import baostock as bs
|
||||
import pandas as pd
|
||||
|
||||
# 登录系统
|
||||
lg = bs.login()
|
||||
|
||||
# 根据市场添加前缀
|
||||
market = get_stock_market(stock_code)
|
||||
if market == '0':
|
||||
full_code = f"sz.{stock_code}"
|
||||
else:
|
||||
full_code = f"sh.{stock_code}"
|
||||
|
||||
# 获取上市日期
|
||||
rs = bs.query_stock_basic(code=full_code)
|
||||
if rs.error_code != '0':
|
||||
print(f"❌ 获取股票基本信息失败: {rs.error_msg}")
|
||||
bs.logout()
|
||||
return None
|
||||
|
||||
# 获取上市日期
|
||||
list_date = None
|
||||
while (rs.error_code == '0') & rs.next():
|
||||
list_date = rs.get_row_data()[2] # 上市日期在第三个字段
|
||||
|
||||
if not list_date:
|
||||
print("❌ 无法获取上市日期")
|
||||
bs.logout()
|
||||
return None
|
||||
|
||||
print(f"股票上市日期: {list_date}")
|
||||
|
||||
# 获取从上市日期到现在的所有数据
|
||||
end_date = datetime.now().strftime('%Y-%m-%d')
|
||||
|
||||
# 获取数据
|
||||
rs = bs.query_history_k_data_plus(
|
||||
full_code,
|
||||
"date,open,high,low,close,volume,amount,turn,pctChg",
|
||||
start_date=list_date,
|
||||
end_date=end_date,
|
||||
frequency="d",
|
||||
adjustflag="2" # 前复权
|
||||
)
|
||||
|
||||
data_list = []
|
||||
while (rs.error_code == '0') & rs.next():
|
||||
data_list.append(rs.get_row_data())
|
||||
|
||||
# 退出系统
|
||||
bs.logout()
|
||||
|
||||
if data_list:
|
||||
df = pd.DataFrame(data_list, columns=rs.fields)
|
||||
|
||||
# 数据类型转换
|
||||
df['date'] = pd.to_datetime(df['date'])
|
||||
df['open'] = pd.to_numeric(df['open'], errors='coerce')
|
||||
df['high'] = pd.to_numeric(df['high'], errors='coerce')
|
||||
df['low'] = pd.to_numeric(df['low'], errors='coerce')
|
||||
df['close'] = pd.to_numeric(df['close'], errors='coerce')
|
||||
df['volume'] = pd.to_numeric(df['volume'], errors='coerce')
|
||||
df['amount'] = pd.to_numeric(df['amount'], errors='coerce')
|
||||
df['turn'] = pd.to_numeric(df['turn'], errors='coerce')
|
||||
df['pctChg'] = pd.to_numeric(df['pctChg'], errors='coerce')
|
||||
|
||||
# 重命名列
|
||||
df = df.rename(columns={
|
||||
'date': '日期',
|
||||
'open': '开盘价',
|
||||
'high': '最高价',
|
||||
'low': '最低价',
|
||||
'close': '收盘价',
|
||||
'volume': '成交量',
|
||||
'amount': '成交额',
|
||||
'turn': '换手率',
|
||||
'pctChg': '涨跌幅'
|
||||
})
|
||||
|
||||
# 添加股票代码列
|
||||
df['股票代码'] = stock_code
|
||||
df.set_index('日期', inplace=True)
|
||||
df = df.sort_index()
|
||||
|
||||
# 计算涨跌额
|
||||
df['涨跌额'] = df['收盘价'].diff()
|
||||
|
||||
# 清理无效数据
|
||||
df = df.dropna()
|
||||
|
||||
print(f"✅ Baostock成功获取 {len(df)} 条历史数据")
|
||||
print(f"时间范围: {df.index.min().strftime('%Y-%m-%d')} 到 {df.index.max().strftime('%Y-%m-%d')}")
|
||||
return df
|
||||
else:
|
||||
print("❌ Baostock未返回数据")
|
||||
return None
|
||||
|
||||
except ImportError:
|
||||
print("⚠️ Baostock未安装,使用 pip install baostock 安装")
|
||||
return None
|
||||
except Exception as e:
|
||||
print(f"❌ Baostock获取历史数据失败: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def get_stock_data_with_retry_all_history(stock_code="002354", retry_count=2):
|
||||
"""
|
||||
带重试机制的数据获取 - 多数据源全历史版本
|
||||
"""
|
||||
data_sources = [
|
||||
("AKShare", get_stock_data_akshare_all_history),
|
||||
("Baostock", get_stock_data_baostock_all_history),
|
||||
("东方财富", get_stock_data_eastmoney_all_history)
|
||||
]
|
||||
|
||||
for source_name, data_func in data_sources:
|
||||
print(f"\n🔍 尝试从 {source_name} 获取全部历史数据...")
|
||||
data = data_func(stock_code)
|
||||
|
||||
if data is not None and not data.empty:
|
||||
print(f"✅ {source_name} 历史数据获取成功!")
|
||||
# 标记数据来源
|
||||
data.attrs['data_source'] = source_name
|
||||
return data
|
||||
|
||||
print("❌ 所有真实数据源都失败,使用示例数据...")
|
||||
return create_sample_data_all_history(stock_code)
|
||||
|
||||
|
||||
def create_sample_data_all_history(stock_code="002354"):
|
||||
"""
|
||||
创建更真实的历史示例数据 - 从上市年份开始
|
||||
"""
|
||||
# 模拟不同股票的上市年份
|
||||
list_years = {
|
||||
'600580': 2002, # 卧龙电驱
|
||||
'002354': 2010, # 天娱数科
|
||||
'300418': 2015, # 昆仑万维
|
||||
'300207': 2011, # 欣旺达
|
||||
}
|
||||
|
||||
list_year = list_years.get(stock_code, 2010)
|
||||
current_year = datetime.now().year
|
||||
|
||||
print(f"📊 创建 {stock_code} 从 {list_year} 年上市至今的示例数据...")
|
||||
|
||||
# 生成从上市年份到现在的交易日(排除周末)
|
||||
start_date = datetime(list_year, 1, 1)
|
||||
end_date = datetime.now()
|
||||
all_dates = pd.bdate_range(start=start_date, end=end_date, freq='B')
|
||||
|
||||
# 生成更真实的股价数据
|
||||
import numpy as np
|
||||
np.random.seed(42)
|
||||
|
||||
# 设置合理的基准价格(根据股票类型)
|
||||
base_prices = {
|
||||
'600580': 8.0, # 卧龙电驱
|
||||
'002354': 15.0, # 天娱数科 - 上市时价格较高
|
||||
'300418': 20.0, # 昆仑万维
|
||||
'300207': 12.0, # 欣旺达
|
||||
}
|
||||
base_price = base_prices.get(stock_code, 10.0)
|
||||
|
||||
stock_data = []
|
||||
current_price = base_price
|
||||
|
||||
for i, date in enumerate(all_dates):
|
||||
# 模拟真实的市场波动
|
||||
volatility = 0.02 # 2%的日波动率
|
||||
|
||||
if i > 0:
|
||||
# 使用随机游走模拟价格变化
|
||||
daily_return = np.random.normal(0, volatility)
|
||||
|
||||
# 模拟不同年份的市场趋势
|
||||
year = date.year
|
||||
if year <= list_year + 2: # 上市初期波动较大
|
||||
daily_return += np.random.normal(0.001, 0.01)
|
||||
elif year <= list_year + 5: # 成长期
|
||||
daily_return += np.random.normal(0.0005, 0.005)
|
||||
else: # 成熟期
|
||||
daily_return += np.random.normal(0.0002, 0.003)
|
||||
|
||||
current_price = current_price * (1 + daily_return)
|
||||
|
||||
# 价格边界限制
|
||||
current_price = max(base_price * 0.3, min(base_price * 10.0, current_price))
|
||||
else:
|
||||
current_price = base_price
|
||||
|
||||
# 生成OHLC数据
|
||||
open_variation = np.random.normal(0, volatility * 0.2)
|
||||
open_price = current_price * (1 + open_variation)
|
||||
|
||||
daily_range = abs(np.random.normal(volatility * 0.8, volatility * 0.3))
|
||||
high_price = max(open_price, current_price) * (1 + daily_range)
|
||||
low_price = min(open_price, current_price) * (1 - daily_range)
|
||||
close_price = current_price
|
||||
|
||||
# 确保价格合理性
|
||||
high_price = max(open_price, close_price, low_price, high_price)
|
||||
low_price = min(open_price, close_price, high_price, low_price)
|
||||
|
||||
# 生成成交量(随年份增长)
|
||||
base_volume = 100000 + (year - list_year) * 50000 # 成交量逐年增长
|
||||
volume_variation = abs(daily_return) * 5000000 if i > 0 else 0
|
||||
volume = int(base_volume + volume_variation + np.random.randint(-200000, 400000))
|
||||
volume = max(50000, volume)
|
||||
|
||||
# 计算成交额(万元)
|
||||
amount = volume * close_price / 10000
|
||||
|
||||
# 计算涨跌幅和涨跌额
|
||||
if i > 0:
|
||||
prev_close = stock_data[-1]['收盘价']
|
||||
price_change = close_price - prev_close
|
||||
pct_change = (price_change / prev_close) * 100
|
||||
else:
|
||||
price_change = 0
|
||||
pct_change = 0
|
||||
|
||||
# 计算振幅
|
||||
amplitude = ((high_price - low_price) / open_price) * 100
|
||||
|
||||
# 生成换手率(1%-15%之间)
|
||||
turnover_rate = np.random.uniform(1.0, 15.0)
|
||||
|
||||
stock_data.append({
|
||||
'日期': date,
|
||||
'股票代码': stock_code,
|
||||
'开盘价': round(open_price, 2),
|
||||
'收盘价': round(close_price, 2),
|
||||
'最高价': round(high_price, 2),
|
||||
'最低价': round(low_price, 2),
|
||||
'成交量': volume,
|
||||
'成交额': round(amount, 2),
|
||||
'振幅': round(amplitude, 2),
|
||||
'涨跌幅': round(pct_change, 2),
|
||||
'涨跌额': round(price_change, 2),
|
||||
'换手率': round(turnover_rate, 2)
|
||||
})
|
||||
|
||||
df = pd.DataFrame(stock_data)
|
||||
df.set_index('日期', inplace=True)
|
||||
|
||||
print(f"✅ 已创建 {len(df)} 条从 {list_year} 年至今的模拟历史数据")
|
||||
print(f"时间范围: {df.index.min().strftime('%Y-%m-%d')} 到 {df.index.max().strftime('%Y-%m-%d')}")
|
||||
|
||||
# 标记为模拟数据
|
||||
df.attrs['data_source'] = '模拟历史数据'
|
||||
|
||||
return df
|
||||
|
||||
|
||||
def display_all_history_data_info(df, stock_code):
|
||||
"""显示全历史数据信息"""
|
||||
if df is None or df.empty:
|
||||
print("没有数据可显示")
|
||||
return
|
||||
|
||||
# 获取数据来源
|
||||
data_source = df.attrs.get('data_source', '未知来源')
|
||||
|
||||
print(f"\n{'=' * 60}")
|
||||
print(f"股票 {stock_code} 全部历史数据摘要")
|
||||
print(f"{'=' * 60}")
|
||||
|
||||
print(f"数据时间范围: {df.index.min().strftime('%Y-%m-%d')} 到 {df.index.max().strftime('%Y-%m-%d')}")
|
||||
print(f"总交易天数: {len(df):,}")
|
||||
print(f"数据来源: {data_source}")
|
||||
|
||||
# 按年份显示统计
|
||||
years = sorted(df.index.year.unique())
|
||||
print(f"\n历史年份: {years}")
|
||||
|
||||
# 显示关键年份统计
|
||||
key_years = [years[0]] # 上市年份
|
||||
if len(years) > 1:
|
||||
key_years.append(years[-1]) # 最新年份
|
||||
if len(years) > 5:
|
||||
key_years.extend([years[len(years) // 2], years[len(years) // 4], years[3 * len(years) // 4]])
|
||||
|
||||
for year in sorted(set(key_years)):
|
||||
year_data = df[df.index.year == year]
|
||||
if len(year_data) > 0:
|
||||
print(f"\n{year}年统计:")
|
||||
print(f" 交易天数: {len(year_data)}")
|
||||
print(f" 平均收盘价: {year_data['收盘价'].mean():.2f} 元")
|
||||
print(f" 最高价: {year_data['最高价'].max():.2f} 元")
|
||||
print(f" 最低价: {year_data['最低价'].min():.2f} 元")
|
||||
if len(year_data) > 1:
|
||||
year_return = (year_data['收盘价'].iloc[-1] / year_data['收盘价'].iloc[0] - 1) * 100
|
||||
print(f" 年度涨跌幅: {year_return:+.2f}%")
|
||||
|
||||
# 显示整体统计
|
||||
print(f"\n整体统计:")
|
||||
total_return = (df['收盘价'].iloc[-1] / df['收盘价'].iloc[0] - 1) * 100
|
||||
print(f" 总涨跌幅: {total_return:+.2f}%")
|
||||
print(f" 历史最高价: {df['最高价'].max():.2f} 元")
|
||||
print(f" 历史最低价: {df['最低价'].min():.2f} 元")
|
||||
print(f" 平均日成交量: {df['成交量'].mean():,.0f} 股")
|
||||
|
||||
# 显示最新交易日数据
|
||||
latest_date = df.index.max()
|
||||
print(f"\n最新交易日 ({latest_date.strftime('%Y-%m-%d')}) 数据:")
|
||||
latest_data = df.loc[latest_date]
|
||||
for col, value in latest_data.items():
|
||||
if col != '股票代码':
|
||||
if col in ['成交量']:
|
||||
print(f" {col}: {value:,.0f}")
|
||||
elif col in ['成交额']:
|
||||
print(f" {col}: {value:,.2f} 万元")
|
||||
else:
|
||||
print(f" {col}: {value}")
|
||||
|
||||
|
||||
def save_all_history_stock_data(df, stock_code, save_dir="D:/lianghuajiaoyi/Kronos/examples/data"):
|
||||
"""
|
||||
保存全历史股票数据到指定目录
|
||||
"""
|
||||
if df is not None and not df.empty:
|
||||
# 确保保存目录存在
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
# 保存CSV文件 - 使用全历史命名
|
||||
csv_file = os.path.join(save_dir, f"{stock_code}_all_history.csv")
|
||||
|
||||
# 重置索引以便保存日期列
|
||||
df_reset = df.reset_index()
|
||||
df_reset.to_csv(csv_file, encoding='utf-8-sig', index=False)
|
||||
|
||||
print(f"\n📁 全历史股票数据已保存: {csv_file}")
|
||||
|
||||
# 同时保存一个按年份分割的版本
|
||||
years = df_reset['日期'].dt.year.unique()
|
||||
for year in years:
|
||||
year_data = df_reset[df_reset['日期'].dt.year == year]
|
||||
year_file = os.path.join(save_dir, f"{stock_code}_{year}.csv")
|
||||
year_data.to_csv(year_file, encoding='utf-8-sig', index=False)
|
||||
|
||||
print(f"📁 同时保存了 {len(years)} 个年份的单独数据文件")
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def main_all_history(stock_code="002354"):
|
||||
"""
|
||||
主函数:获取并保存股票全历史数据
|
||||
"""
|
||||
# 设置保存目录
|
||||
save_directory = "D:/lianghuajiaoyi/Kronos/examples/data"
|
||||
|
||||
print("=" * 60)
|
||||
print(f"开始获取股票 {stock_code} 的全部历史数据")
|
||||
print("=" * 60)
|
||||
print(f"数据将保存到: {save_directory}")
|
||||
|
||||
# 检查必要库
|
||||
try:
|
||||
import requests
|
||||
import numpy as np
|
||||
except ImportError:
|
||||
print("正在安装必要库...")
|
||||
import subprocess
|
||||
subprocess.check_call(["pip", "install", "requests", "numpy", "pandas"])
|
||||
import requests
|
||||
import numpy as np
|
||||
|
||||
# 获取全历史数据(多数据源)
|
||||
stock_data = get_stock_data_with_retry_all_history(stock_code)
|
||||
|
||||
if stock_data is not None:
|
||||
# 显示数据信息
|
||||
display_all_history_data_info(stock_data, stock_code)
|
||||
|
||||
# 保存全历史数据到指定目录
|
||||
save_all_history_stock_data(stock_data, stock_code, save_directory)
|
||||
|
||||
print(f"\n🎉 股票 {stock_code} 全历史数据处理完成!")
|
||||
print(
|
||||
f"数据时间跨度: {stock_data.index.min().strftime('%Y-%m-%d')} 到 {stock_data.index.max().strftime('%Y-%m-%d')}")
|
||||
print(f"总交易天数: {len(stock_data):,}")
|
||||
|
||||
# 显示保存的文件
|
||||
csv_file = os.path.join(save_directory, f"{stock_code}_all_history.csv")
|
||||
if os.path.exists(csv_file):
|
||||
file_size = os.path.getsize(csv_file) / 1024 # KB
|
||||
print(f"📄 生成的文件: {csv_file} ({file_size:.1f} KB)")
|
||||
else:
|
||||
print("❌ 未能获取股票全历史数据")
|
||||
|
||||
|
||||
# 使用方法说明
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
使用方法:
|
||||
修改下面的参数来获取不同股票的全历史数据
|
||||
"""
|
||||
|
||||
# ==================== 在这里修改参数 ====================
|
||||
TARGET_STOCK_CODE = "300418" # 股票代码
|
||||
# =====================================================
|
||||
|
||||
print("股票全历史数据获取工具")
|
||||
print("说明:修改代码中的 TARGET_STOCK_CODE 来获取不同股票的全部历史数据")
|
||||
print(f"当前设置: 股票代码={TARGET_STOCK_CODE}")
|
||||
print()
|
||||
|
||||
# 运行主程序
|
||||
main_all_history(stock_code=TARGET_STOCK_CODE)
|
||||
|
||||
print(f"\n💡 提示:要获取其他股票的全历史数据,请修改代码中的 TARGET_STOCK_CODE 变量")
|
||||
@@ -0,0 +1,545 @@
|
||||
import pandas as pd
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import sys
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('ignore')
|
||||
|
||||
# 添加项目路径以便导入自定义模块
|
||||
sys.path.append("../")
|
||||
from model import Kronos, KronosTokenizer, KronosPredictor
|
||||
|
||||
# 设置中文字体
|
||||
plt.rcParams['font.sans-serif'] = ['SimHei'] # 用来正常显示中文标签
|
||||
plt.rcParams['axes.unicode_minus'] = False # 用来正常显示负号
|
||||
|
||||
|
||||
def ensure_output_directory(output_dir):
|
||||
"""确保输出目录存在,如果不存在则创建"""
|
||||
if not os.path.exists(output_dir):
|
||||
os.makedirs(output_dir)
|
||||
print(f"✅ 创建输出目录: {output_dir}")
|
||||
return output_dir
|
||||
|
||||
|
||||
def prepare_stock_data(csv_file_path, stock_code):
|
||||
"""
|
||||
准备股票数据,转换为Kronos模型需要的格式
|
||||
|
||||
参数:
|
||||
csv_file_path: CSV文件路径
|
||||
stock_code: 股票代码,用于显示信息
|
||||
|
||||
返回:
|
||||
df: 处理后的DataFrame
|
||||
"""
|
||||
print(f"正在加载和预处理股票 {stock_code} 数据...")
|
||||
|
||||
# 读取CSV文件
|
||||
df = pd.read_csv(csv_file_path, encoding='utf-8-sig')
|
||||
|
||||
# 检查数据列名并重命名为标准格式
|
||||
column_mapping = {
|
||||
'日期': 'timestamps',
|
||||
'开盘价': 'open',
|
||||
'最高价': 'high',
|
||||
'最低价': 'low',
|
||||
'收盘价': 'close',
|
||||
'成交量': 'volume',
|
||||
'成交额': 'amount'
|
||||
}
|
||||
|
||||
# 只重命名存在的列
|
||||
actual_mapping = {k: v for k, v in column_mapping.items() if k in df.columns}
|
||||
df = df.rename(columns=actual_mapping)
|
||||
|
||||
# 确保时间戳列存在并转换为datetime格式
|
||||
if 'timestamps' not in df.columns:
|
||||
# 如果数据有日期索引,重置索引
|
||||
if df.index.name == '日期':
|
||||
df = df.reset_index()
|
||||
df = df.rename(columns={'日期': 'timestamps'})
|
||||
|
||||
df['timestamps'] = pd.to_datetime(df['timestamps'])
|
||||
|
||||
# 按时间排序
|
||||
df = df.sort_values('timestamps').reset_index(drop=True)
|
||||
|
||||
print(f"✅ 数据加载完成,共 {len(df)} 条记录")
|
||||
print(f"时间范围: {df['timestamps'].min()} 到 {df['timestamps'].max()}")
|
||||
print(f"数据列: {df.columns.tolist()}")
|
||||
|
||||
return df
|
||||
|
||||
|
||||
def calculate_prediction_parameters(df, target_days=100):
|
||||
"""
|
||||
根据目标预测天数计算合适的参数
|
||||
|
||||
参数:
|
||||
df: 股票数据DataFrame
|
||||
target_days: 目标预测天数(自然日)
|
||||
|
||||
返回:
|
||||
lookback: 回看期数
|
||||
pred_len: 预测期数
|
||||
"""
|
||||
# 计算平均交易日数量(考虑节假日)
|
||||
total_days = (df['timestamps'].max() - df['timestamps'].min()).days
|
||||
trading_days = len(df)
|
||||
trading_ratio = trading_days / total_days if total_days > 0 else 0.7 # 交易日比例
|
||||
|
||||
# 计算目标预测的交易日数量
|
||||
pred_trading_days = int(target_days * trading_ratio)
|
||||
|
||||
# 设置回看期数为预测期数的2-3倍,但不超过数据总量的70%
|
||||
max_lookback = int(len(df) * 0.7)
|
||||
lookback = min(pred_trading_days * 2, max_lookback, len(df) - pred_trading_days)
|
||||
pred_len = min(pred_trading_days, len(df) - lookback)
|
||||
|
||||
print(f"📊 参数计算:")
|
||||
print(f" 目标预测天数: {target_days} 天(自然日)")
|
||||
print(f" 预计交易日数量: {pred_trading_days} 天")
|
||||
print(f" 回看期数 (lookback): {lookback}")
|
||||
print(f" 预测期数 (pred_len): {pred_len}")
|
||||
|
||||
return lookback, pred_len
|
||||
|
||||
|
||||
def generate_future_dates_with_holidays(last_date, pred_len):
|
||||
"""
|
||||
生成未来的交易日日期,考虑中国节假日
|
||||
|
||||
参数:
|
||||
last_date: 最后一个历史数据的日期
|
||||
pred_len: 预测期数
|
||||
|
||||
返回:
|
||||
future_dates: 未来的交易日日期列表
|
||||
"""
|
||||
# 中国主要节假日(需要根据实际情况调整)
|
||||
holidays_2025 = [
|
||||
# 2025年国庆节假期(通常为10月1日-10月8日)
|
||||
datetime(2025, 10, 1), datetime(2025, 10, 2), datetime(2025, 10, 3),
|
||||
datetime(2025, 10, 4), datetime(2025, 10, 5), datetime(2025, 10, 6),
|
||||
datetime(2025, 10, 7), datetime(2025, 10, 8), # 添加10月8日
|
||||
# 周末调休等可以根据需要添加
|
||||
]
|
||||
|
||||
future_dates = []
|
||||
current_date = last_date + timedelta(days=1)
|
||||
|
||||
while len(future_dates) < pred_len:
|
||||
# 如果是工作日(周一到周五)且不是节假日
|
||||
if current_date.weekday() < 5 and current_date not in holidays_2025:
|
||||
future_dates.append(current_date)
|
||||
current_date += timedelta(days=1)
|
||||
|
||||
print(f"📅 生成的未来交易日: 共 {len(future_dates)} 天")
|
||||
print(f" 起始日期: {future_dates[0].strftime('%Y-%m-%d')}")
|
||||
print(f" 结束日期: {future_dates[-1].strftime('%Y-%m-%d')}")
|
||||
|
||||
# 显示节假日信息
|
||||
holiday_count = sum(1 for date in holidays_2025 if date > last_date)
|
||||
print(f" 包含节假日: {holiday_count} 天")
|
||||
|
||||
return future_dates[:pred_len]
|
||||
|
||||
|
||||
def plot_prediction_with_details(kline_df, pred_df, future_dates, stock_code="002354", stock_name="股票", pred_len=100,
|
||||
output_dir="."):
|
||||
"""
|
||||
绘制详细的预测结果图表 - 优化版,图表更大更清晰
|
||||
|
||||
参数:
|
||||
kline_df: 历史K线数据
|
||||
pred_df: 预测数据
|
||||
future_dates: 未来日期列表
|
||||
stock_code: 股票代码
|
||||
stock_name: 股票名称
|
||||
pred_len: 预测期数
|
||||
output_dir: 输出目录
|
||||
"""
|
||||
# 确保输出目录存在
|
||||
ensure_output_directory(output_dir)
|
||||
|
||||
# 确保数据长度一致
|
||||
min_len = min(len(pred_df), len(future_dates))
|
||||
pred_df = pred_df.iloc[:min_len]
|
||||
future_dates = future_dates[:min_len]
|
||||
|
||||
# 设置预测数据的索引为未来日期
|
||||
pred_df.index = future_dates
|
||||
|
||||
# 准备价格数据
|
||||
sr_close = kline_df.set_index('timestamps')['close']
|
||||
sr_pred_close = pred_df['close']
|
||||
sr_close.name = '历史数据'
|
||||
sr_pred_close.name = "预测数据"
|
||||
|
||||
# 准备成交量数据
|
||||
sr_volume = kline_df.set_index('timestamps')['volume']
|
||||
sr_pred_volume = pred_df['volume']
|
||||
sr_volume.name = '历史数据'
|
||||
sr_pred_volume.name = "预测数据"
|
||||
|
||||
# 合并数据
|
||||
close_df = pd.concat([sr_close, sr_pred_close], axis=1)
|
||||
volume_df = pd.concat([sr_volume, sr_pred_volume], axis=1)
|
||||
|
||||
# 创建更大的图表
|
||||
fig = plt.figure(figsize=(18, 14))
|
||||
|
||||
# 使用GridSpec创建更灵活的布局
|
||||
gs = plt.GridSpec(3, 1, figure=fig, height_ratios=[3, 1, 1])
|
||||
|
||||
ax1 = fig.add_subplot(gs[0]) # 价格图表
|
||||
ax2 = fig.add_subplot(gs[1]) # 成交量图表
|
||||
ax3 = fig.add_subplot(gs[2]) # 价格变动图表
|
||||
|
||||
# 1. 价格图表 - 更大更清晰
|
||||
# 只显示最近200个交易日的历史数据,避免图表过于拥挤
|
||||
recent_history = close_df['历史数据'].iloc[-min(200, len(close_df['历史数据'])):]
|
||||
ax1.plot(recent_history.index, recent_history.values, label='历史价格', color='#1f77b4', linewidth=2.5, alpha=0.9)
|
||||
ax1.plot(close_df['预测数据'].index, close_df['预测数据'].values, label='预测价格',
|
||||
color='#ff7f0e', linewidth=2.5, linestyle='-', marker='o', markersize=3)
|
||||
|
||||
# 添加预测起始点的标记
|
||||
prediction_start_date = close_df['预测数据'].index[0] if len(close_df['预测数据']) > 0 else close_df.index[-1]
|
||||
prediction_start_price = close_df['历史数据'].iloc[-1]
|
||||
ax1.axvline(x=prediction_start_date, color='red', linestyle='--', alpha=0.7, linewidth=1.5)
|
||||
ax1.annotate('预测起点', xy=(prediction_start_date, prediction_start_price),
|
||||
xytext=(10, 10), textcoords='offset points',
|
||||
bbox=dict(boxstyle='round,pad=0.3', facecolor='yellow', alpha=0.7),
|
||||
arrowprops=dict(arrowstyle='->', connectionstyle='arc3,rad=0'))
|
||||
|
||||
ax1.set_ylabel('收盘价 (元)', fontsize=14, fontweight='bold')
|
||||
ax1.legend(loc='upper left', fontsize=12)
|
||||
ax1.grid(True, alpha=0.3)
|
||||
ax1.set_title(f'{stock_name}({stock_code}) 股票价格预测 - 未来{pred_len}个交易日',
|
||||
fontsize=16, fontweight='bold', pad=20)
|
||||
|
||||
# 设置x轴日期格式
|
||||
ax1.xaxis.set_major_formatter(plt.matplotlib.dates.DateFormatter('%Y-%m-%d'))
|
||||
plt.setp(ax1.xaxis.get_majorticklabels(), rotation=45)
|
||||
|
||||
# 设置y轴格式
|
||||
ax1.yaxis.set_major_formatter(plt.FuncFormatter(lambda x, p: f'{x:.2f}'))
|
||||
|
||||
# 2. 成交量图表 - 优化显示
|
||||
# 只显示预测期的成交量
|
||||
pred_volumes = volume_df['预测数据'].dropna()
|
||||
if len(pred_volumes) > 0:
|
||||
ax2.bar(pred_volumes.index, pred_volumes.values,
|
||||
alpha=0.7, color='#ff7f0e', label='预测成交量', width=0.8)
|
||||
|
||||
ax2.set_ylabel('成交量 (手)', fontsize=14, fontweight='bold')
|
||||
ax2.legend(loc='upper left', fontsize=12)
|
||||
ax2.grid(True, alpha=0.3)
|
||||
|
||||
# 设置x轴标签
|
||||
if len(pred_volumes) > 0:
|
||||
ax2.xaxis.set_major_formatter(plt.matplotlib.dates.DateFormatter('%m-%d'))
|
||||
plt.setp(ax2.xaxis.get_majorticklabels(), rotation=45)
|
||||
|
||||
# 3. 价格变动图表 - 优化显示
|
||||
if len(close_df['预测数据']) > 0:
|
||||
price_change = close_df['预测数据'] - close_df['历史数据'].iloc[-1]
|
||||
colors = ['green' if x >= 0 else 'red' for x in price_change]
|
||||
|
||||
# 每5个交易日显示一个标签,避免过于拥挤
|
||||
bars = ax3.bar(range(len(price_change)), price_change, alpha=0.8, color=colors)
|
||||
|
||||
# 在关键点添加数值标签
|
||||
for i, bar in enumerate(bars):
|
||||
height = bar.get_height()
|
||||
if i % 10 == 0 or i == len(bars) - 1 or abs(height) > price_change.std(): # 每10天或最后一天或显著波动
|
||||
ax3.text(bar.get_x() + bar.get_width() / 2., height,
|
||||
f'{height:+.2f}', ha='center', va='bottom' if height >= 0 else 'top',
|
||||
fontsize=8, fontweight='bold')
|
||||
|
||||
ax3.axhline(y=0, color='black', linestyle='-', alpha=0.5, linewidth=1)
|
||||
|
||||
ax3.set_ylabel('价格变动 (元)', fontsize=14, fontweight='bold')
|
||||
ax3.set_xlabel('交易日', fontsize=14, fontweight='bold')
|
||||
ax3.grid(True, alpha=0.3)
|
||||
|
||||
# 设置x轴标签
|
||||
if len(price_change) > 0:
|
||||
# 每10个交易日显示一个标签
|
||||
xticks_positions = list(range(0, len(price_change), max(1, len(price_change) // 10)))
|
||||
if len(price_change) - 1 not in xticks_positions:
|
||||
xticks_positions.append(len(price_change) - 1)
|
||||
ax3.set_xticks(xticks_positions)
|
||||
ax3.set_xticklabels([f'D{i + 1}' for i in xticks_positions])
|
||||
|
||||
# 添加详细的统计信息框
|
||||
if len(close_df['预测数据']) > 0 and not np.isnan(close_df['历史数据'].iloc[-1]):
|
||||
pred_stats = {
|
||||
'股票代码': stock_code,
|
||||
'股票名称': stock_name,
|
||||
'当前价格': f"{close_df['历史数据'].iloc[-1]:.2f} 元",
|
||||
'预测结束价格': f"{close_df['预测数据'].iloc[-1]:.2f} 元",
|
||||
'预测涨跌幅': f"{(close_df['预测数据'].iloc[-1] / close_df['历史数据'].iloc[-1] - 1) * 100:+.2f}%",
|
||||
'预测期间最高价': f"{close_df['预测数据'].max():.2f} 元",
|
||||
'预测期间最低价': f"{close_df['预测数据'].min():.2f} 元",
|
||||
'预测波动率': f"{close_df['预测数据'].std():.2f} 元",
|
||||
'预测起始日期': f"{close_df['预测数据'].index[0].strftime('%Y-%m-%d')}",
|
||||
'预测结束日期': f"{close_df['预测数据'].index[-1].strftime('%Y-%m-%d')}",
|
||||
'预测交易日数': f"{len(close_df['预测数据'])} 天"
|
||||
}
|
||||
|
||||
stats_text = "\n".join([f"{k}: {v}" for k, v in pred_stats.items()])
|
||||
fig.text(0.02, 0.02, stats_text, fontsize=10,
|
||||
bbox=dict(boxstyle="round,pad=0.5", facecolor="lightblue", alpha=0.8),
|
||||
verticalalignment='bottom')
|
||||
|
||||
plt.tight_layout()
|
||||
|
||||
# 保存高分辨率图片到指定目录
|
||||
chart_filename = os.path.join(output_dir, f'{stock_code}_prediction_chart.png')
|
||||
plt.savefig(chart_filename, dpi=300, bbox_inches='tight', facecolor='white')
|
||||
print(f"📊 预测图表已保存: {chart_filename}")
|
||||
|
||||
plt.show()
|
||||
|
||||
return close_df, volume_df
|
||||
|
||||
|
||||
def generate_prediction_report(close_df, volume_df, pred_df, future_dates, stock_code="002354", stock_name="股票",
|
||||
output_dir="."):
|
||||
"""
|
||||
生成预测报告
|
||||
"""
|
||||
# 确保输出目录存在
|
||||
ensure_output_directory(output_dir)
|
||||
|
||||
print(f"\n{'=' * 70}")
|
||||
print(f"📊 {stock_name}({stock_code}) 股票预测报告")
|
||||
print(f"{'=' * 70}")
|
||||
|
||||
if len(close_df['预测数据']) == 0 or np.isnan(close_df['历史数据'].iloc[-1]):
|
||||
print("❌ 没有有效的预测数据可生成报告")
|
||||
return
|
||||
|
||||
# 确保所有数组长度一致
|
||||
min_len = min(len(close_df['预测数据']), len(volume_df['预测数据']), len(future_dates))
|
||||
|
||||
# 基本统计
|
||||
historical_close = close_df['历史数据'].iloc[-1]
|
||||
predicted_close = close_df['预测数据'].iloc[-1]
|
||||
price_change_pct = (predicted_close / historical_close - 1) * 100
|
||||
|
||||
print(f"🔮 预测概览:")
|
||||
print(f" 当前价格: {historical_close:.2f} 元")
|
||||
print(f" 预测结束价格: {predicted_close:.2f} 元")
|
||||
print(f" 预测涨跌幅: {price_change_pct:+.2f}%")
|
||||
print(f" 预测期间: {min_len} 个交易日")
|
||||
print(
|
||||
f" 预测时间范围: {future_dates[0].strftime('%Y-%m-%d')} 到 {future_dates[min_len - 1].strftime('%Y-%m-%d')}")
|
||||
|
||||
print(f"\n📈 价格预测统计:")
|
||||
print(f" 预测最高价: {close_df['预测数据'].max():.2f} 元")
|
||||
print(f" 预测最低价: {close_df['预测数据'].min():.2f} 元")
|
||||
print(f" 预测平均价: {close_df['预测数据'].mean():.2f} 元")
|
||||
print(f" 价格波动率: {close_df['预测数据'].std():.2f} 元")
|
||||
|
||||
print(f"\n📊 成交量预测统计:")
|
||||
print(f" 预测平均成交量: {volume_df['预测数据'].mean():,.0f} 手")
|
||||
print(f" 预测最大成交量: {volume_df['预测数据'].max():,.0f} 手")
|
||||
print(f" 预测最小成交量: {volume_df['预测数据'].min():,.0f} 手")
|
||||
|
||||
# 保存详细预测数据到指定目录 - 确保所有数组长度一致
|
||||
prediction_details = pd.DataFrame({
|
||||
'日期': future_dates[:min_len],
|
||||
'预测收盘价': close_df['预测数据'].values[:min_len],
|
||||
'预测成交量': volume_df['预测数据'].values[:min_len],
|
||||
'价格变动(元)': (close_df['预测数据'].values[:min_len] - historical_close),
|
||||
'价格变动(%)': ((close_df['预测数据'].values[:min_len] / historical_close - 1) * 100)
|
||||
})
|
||||
|
||||
prediction_file = os.path.join(output_dir, f'{stock_code}_detailed_predictions.csv')
|
||||
prediction_details.to_csv(prediction_file, index=False, encoding='utf-8-sig')
|
||||
print(f"\n💾 详细预测数据已保存: {prediction_file}")
|
||||
|
||||
|
||||
def main(stock_code="002354", stock_name="天娱数科", data_dir="./data", pred_days=100, output_dir="./output"):
|
||||
"""
|
||||
主函数:执行股票价格预测
|
||||
|
||||
参数:
|
||||
stock_code: 股票代码
|
||||
stock_name: 股票名称
|
||||
data_dir: 数据文件目录
|
||||
pred_days: 预测天数(自然日)
|
||||
output_dir: 输出文件目录
|
||||
"""
|
||||
# 构建数据文件路径
|
||||
csv_file_path = os.path.join(data_dir, f"{stock_code}_stock_data.csv")
|
||||
|
||||
print(f"🎯 开始 {stock_name}({stock_code}) 股票价格预测")
|
||||
print("=" * 70)
|
||||
print(f"数据文件: {csv_file_path}")
|
||||
print(f"预测天数: {pred_days} 天(自然日)")
|
||||
print(f"输出目录: {output_dir}")
|
||||
|
||||
# 检查数据文件是否存在
|
||||
if not os.path.exists(csv_file_path):
|
||||
print(f"❌ 数据文件不存在: {csv_file_path}")
|
||||
print("请先运行数据获取脚本生成股票数据文件")
|
||||
return
|
||||
|
||||
# 确保输出目录存在
|
||||
ensure_output_directory(output_dir)
|
||||
|
||||
try:
|
||||
# 1. 加载模型和分词器
|
||||
print("\n步骤1: 加载Kronos模型和分词器...")
|
||||
tokenizer = KronosTokenizer.from_pretrained("NeoQuasar/Kronos-Tokenizer-base")
|
||||
model = Kronos.from_pretrained("NeoQuasar/Kronos-base")
|
||||
print("✅ 模型加载完成")
|
||||
|
||||
# 2. 实例化预测器
|
||||
print("步骤2: 初始化预测器...")
|
||||
predictor = KronosPredictor(model, tokenizer, device="cuda:0", max_context=512)
|
||||
print("✅ 预测器初始化完成")
|
||||
|
||||
# 3. 准备数据
|
||||
print("步骤3: 准备股票数据...")
|
||||
df = prepare_stock_data(csv_file_path, stock_code)
|
||||
|
||||
# 4. 计算预测参数
|
||||
print("步骤4: 计算预测参数...")
|
||||
lookback, pred_len = calculate_prediction_parameters(df, target_days=pred_days)
|
||||
|
||||
if pred_len <= 0:
|
||||
print("❌ 数据量不足,无法进行预测")
|
||||
return
|
||||
|
||||
print(f"✅ 最终参数 - 回看期: {lookback}, 预测期: {pred_len}")
|
||||
|
||||
# 5. 准备输入数据
|
||||
print("步骤5: 准备输入数据...")
|
||||
# 使用最新的数据作为输入
|
||||
x_df = df.loc[-lookback:, ['open', 'high', 'low', 'close', 'volume', 'amount']].reset_index(drop=True)
|
||||
x_timestamp = df.loc[-lookback:, 'timestamps'].reset_index(drop=True)
|
||||
|
||||
# 生成未来日期(考虑节假日)
|
||||
last_historical_date = df['timestamps'].iloc[-1]
|
||||
future_dates = generate_future_dates_with_holidays(last_historical_date, pred_len)
|
||||
|
||||
print(f"输入数据形状: {x_df.shape}")
|
||||
print(f"历史数据时间范围: {x_timestamp.iloc[0]} 到 {x_timestamp.iloc[-1]}")
|
||||
print(f"预测时间范围: {future_dates[0]} 到 {future_dates[-1]}")
|
||||
|
||||
# 6. 执行预测
|
||||
print("步骤6: 执行价格预测...")
|
||||
pred_df = predictor.predict(
|
||||
df=x_df,
|
||||
x_timestamp=x_timestamp,
|
||||
y_timestamp=pd.Series(future_dates), # 使用未来日期作为预测时间戳
|
||||
pred_len=pred_len,
|
||||
T=1.0,
|
||||
top_p=0.9,
|
||||
sample_count=1,
|
||||
verbose=True
|
||||
)
|
||||
|
||||
print("✅ 预测完成")
|
||||
|
||||
# 7. 显示预测结果
|
||||
print("\n步骤7: 显示预测结果...")
|
||||
print("预测数据前5行:")
|
||||
# 确保预测数据长度与未来日期一致
|
||||
min_len = min(len(pred_df), len(future_dates))
|
||||
pred_df = pred_df.iloc[:min_len]
|
||||
pred_df.index = future_dates[:min_len]
|
||||
print(pred_df.head())
|
||||
|
||||
# 8. 可视化结果
|
||||
print("步骤8: 生成可视化图表...")
|
||||
# 使用最后一部分历史数据和预测数据
|
||||
kline_df = df.loc[-lookback:].reset_index(drop=True)
|
||||
close_df, volume_df = plot_prediction_with_details(kline_df, pred_df, future_dates, stock_code, stock_name,
|
||||
pred_len, output_dir)
|
||||
|
||||
# 9. 生成预测报告
|
||||
print("步骤9: 生成预测报告...")
|
||||
generate_prediction_report(close_df, volume_df, pred_df, future_dates, stock_code, stock_name, output_dir)
|
||||
|
||||
print(f"\n🎉 {stock_name}({stock_code}) 股票预测完成!")
|
||||
print("生成的文件:")
|
||||
print(f" 📊 {os.path.join(output_dir, stock_code + '_prediction_chart.png')} - 预测图表")
|
||||
print(f" 📋 {os.path.join(output_dir, stock_code + '_detailed_predictions.csv')} - 详细预测数据")
|
||||
|
||||
# 显示预测总结
|
||||
if len(close_df['预测数据']) > 0 and not np.isnan(close_df['历史数据'].iloc[-1]):
|
||||
print(f"\n📈 预测总结:")
|
||||
historical_price = close_df['历史数据'].iloc[-1]
|
||||
predicted_price = close_df['预测数据'].iloc[-1]
|
||||
change_pct = (predicted_price / historical_price - 1) * 100
|
||||
|
||||
print(f" 当前价格: {historical_price:.2f} 元")
|
||||
print(f" 预测价格: {predicted_price:.2f} 元")
|
||||
print(f" 预期涨跌: {change_pct:+.2f}%")
|
||||
print(
|
||||
f" 预测时间: {future_dates[0].strftime('%Y-%m-%d')} 到 {future_dates[min_len - 1].strftime('%Y-%m-%d')}")
|
||||
|
||||
if change_pct > 10:
|
||||
print(f" 🚀 模型预测未来{pred_len}个交易日大幅看涨 (+{change_pct:.1f}%)")
|
||||
elif change_pct > 5:
|
||||
print(f" 📈 模型预测未来{pred_len}个交易日看涨 (+{change_pct:.1f}%)")
|
||||
elif change_pct > 0:
|
||||
print(f" ↗️ 模型预测未来{pred_len}个交易日微涨 (+{change_pct:.1f}%)")
|
||||
elif change_pct > -5:
|
||||
print(f" ↘️ 模型预测未来{pred_len}个交易日微跌 ({change_pct:.1f}%)")
|
||||
elif change_pct > -10:
|
||||
print(f" 📉 模型预测未来{pred_len}个交易日看跌 ({change_pct:.1f}%)")
|
||||
else:
|
||||
print(f" 🔻 模型预测未来{pred_len}个交易日大幅看跌 ({change_pct:.1f}%)")
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ 预测过程中出现错误: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
|
||||
# 使用方法说明
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
股票预测工具 - 支持多股票预测
|
||||
|
||||
使用方法:
|
||||
修改下面的 STOCK_CONFIG 来预测不同的股票
|
||||
"""
|
||||
|
||||
# ==================== 在这里修改股票配置 ====================
|
||||
STOCK_CONFIG = {
|
||||
"stock_code": "300418", # 股票代码
|
||||
"stock_name": "昆仑万维", # 股票名称
|
||||
"data_dir": "./data", # 数据文件目录
|
||||
"pred_days": 100, # 预测100个自然日
|
||||
"output_dir": r"D:\lianghuajiaoyi\Kronos\examples\yuce" # 输出文件目录
|
||||
}
|
||||
|
||||
# 其他股票配置示例:
|
||||
# STOCK_CONFIG = {"stock_code": "000001", "stock_name": "平安银行", "data_dir": "./data", "pred_days": 100, "output_dir": r"D:\lianghuajiaoyi\Kronos\examples\yuce"}
|
||||
# STOCK_CONFIG = {"stock_code": "600036", "stock_name": "招商银行", "data_dir": "./data", "pred_days": 100, "output_dir": r"D:\lianghuajiaoyi\Kronos\examples\yuce"}
|
||||
# STOCK_CONFIG = {"stock_code": "300750", "stock_name": "宁德时代", "data_dir": "./data", "pred_days": 100, "output_dir": r"D:\lianghuajiaoyi\Kronos\examples\yuce"}
|
||||
# =========================================================
|
||||
|
||||
print("🤖 智能股票预测工具")
|
||||
print("=" * 70)
|
||||
print(f"当前预测股票: {STOCK_CONFIG['stock_name']}({STOCK_CONFIG['stock_code']})")
|
||||
print(f"数据目录: {STOCK_CONFIG['data_dir']}")
|
||||
print(f"预测天数: {STOCK_CONFIG['pred_days']} 天(自然日)")
|
||||
print(f"输出目录: {STOCK_CONFIG['output_dir']}")
|
||||
print()
|
||||
|
||||
# 运行主程序
|
||||
main(**STOCK_CONFIG)
|
||||
|
||||
print(f"\n💡 提示:要预测其他股票,请修改代码中的 STOCK_CONFIG 变量")
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,455 @@
|
||||
# run_backtest.py
|
||||
import pandas as pd
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('ignore')
|
||||
|
||||
# 设置中文字体
|
||||
plt.rcParams['font.sans-serif'] = ['SimHei']
|
||||
plt.rcParams['axes.unicode_minus'] = False
|
||||
|
||||
|
||||
class KronosBacktester:
|
||||
"""
|
||||
Kronos模型回测类
|
||||
"""
|
||||
|
||||
def __init__(self, data_dir, model_dir, initial_capital=100000):
|
||||
"""
|
||||
初始化回测器
|
||||
|
||||
参数:
|
||||
data_dir: 数据目录
|
||||
model_dir: 模型预测结果目录
|
||||
initial_capital: 初始资金
|
||||
"""
|
||||
self.data_dir = data_dir
|
||||
self.model_dir = model_dir
|
||||
self.initial_capital = initial_capital
|
||||
self.results = {}
|
||||
|
||||
def load_historical_data(self, stock_code):
|
||||
"""
|
||||
加载历史数据
|
||||
"""
|
||||
csv_file = os.path.join(self.data_dir, f"{stock_code}_stock_data.csv")
|
||||
if not os.path.exists(csv_file):
|
||||
raise FileNotFoundError(f"数据文件不存在: {csv_file}")
|
||||
|
||||
df = pd.read_csv(csv_file, encoding='utf-8-sig')
|
||||
|
||||
# 检查列名并标准化
|
||||
column_mapping = {
|
||||
'日期': 'date',
|
||||
'开盘价': 'open',
|
||||
'最高价': 'high',
|
||||
'最低价': 'low',
|
||||
'收盘价': 'close',
|
||||
'成交量': 'volume',
|
||||
'成交额': 'amount'
|
||||
}
|
||||
|
||||
# 重命名列
|
||||
for old_col, new_col in column_mapping.items():
|
||||
if old_col in df.columns:
|
||||
df = df.rename(columns={old_col: new_col})
|
||||
|
||||
df['date'] = pd.to_datetime(df['date'])
|
||||
df.set_index('date', inplace=True)
|
||||
df = df.sort_index()
|
||||
|
||||
print(f"✅ 加载历史数据: {len(df)} 条记录")
|
||||
print(f"时间范围: {df.index.min()} 到 {df.index.max()}")
|
||||
|
||||
return df
|
||||
|
||||
def load_predictions(self, stock_code):
|
||||
"""
|
||||
加载模型预测结果
|
||||
"""
|
||||
# 尝试不同的预测文件命名
|
||||
pred_files = [
|
||||
os.path.join(self.model_dir, f"{stock_code}_kronos_predictions.csv"),
|
||||
os.path.join(self.model_dir, f"{stock_code}_detailed_predictions.csv"),
|
||||
os.path.join(self.model_dir, f"{stock_code}_predictions.csv")
|
||||
]
|
||||
|
||||
pred_df = None
|
||||
for pred_file in pred_files:
|
||||
if os.path.exists(pred_file):
|
||||
pred_df = pd.read_csv(pred_file, encoding='utf-8-sig')
|
||||
print(f"✅ 找到预测文件: {pred_file}")
|
||||
break
|
||||
|
||||
if pred_df is None:
|
||||
raise FileNotFoundError(f"未找到预测文件,请检查目录: {self.model_dir}")
|
||||
|
||||
# 标准化列名
|
||||
column_mapping = {
|
||||
'日期': 'date',
|
||||
'预测收盘价': 'predicted_close',
|
||||
'收盘价': 'predicted_close',
|
||||
'预测成交量': 'predicted_volume',
|
||||
'成交量': 'predicted_volume'
|
||||
}
|
||||
|
||||
for old_col, new_col in column_mapping.items():
|
||||
if old_col in pred_df.columns:
|
||||
pred_df = pred_df.rename(columns={old_col: new_col})
|
||||
|
||||
pred_df['date'] = pd.to_datetime(pred_df['date'])
|
||||
pred_df.set_index('date', inplace=True)
|
||||
pred_df = pred_df.sort_index()
|
||||
|
||||
print(f"✅ 加载预测数据: {len(pred_df)} 条记录")
|
||||
print(f"预测时间范围: {pred_df.index.min()} 到 {pred_df.index.max()}")
|
||||
|
||||
return pred_df
|
||||
|
||||
def align_data(self, hist_df, pred_df):
|
||||
"""
|
||||
对齐历史数据和预测数据的时间范围
|
||||
"""
|
||||
# 找到历史数据的最后日期
|
||||
last_hist_date = hist_df.index.max()
|
||||
|
||||
# 筛选预测数据,从历史数据结束后开始
|
||||
pred_df_aligned = pred_df[pred_df.index > last_hist_date]
|
||||
|
||||
if len(pred_df_aligned) == 0:
|
||||
# 如果没有未来的预测数据,使用所有预测数据
|
||||
pred_df_aligned = pred_df.copy()
|
||||
print("⚠️ 警告:预测数据没有未来的日期,使用所有预测数据")
|
||||
|
||||
print(f"✅ 数据对齐: 历史数据结束于 {last_hist_date}, 预测数据从 {pred_df_aligned.index.min()} 开始")
|
||||
|
||||
return pred_df_aligned
|
||||
|
||||
def calculate_trading_signals(self, hist_df, pred_df, threshold=0.02):
|
||||
"""
|
||||
计算交易信号
|
||||
"""
|
||||
# 对齐数据
|
||||
pred_df = self.align_data(hist_df, pred_df)
|
||||
|
||||
# 合并历史数据和预测数据
|
||||
combined = pd.concat([
|
||||
hist_df[['close']].rename(columns={'close': 'actual'}),
|
||||
pred_df[['predicted_close']].rename(columns={'predicted_close': 'predicted'})
|
||||
], axis=1)
|
||||
|
||||
# 计算预测收益率
|
||||
combined['pred_return'] = combined['predicted'].pct_change()
|
||||
|
||||
# 生成交易信号
|
||||
combined['signal'] = 0
|
||||
combined['signal'] = np.where(combined['pred_return'] > threshold, 1, # 买入信号
|
||||
np.where(combined['pred_return'] < -threshold, -1, 0)) # 卖出信号
|
||||
|
||||
# 过滤信号:避免频繁交易
|
||||
combined['position'] = combined['signal'].replace(to_replace=0, method='ffill').fillna(0)
|
||||
|
||||
return combined
|
||||
|
||||
def run_backtest(self, combined_df):
|
||||
"""
|
||||
运行回测
|
||||
"""
|
||||
# 初始化资金和持仓
|
||||
capital = self.initial_capital
|
||||
position = 0
|
||||
trades = []
|
||||
|
||||
# 回测记录
|
||||
backtest_results = pd.DataFrame(index=combined_df.index)
|
||||
backtest_results['capital'] = capital
|
||||
backtest_results['position'] = 0
|
||||
backtest_results['returns'] = 0.0
|
||||
backtest_results['price'] = combined_df['actual'].combine_first(combined_df['predicted'])
|
||||
|
||||
for i, (date, row) in enumerate(combined_df.iterrows()):
|
||||
current_price = row['actual'] if not pd.isna(row['actual']) else row['predicted']
|
||||
signal = row['position']
|
||||
|
||||
# 跳过无效价格
|
||||
if pd.isna(current_price):
|
||||
continue
|
||||
|
||||
# 执行交易
|
||||
if i > 0: # 从第二天开始
|
||||
prev_position = backtest_results['position'].iloc[i - 1] if i > 0 else 0
|
||||
|
||||
# 平仓信号
|
||||
if prev_position != 0 and signal == 0:
|
||||
# 平仓
|
||||
capital = position * current_price
|
||||
position = 0
|
||||
trades.append({
|
||||
'date': date,
|
||||
'action': 'SELL',
|
||||
'price': current_price,
|
||||
'shares': prev_position,
|
||||
'capital': capital
|
||||
})
|
||||
|
||||
# 开仓信号
|
||||
elif prev_position == 0 and signal != 0:
|
||||
# 计算可买股数(假设全仓交易)
|
||||
shares = int(capital / current_price)
|
||||
if shares > 0:
|
||||
position = shares * signal
|
||||
capital -= shares * current_price
|
||||
trades.append({
|
||||
'date': date,
|
||||
'action': 'BUY',
|
||||
'price': current_price,
|
||||
'shares': shares * signal,
|
||||
'capital': capital
|
||||
})
|
||||
|
||||
# 更新持仓市值
|
||||
portfolio_value = capital + position * current_price
|
||||
|
||||
# 记录结果
|
||||
backtest_results.loc[date, 'capital'] = portfolio_value
|
||||
backtest_results.loc[date, 'position'] = position
|
||||
backtest_results.loc[date, 'price'] = current_price
|
||||
|
||||
# 计算日收益率
|
||||
if i > 0:
|
||||
prev_value = backtest_results['capital'].iloc[i - 1]
|
||||
if prev_value > 0:
|
||||
backtest_results.loc[date, 'returns'] = (portfolio_value - prev_value) / prev_value
|
||||
|
||||
return backtest_results, trades
|
||||
|
||||
def calculate_metrics(self, backtest_results, trades):
|
||||
"""
|
||||
计算回测指标
|
||||
"""
|
||||
returns = backtest_results['returns'].replace([np.inf, -np.inf], np.nan).dropna()
|
||||
|
||||
if len(returns) == 0:
|
||||
return {
|
||||
'总收益率': 0,
|
||||
'年化收益率': 0,
|
||||
'波动率': 0,
|
||||
'夏普比率': 0,
|
||||
'最大回撤': 0,
|
||||
'胜率': 0,
|
||||
'平均交易收益': 0,
|
||||
'交易次数': 0,
|
||||
'最终资金': self.initial_capital
|
||||
}
|
||||
|
||||
total_return = (backtest_results['capital'].iloc[-1] - self.initial_capital) / self.initial_capital
|
||||
annual_return = (1 + total_return) ** (252 / len(returns)) - 1
|
||||
|
||||
# 波动率
|
||||
volatility = returns.std() * np.sqrt(252)
|
||||
|
||||
# 夏普比率(假设无风险利率为3%)
|
||||
risk_free_rate = 0.03
|
||||
sharpe_ratio = (annual_return - risk_free_rate) / volatility if volatility > 0 else 0
|
||||
|
||||
# 最大回撤
|
||||
cumulative_returns = (1 + returns).cumprod()
|
||||
peak = cumulative_returns.expanding().max()
|
||||
drawdown = (cumulative_returns - peak) / peak
|
||||
max_drawdown = drawdown.min()
|
||||
|
||||
# 交易统计
|
||||
trade_returns = []
|
||||
buy_trades = [t for t in trades if t['action'] == 'BUY']
|
||||
sell_trades = [t for t in trades if t['action'] == 'SELL']
|
||||
|
||||
for i in range(min(len(buy_trades), len(sell_trades))):
|
||||
buy = buy_trades[i]
|
||||
sell = sell_trades[i]
|
||||
trade_return = (sell['price'] - buy['price']) / buy['price']
|
||||
trade_returns.append(trade_return)
|
||||
|
||||
win_rate = len([r for r in trade_returns if r > 0]) / len(trade_returns) if trade_returns else 0
|
||||
avg_trade_return = np.mean(trade_returns) if trade_returns else 0
|
||||
|
||||
metrics = {
|
||||
'总收益率': total_return,
|
||||
'年化收益率': annual_return,
|
||||
'波动率': volatility,
|
||||
'夏普比率': sharpe_ratio,
|
||||
'最大回撤': max_drawdown,
|
||||
'胜率': win_rate,
|
||||
'平均交易收益': avg_trade_return,
|
||||
'交易次数': len(trades),
|
||||
'最终资金': backtest_results['capital'].iloc[-1]
|
||||
}
|
||||
|
||||
return metrics
|
||||
|
||||
def plot_backtest_results(self, backtest_results, metrics, stock_code, output_dir):
|
||||
"""
|
||||
绘制回测结果图表
|
||||
"""
|
||||
fig, (ax1, ax2, ax3) = plt.subplots(3, 1, figsize=(15, 12))
|
||||
|
||||
# 1. 资金曲线
|
||||
ax1.plot(backtest_results.index, backtest_results['capital'],
|
||||
linewidth=2, label='策略资金曲线', color='#1f77b4')
|
||||
ax1.axhline(y=self.initial_capital, color='red', linestyle='--',
|
||||
label=f'初始资金 ({self.initial_capital:,.0f}元)')
|
||||
ax1.set_ylabel('资金 (元)', fontsize=12)
|
||||
ax1.legend()
|
||||
ax1.grid(True, alpha=0.3)
|
||||
ax1.set_title(f'{stock_code} Kronos模型回测结果', fontsize=14, fontweight='bold')
|
||||
|
||||
# 2. 收益率曲线
|
||||
cumulative_returns = (1 + backtest_results['returns'].fillna(0)).cumprod()
|
||||
ax2.plot(backtest_results.index, cumulative_returns,
|
||||
linewidth=2, label='策略累计收益', color='#2ca02c')
|
||||
|
||||
# 基准收益(买入持有)
|
||||
price_returns = backtest_results['price'].pct_change().fillna(0)
|
||||
benchmark_returns = (1 + price_returns).cumprod()
|
||||
ax2.plot(backtest_results.index, benchmark_returns,
|
||||
linewidth=2, label='基准收益(买入持有)', color='#ff7f0e', alpha=0.7)
|
||||
|
||||
ax2.set_ylabel('累计收益', fontsize=12)
|
||||
ax2.legend()
|
||||
ax2.grid(True, alpha=0.3)
|
||||
|
||||
# 3. 回撤曲线
|
||||
peak = cumulative_returns.expanding().max()
|
||||
drawdown = (cumulative_returns - peak) / peak
|
||||
ax3.fill_between(backtest_results.index, drawdown, 0,
|
||||
alpha=0.3, color='red', label='回撤')
|
||||
ax3.set_ylabel('回撤', fontsize=12)
|
||||
ax3.set_xlabel('日期', fontsize=12)
|
||||
ax3.legend()
|
||||
ax3.grid(True, alpha=0.3)
|
||||
|
||||
# 添加指标文本
|
||||
metrics_text = (
|
||||
f"总收益率: {metrics['总收益率']:.2%}\n"
|
||||
f"年化收益率: {metrics['年化收益率']:.2%}\n"
|
||||
f"夏普比率: {metrics['夏普比率']:.2f}\n"
|
||||
f"最大回撤: {metrics['最大回撤']:.2%}\n"
|
||||
f"胜率: {metrics['胜率']:.2%}\n"
|
||||
f"交易次数: {metrics['交易次数']}\n"
|
||||
f"最终资金: {metrics['最终资金']:,.0f}元"
|
||||
)
|
||||
|
||||
ax1.text(0.02, 0.98, metrics_text, transform=ax1.transAxes, fontsize=10,
|
||||
verticalalignment='top', bbox=dict(boxstyle="round,pad=0.3",
|
||||
facecolor="lightyellow", alpha=0.8))
|
||||
|
||||
plt.tight_layout()
|
||||
|
||||
# 保存图表
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
chart_file = os.path.join(output_dir, f'{stock_code}_backtest_results.png')
|
||||
plt.savefig(chart_file, dpi=300, bbox_inches='tight')
|
||||
print(f"📊 回测图表已保存: {chart_file}")
|
||||
|
||||
plt.show()
|
||||
|
||||
def run_complete_backtest(self, stock_code, output_dir, threshold=0.02):
|
||||
"""
|
||||
运行完整的回测流程
|
||||
"""
|
||||
print(f"🎯 开始 {stock_code} 回测分析")
|
||||
print("=" * 50)
|
||||
|
||||
try:
|
||||
# 1. 加载数据
|
||||
print("步骤1: 加载历史数据和预测数据...")
|
||||
hist_df = self.load_historical_data(stock_code)
|
||||
pred_df = self.load_predictions(stock_code)
|
||||
|
||||
# 2. 计算交易信号
|
||||
print("步骤2: 计算交易信号...")
|
||||
combined_df = self.calculate_trading_signals(hist_df, pred_df, threshold)
|
||||
|
||||
# 3. 运行回测
|
||||
print("步骤3: 运行回测...")
|
||||
backtest_results, trades = self.run_backtest(combined_df)
|
||||
|
||||
# 4. 计算指标
|
||||
print("步骤4: 计算回测指标...")
|
||||
metrics = self.calculate_metrics(backtest_results, trades)
|
||||
|
||||
# 5. 绘制结果
|
||||
print("步骤5: 生成回测图表...")
|
||||
self.plot_backtest_results(backtest_results, metrics, stock_code, output_dir)
|
||||
|
||||
# 6. 打印详细报告
|
||||
print("\n" + "=" * 70)
|
||||
print(f"📊 {stock_code} 回测报告")
|
||||
print("=" * 70)
|
||||
for key, value in metrics.items():
|
||||
if isinstance(value, float):
|
||||
if '率' in key or '收益' in key or '回撤' in key:
|
||||
print(f" {key}: {value:.2%}")
|
||||
else:
|
||||
print(f" {key}: {value:.2f}")
|
||||
else:
|
||||
print(f" {key}: {value}")
|
||||
|
||||
print(f"\n交易记录 (共{len(trades)}次交易):")
|
||||
for i, trade in enumerate(trades[-10:], 1): # 显示最后10次交易
|
||||
print(f" 交易{i}: {trade['date'].strftime('%Y-%m-%d')} "
|
||||
f"{trade['action']} {abs(trade['shares'])}股 @ {trade['price']:.2f}元")
|
||||
|
||||
return metrics, backtest_results, trades
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ 回测过程中出现错误: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return None, None, None
|
||||
|
||||
|
||||
def main():
|
||||
"""
|
||||
主函数:运行Kronos模型回测
|
||||
"""
|
||||
# 配置参数
|
||||
BACKTEST_CONFIG = {
|
||||
"stock_code": "000831", # 要回测的股票代码
|
||||
"data_dir": r"D:\lianghuajiaoyi\Kronos\examples\data", # 历史数据目录
|
||||
"model_dir": r"D:\lianghuajiaoyi\Kronos\examples\yuce", # 模型预测结果目录
|
||||
"output_dir": r"D:\lianghuajiaoyi\Kronos\examples\backtest", # 回测结果输出目录
|
||||
"initial_capital": 100000, # 初始资金
|
||||
"threshold": 0.02 # 交易阈值(2%)
|
||||
}
|
||||
|
||||
print("🤖 Kronos模型回测系统")
|
||||
print("=" * 50)
|
||||
print(f"回测股票: {BACKTEST_CONFIG['stock_code']}")
|
||||
print(f"初始资金: {BACKTEST_CONFIG['initial_capital']:,.0f}元")
|
||||
print(f"交易阈值: {BACKTEST_CONFIG['threshold']:.1%}")
|
||||
print()
|
||||
|
||||
# 创建回测器并运行
|
||||
backtester = KronosBacktester(
|
||||
data_dir=BACKTEST_CONFIG["data_dir"],
|
||||
model_dir=BACKTEST_CONFIG["model_dir"],
|
||||
initial_capital=BACKTEST_CONFIG["initial_capital"]
|
||||
)
|
||||
|
||||
metrics, results, trades = backtester.run_complete_backtest(
|
||||
stock_code=BACKTEST_CONFIG["stock_code"],
|
||||
output_dir=BACKTEST_CONFIG["output_dir"],
|
||||
threshold=BACKTEST_CONFIG["threshold"]
|
||||
)
|
||||
|
||||
if metrics:
|
||||
print(f"\n✅ {BACKTEST_CONFIG['stock_code']} 回测完成!")
|
||||
print(f"📁 结果保存在: {BACKTEST_CONFIG['output_dir']}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,60 @@
|
||||
{
|
||||
"timestamp": "2025-10-10 16:09:38",
|
||||
"stock_code": "000021",
|
||||
"market_analysis": {
|
||||
"overall_is_main_uptrend": false,
|
||||
"overall_trend_strength": 0.5,
|
||||
"market_status": "未知",
|
||||
"detailed_analysis": {}
|
||||
},
|
||||
"sector_analysis": {
|
||||
"industry": "消费电子",
|
||||
"matched_sectors": [],
|
||||
"main_sector": {
|
||||
"sector": "传统行业",
|
||||
"momentum": 0.5,
|
||||
"description": "无热门概念"
|
||||
},
|
||||
"is_sector_hot": false,
|
||||
"resonance_score": 0.5,
|
||||
"sector_count": 0
|
||||
},
|
||||
"macro_analysis": {
|
||||
"us_rate_cycle": {
|
||||
"current_rate": 4.25,
|
||||
"trend": "降息周期",
|
||||
"recent_cut": "2025年9月降息25个基点",
|
||||
"expected_cuts_2025": 2,
|
||||
"expected_cuts_2026": 2,
|
||||
"impact_on_emerging_markets": "positive",
|
||||
"usd_index_support": 95.0,
|
||||
"analysis": "美联储开启宽松周期,利好全球流动性"
|
||||
},
|
||||
"domestic_policy": {
|
||||
"monetary_policy": "稳健偏松",
|
||||
"fiscal_policy": "积极财政",
|
||||
"market_liquidity": "合理充裕",
|
||||
"industrial_policy": "设备更新、以旧换新",
|
||||
"employment_policy": "稳就业政策加力",
|
||||
"analysis": "政策组合拳发力,经济稳中向好"
|
||||
},
|
||||
"industry_policy": {
|
||||
"robot_policy": "机器人产业政策支持",
|
||||
"chip_policy": "国产替代加速推进",
|
||||
"AI_policy": "人工智能发展规划",
|
||||
"low_altitude": "低空经济发展规划"
|
||||
},
|
||||
"global_liquidity_outlook": "改善",
|
||||
"overall_macro_score": 0.75
|
||||
},
|
||||
"fundamental_analysis": {
|
||||
"company_name": "未知",
|
||||
"business_areas": [],
|
||||
"recent_developments": [],
|
||||
"growth_drivers": [],
|
||||
"risk_factors": [],
|
||||
"investment_rating": "中性",
|
||||
"fundamental_score": 0.5
|
||||
},
|
||||
"adjustment_factor": 1.04545
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 857 KiB |
@@ -0,0 +1,60 @@
|
||||
{
|
||||
"timestamp": "2025-10-01 01:55:50",
|
||||
"stock_code": "002354",
|
||||
"market_analysis": {
|
||||
"overall_is_main_uptrend": false,
|
||||
"overall_trend_strength": 0.5,
|
||||
"market_status": "未知",
|
||||
"detailed_analysis": {}
|
||||
},
|
||||
"sector_analysis": {
|
||||
"industry": "互联网服务",
|
||||
"matched_sectors": [],
|
||||
"main_sector": {
|
||||
"sector": "传统行业",
|
||||
"momentum": 0.5,
|
||||
"description": "无热门概念"
|
||||
},
|
||||
"is_sector_hot": false,
|
||||
"resonance_score": 0.5,
|
||||
"sector_count": 0
|
||||
},
|
||||
"macro_analysis": {
|
||||
"us_rate_cycle": {
|
||||
"current_rate": 4.25,
|
||||
"trend": "降息周期",
|
||||
"recent_cut": "2025年9月降息25个基点",
|
||||
"expected_cuts_2025": 2,
|
||||
"expected_cuts_2026": 2,
|
||||
"impact_on_emerging_markets": "positive",
|
||||
"usd_index_support": 95.0,
|
||||
"analysis": "美联储开启宽松周期,利好全球流动性"
|
||||
},
|
||||
"domestic_policy": {
|
||||
"monetary_policy": "稳健偏松",
|
||||
"fiscal_policy": "积极财政",
|
||||
"market_liquidity": "合理充裕",
|
||||
"industrial_policy": "设备更新、以旧换新",
|
||||
"employment_policy": "稳就业政策加力",
|
||||
"analysis": "政策组合拳发力,经济稳中向好"
|
||||
},
|
||||
"industry_policy": {
|
||||
"robot_policy": "机器人产业政策支持",
|
||||
"chip_policy": "国产替代加速推进",
|
||||
"AI_policy": "人工智能发展规划",
|
||||
"low_altitude": "低空经济发展规划"
|
||||
},
|
||||
"global_liquidity_outlook": "改善",
|
||||
"overall_macro_score": 0.75
|
||||
},
|
||||
"fundamental_analysis": {
|
||||
"company_name": "未知",
|
||||
"business_areas": [],
|
||||
"recent_developments": [],
|
||||
"growth_drivers": [],
|
||||
"risk_factors": [],
|
||||
"investment_rating": "中性",
|
||||
"fundamental_score": 0.5
|
||||
},
|
||||
"adjustment_factor": 1.04545
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 940 KiB |
@@ -0,0 +1,70 @@
|
||||
{
|
||||
"timestamp": "2025-10-01 01:54:37",
|
||||
"stock_code": "300207",
|
||||
"market_analysis": {
|
||||
"overall_is_main_uptrend": false,
|
||||
"overall_trend_strength": 0.5,
|
||||
"market_status": "未知",
|
||||
"detailed_analysis": {}
|
||||
},
|
||||
"sector_analysis": {
|
||||
"industry": "电池",
|
||||
"matched_sectors": [
|
||||
{
|
||||
"sector": "新能源",
|
||||
"momentum": 0.6,
|
||||
"limit_up_stocks": 8,
|
||||
"is_active": true,
|
||||
"description": "光伏、储能"
|
||||
}
|
||||
],
|
||||
"main_sector": {
|
||||
"sector": "新能源",
|
||||
"momentum": 0.6,
|
||||
"limit_up_stocks": 8,
|
||||
"is_active": true,
|
||||
"description": "光伏、储能"
|
||||
},
|
||||
"is_sector_hot": true,
|
||||
"resonance_score": 0.6,
|
||||
"sector_count": 1
|
||||
},
|
||||
"macro_analysis": {
|
||||
"us_rate_cycle": {
|
||||
"current_rate": 4.25,
|
||||
"trend": "降息周期",
|
||||
"recent_cut": "2025年9月降息25个基点",
|
||||
"expected_cuts_2025": 2,
|
||||
"expected_cuts_2026": 2,
|
||||
"impact_on_emerging_markets": "positive",
|
||||
"usd_index_support": 95.0,
|
||||
"analysis": "美联储开启宽松周期,利好全球流动性"
|
||||
},
|
||||
"domestic_policy": {
|
||||
"monetary_policy": "稳健偏松",
|
||||
"fiscal_policy": "积极财政",
|
||||
"market_liquidity": "合理充裕",
|
||||
"industrial_policy": "设备更新、以旧换新",
|
||||
"employment_policy": "稳就业政策加力",
|
||||
"analysis": "政策组合拳发力,经济稳中向好"
|
||||
},
|
||||
"industry_policy": {
|
||||
"robot_policy": "机器人产业政策支持",
|
||||
"chip_policy": "国产替代加速推进",
|
||||
"AI_policy": "人工智能发展规划",
|
||||
"low_altitude": "低空经济发展规划"
|
||||
},
|
||||
"global_liquidity_outlook": "改善",
|
||||
"overall_macro_score": 0.75
|
||||
},
|
||||
"fundamental_analysis": {
|
||||
"company_name": "未知",
|
||||
"business_areas": [],
|
||||
"recent_developments": [],
|
||||
"growth_drivers": [],
|
||||
"risk_factors": [],
|
||||
"investment_rating": "中性",
|
||||
"fundamental_score": 0.5
|
||||
},
|
||||
"adjustment_factor": 1.0935407000000001
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 776 KiB |
@@ -0,0 +1,96 @@
|
||||
{
|
||||
"timestamp": "2025-10-01 01:53:45",
|
||||
"stock_code": "600580",
|
||||
"market_analysis": {
|
||||
"overall_is_main_uptrend": false,
|
||||
"overall_trend_strength": 0.5,
|
||||
"market_status": "未知",
|
||||
"detailed_analysis": {}
|
||||
},
|
||||
"sector_analysis": {
|
||||
"industry": "电机",
|
||||
"matched_sectors": [
|
||||
{
|
||||
"sector": "机器人",
|
||||
"momentum": 0.85,
|
||||
"limit_up_stocks": 18,
|
||||
"is_active": true,
|
||||
"description": "人形机器人、工业自动化"
|
||||
},
|
||||
{
|
||||
"sector": "低空经济",
|
||||
"momentum": 0.7,
|
||||
"limit_up_stocks": 10,
|
||||
"is_active": true,
|
||||
"description": "无人机、eVTOL"
|
||||
}
|
||||
],
|
||||
"main_sector": {
|
||||
"sector": "机器人",
|
||||
"momentum": 0.85,
|
||||
"limit_up_stocks": 18,
|
||||
"is_active": true,
|
||||
"description": "人形机器人、工业自动化"
|
||||
},
|
||||
"is_sector_hot": true,
|
||||
"resonance_score": 0.7749999999999999,
|
||||
"sector_count": 2
|
||||
},
|
||||
"macro_analysis": {
|
||||
"us_rate_cycle": {
|
||||
"current_rate": 4.25,
|
||||
"trend": "降息周期",
|
||||
"recent_cut": "2025年9月降息25个基点",
|
||||
"expected_cuts_2025": 2,
|
||||
"expected_cuts_2026": 2,
|
||||
"impact_on_emerging_markets": "positive",
|
||||
"usd_index_support": 95.0,
|
||||
"analysis": "美联储开启宽松周期,利好全球流动性"
|
||||
},
|
||||
"domestic_policy": {
|
||||
"monetary_policy": "稳健偏松",
|
||||
"fiscal_policy": "积极财政",
|
||||
"market_liquidity": "合理充裕",
|
||||
"industrial_policy": "设备更新、以旧换新",
|
||||
"employment_policy": "稳就业政策加力",
|
||||
"analysis": "政策组合拳发力,经济稳中向好"
|
||||
},
|
||||
"industry_policy": {
|
||||
"robot_policy": "机器人产业政策支持",
|
||||
"chip_policy": "国产替代加速推进",
|
||||
"AI_policy": "人工智能发展规划",
|
||||
"low_altitude": "低空经济发展规划"
|
||||
},
|
||||
"global_liquidity_outlook": "改善",
|
||||
"overall_macro_score": 0.75
|
||||
},
|
||||
"fundamental_analysis": {
|
||||
"company_name": "卧龙电驱",
|
||||
"business_areas": [
|
||||
"工业电机",
|
||||
"机器人关键部件",
|
||||
"航空电机",
|
||||
"新能源汽车驱动"
|
||||
],
|
||||
"recent_developments": [
|
||||
"与智元机器人实现双向持股,推进具身智能机器人技术研发",
|
||||
"成立浙江龙飞电驱,专注航空电机业务",
|
||||
"发布AI外骨骼机器人及灵巧手",
|
||||
"布局高爆发关节模组、伺服驱动器等人形机器人关键部件"
|
||||
],
|
||||
"growth_drivers": [
|
||||
"设备更新政策推动工业电机需求",
|
||||
"机器人产业快速发展",
|
||||
"低空经济政策支持",
|
||||
"出海战略加速"
|
||||
],
|
||||
"risk_factors": [
|
||||
"机器人业务营收占比仅2.71%,占比较低",
|
||||
"工业需求景气度波动",
|
||||
"原料价格波动风险"
|
||||
],
|
||||
"investment_rating": "积极关注",
|
||||
"fundamental_score": 0.7
|
||||
},
|
||||
"adjustment_factor": 1.1
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 890 KiB |
@@ -0,0 +1,384 @@
|
||||
# historical_backtest.py
|
||||
import pandas as pd
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('ignore')
|
||||
|
||||
# 设置中文字体
|
||||
plt.rcParams['font.sans-serif'] = ['SimHei']
|
||||
plt.rcParams['axes.unicode_minus'] = False
|
||||
|
||||
|
||||
class HistoricalBacktester:
|
||||
"""
|
||||
历史回测类:用历史数据验证模型预测效果
|
||||
"""
|
||||
|
||||
def __init__(self, data_dir, initial_capital=100000):
|
||||
self.data_dir = data_dir
|
||||
self.initial_capital = initial_capital
|
||||
|
||||
def load_historical_data(self, stock_code):
|
||||
"""加载历史数据"""
|
||||
csv_file = os.path.join(self.data_dir, f"{stock_code}_stock_data.csv")
|
||||
if not os.path.exists(csv_file):
|
||||
raise FileNotFoundError(f"数据文件不存在: {csv_file}")
|
||||
|
||||
df = pd.read_csv(csv_file, encoding='utf-8-sig')
|
||||
|
||||
# 标准化列名
|
||||
column_mapping = {
|
||||
'日期': 'date',
|
||||
'开盘价': 'open',
|
||||
'最高价': 'high',
|
||||
'最低价': 'low',
|
||||
'收盘价': 'close',
|
||||
'成交量': 'volume',
|
||||
'成交额': 'amount'
|
||||
}
|
||||
|
||||
for old_col, new_col in column_mapping.items():
|
||||
if old_col in df.columns:
|
||||
df = df.rename(columns={old_col: new_col})
|
||||
|
||||
df['date'] = pd.to_datetime(df['date'])
|
||||
df.set_index('date', inplace=True)
|
||||
df = df.sort_index()
|
||||
|
||||
print(f"✅ 加载历史数据: {len(df)} 条记录")
|
||||
print(f"时间范围: {df.index.min()} 到 {df.index.max()}")
|
||||
|
||||
return df
|
||||
|
||||
def simulate_model_prediction(self, df, lookback_days=60, pred_days=30):
|
||||
"""
|
||||
模拟模型预测:使用历史数据进行"预测",然后与实际结果对比
|
||||
"""
|
||||
results = []
|
||||
|
||||
# 从数据中选取多个时间点进行"预测"
|
||||
test_points = range(lookback_days, len(df) - pred_days, pred_days)
|
||||
|
||||
for start_idx in test_points:
|
||||
# 模拟预测:使用前lookback_days天数据"预测"后pred_days天
|
||||
historical_data = df.iloc[start_idx - lookback_days:start_idx]
|
||||
actual_future = df.iloc[start_idx:start_idx + pred_days]
|
||||
|
||||
# 简单的预测策略(这里应该替换为您的实际模型预测)
|
||||
# 这里使用移动平均作为示例预测
|
||||
pred_close = self.simple_prediction(historical_data, pred_days)
|
||||
|
||||
# 记录结果
|
||||
for i in range(min(len(pred_close), len(actual_future))):
|
||||
results.append({
|
||||
'date': actual_future.index[i],
|
||||
'actual_close': actual_future['close'].iloc[i],
|
||||
'predicted_close': pred_close[i],
|
||||
'lookback_start': historical_data.index[0],
|
||||
'prediction_date': historical_data.index[-1]
|
||||
})
|
||||
|
||||
return pd.DataFrame(results)
|
||||
|
||||
def simple_prediction(self, historical_data, pred_days):
|
||||
"""简单的预测方法(示例)"""
|
||||
# 使用移动平均 + 随机波动作为预测
|
||||
last_price = historical_data['close'].iloc[-1]
|
||||
avg_volatility = historical_data['close'].pct_change().std()
|
||||
|
||||
predictions = []
|
||||
current_price = last_price
|
||||
|
||||
for _ in range(pred_days):
|
||||
# 模拟价格变化(正态分布)
|
||||
change = np.random.normal(0, avg_volatility)
|
||||
current_price = current_price * (1 + change)
|
||||
predictions.append(current_price)
|
||||
|
||||
return predictions
|
||||
|
||||
def calculate_prediction_accuracy(self, results_df):
|
||||
"""计算预测准确率"""
|
||||
results_df['error'] = results_df['predicted_close'] - results_df['actual_close']
|
||||
results_df['error_pct'] = results_df['error'] / results_df['actual_close']
|
||||
results_df['abs_error_pct'] = abs(results_df['error_pct'])
|
||||
|
||||
accuracy_metrics = {
|
||||
'平均绝对误差率': results_df['abs_error_pct'].mean(),
|
||||
'预测准确率': (results_df['abs_error_pct'] < 0.05).mean(), # 误差小于5%算准确
|
||||
'方向准确率': (np.sign(results_df['predicted_close'].diff()) ==
|
||||
np.sign(results_df['actual_close'].diff())).mean(),
|
||||
'相关系数': results_df['predicted_close'].corr(results_df['actual_close'])
|
||||
}
|
||||
|
||||
return accuracy_metrics
|
||||
|
||||
def run_trading_strategy(self, results_df, threshold=0.03):
|
||||
"""基于预测结果运行交易策略"""
|
||||
capital = self.initial_capital
|
||||
position = 0
|
||||
trades = []
|
||||
portfolio_values = []
|
||||
|
||||
# 按日期排序
|
||||
results_df = results_df.sort_index()
|
||||
|
||||
for date, row in results_df.iterrows():
|
||||
current_price = row['actual_close']
|
||||
predicted_price = row['predicted_close']
|
||||
predicted_return = (predicted_price - current_price) / current_price
|
||||
|
||||
# 交易逻辑
|
||||
if position == 0 and predicted_return > threshold:
|
||||
# 买入信号
|
||||
shares = int(capital / current_price)
|
||||
if shares > 0:
|
||||
position = shares
|
||||
capital -= shares * current_price
|
||||
trades.append({
|
||||
'date': date,
|
||||
'action': 'BUY',
|
||||
'price': current_price,
|
||||
'shares': shares,
|
||||
'reason': f'预测上涨{predicted_return:.2%}'
|
||||
})
|
||||
|
||||
elif position > 0 and predicted_return < -threshold:
|
||||
# 卖出信号
|
||||
capital += position * current_price
|
||||
trades.append({
|
||||
'date': date,
|
||||
'action': 'SELL',
|
||||
'price': current_price,
|
||||
'shares': position,
|
||||
'reason': f'预测下跌{predicted_return:.2%}'
|
||||
})
|
||||
position = 0
|
||||
|
||||
# 计算当前资产总值
|
||||
portfolio_value = capital + position * current_price
|
||||
portfolio_values.append({
|
||||
'date': date,
|
||||
'portfolio_value': portfolio_value,
|
||||
'position': position,
|
||||
'price': current_price
|
||||
})
|
||||
|
||||
return pd.DataFrame(portfolio_values), trades
|
||||
|
||||
def calculate_performance(self, portfolio_df, trades):
|
||||
"""计算策略表现"""
|
||||
portfolio_df = portfolio_df.set_index('date')
|
||||
returns = portfolio_df['portfolio_value'].pct_change().dropna()
|
||||
|
||||
total_return = (portfolio_df['portfolio_value'].iloc[-1] - self.initial_capital) / self.initial_capital
|
||||
|
||||
if len(returns) > 0:
|
||||
annual_return = (1 + total_return) ** (252 / len(returns)) - 1
|
||||
volatility = returns.std() * np.sqrt(252)
|
||||
sharpe_ratio = (annual_return - 0.03) / volatility if volatility > 0 else 0
|
||||
|
||||
# 最大回撤
|
||||
cumulative = (1 + returns).cumprod()
|
||||
peak = cumulative.expanding().max()
|
||||
drawdown = (cumulative - peak) / peak
|
||||
max_drawdown = drawdown.min()
|
||||
else:
|
||||
annual_return = 0
|
||||
volatility = 0
|
||||
sharpe_ratio = 0
|
||||
max_drawdown = 0
|
||||
|
||||
# 买入持有策略对比
|
||||
buy_hold_return = (portfolio_df['price'].iloc[-1] - portfolio_df['price'].iloc[0]) / portfolio_df['price'].iloc[
|
||||
0]
|
||||
|
||||
performance = {
|
||||
'策略总收益': total_return,
|
||||
'策略年化收益': annual_return,
|
||||
'买入持有收益': buy_hold_return,
|
||||
'波动率': volatility,
|
||||
'夏普比率': sharpe_ratio,
|
||||
'最大回撤': max_drawdown,
|
||||
'交易次数': len(trades),
|
||||
'最终资金': portfolio_df['portfolio_value'].iloc[-1],
|
||||
'超额收益': total_return - buy_hold_return
|
||||
}
|
||||
|
||||
return performance
|
||||
|
||||
def plot_comparison(self, results_df, portfolio_df, stock_code, output_dir):
|
||||
"""绘制预测对比图表"""
|
||||
fig, (ax1, ax2, ax3) = plt.subplots(3, 1, figsize=(15, 12))
|
||||
|
||||
# 1. 价格预测对比
|
||||
ax1.plot(results_df.index, results_df['actual_close'],
|
||||
label='实际价格', color='blue', linewidth=2)
|
||||
ax1.plot(results_df.index, results_df['predicted_close'],
|
||||
label='预测价格', color='red', linestyle='--', alpha=0.7)
|
||||
ax1.set_ylabel('价格 (元)')
|
||||
ax1.legend()
|
||||
ax1.set_title(f'{stock_code} - 价格预测 vs 实际走势', fontsize=14, fontweight='bold')
|
||||
ax1.grid(True, alpha=0.3)
|
||||
|
||||
# 2. 预测误差
|
||||
ax2.bar(results_df.index, results_df['error_pct'] * 100,
|
||||
alpha=0.6, color='orange')
|
||||
ax2.axhline(y=0, color='black', linestyle='-', linewidth=1)
|
||||
ax2.set_ylabel('预测误差 (%)')
|
||||
ax2.set_title('预测误差分析')
|
||||
ax2.grid(True, alpha=0.3)
|
||||
|
||||
# 3. 策略表现
|
||||
ax3.plot(portfolio_df['date'], portfolio_df['portfolio_value'],
|
||||
label='策略资金曲线', color='green', linewidth=2)
|
||||
ax3.axhline(y=self.initial_capital, color='red', linestyle='--',
|
||||
label=f'初始资金 ({self.initial_capital:,.0f}元)')
|
||||
|
||||
# 买入持有对比
|
||||
initial_shares = self.initial_capital / portfolio_df['price'].iloc[0]
|
||||
buy_hold_values = portfolio_df['price'] * initial_shares
|
||||
ax3.plot(portfolio_df['date'], buy_hold_values,
|
||||
label='买入持有策略', color='blue', linestyle=':', alpha=0.7)
|
||||
|
||||
ax3.set_ylabel('资金 (元)')
|
||||
ax3.set_xlabel('日期')
|
||||
ax3.legend()
|
||||
ax3.set_title('策略表现对比')
|
||||
ax3.grid(True, alpha=0.3)
|
||||
|
||||
plt.tight_layout()
|
||||
|
||||
# 保存图表
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
chart_file = os.path.join(output_dir, f'{stock_code}_historical_backtest.png')
|
||||
plt.savefig(chart_file, dpi=300, bbox_inches='tight')
|
||||
print(f"📊 历史回测图表已保存: {chart_file}")
|
||||
|
||||
plt.show()
|
||||
|
||||
def run_complete_backtest(self, stock_code, output_dir, lookback_days=60, pred_days=30, threshold=0.03):
|
||||
"""运行完整的历史回测"""
|
||||
print(f"🎯 开始 {stock_code} 历史回测分析")
|
||||
print("=" * 60)
|
||||
|
||||
try:
|
||||
# 1. 加载历史数据
|
||||
print("步骤1: 加载历史数据...")
|
||||
df = self.load_historical_data(stock_code)
|
||||
|
||||
# 2. 模拟模型预测
|
||||
print("步骤2: 模拟模型预测...")
|
||||
results_df = self.simulate_model_prediction(df, lookback_days, pred_days)
|
||||
|
||||
# 3. 计算预测准确率
|
||||
print("步骤3: 计算预测准确率...")
|
||||
accuracy_metrics = self.calculate_prediction_accuracy(results_df)
|
||||
|
||||
# 4. 运行交易策略
|
||||
print("步骤4: 运行交易策略...")
|
||||
portfolio_df, trades = self.run_trading_strategy(results_df, threshold)
|
||||
|
||||
# 5. 计算策略表现
|
||||
print("步骤5: 计算策略表现...")
|
||||
performance = self.calculate_performance(portfolio_df, trades)
|
||||
|
||||
# 6. 绘制结果
|
||||
print("步骤6: 生成回测图表...")
|
||||
self.plot_comparison(results_df, portfolio_df, stock_code, output_dir)
|
||||
|
||||
# 7. 打印报告
|
||||
print("\n" + "=" * 70)
|
||||
print(f"📊 {stock_code} 历史回测报告")
|
||||
print("=" * 70)
|
||||
|
||||
print("\n🔍 预测准确率分析:")
|
||||
for metric, value in accuracy_metrics.items():
|
||||
if isinstance(value, float):
|
||||
print(f" {metric}: {value:.2%}")
|
||||
else:
|
||||
print(f" {metric}: {value:.4f}")
|
||||
|
||||
print("\n💰 策略表现分析:")
|
||||
for metric, value in performance.items():
|
||||
if isinstance(value, float):
|
||||
if '收益' in metric or '回撤' in metric:
|
||||
print(f" {metric}: {value:.2%}")
|
||||
else:
|
||||
print(f" {metric}: {value:.4f}")
|
||||
else:
|
||||
print(f" {metric}: {value}")
|
||||
|
||||
print(f"\n📈 交易统计:")
|
||||
print(f" 总交易次数: {len(trades)}")
|
||||
print(f" 买入次数: {len([t for t in trades if t['action'] == 'BUY'])}")
|
||||
print(f" 卖出次数: {len([t for t in trades if t['action'] == 'SELL'])}")
|
||||
|
||||
if len(trades) > 0:
|
||||
print(f"\n最近5次交易:")
|
||||
for trade in trades[-5:]:
|
||||
print(f" {trade['date'].strftime('%Y-%m-%d')} {trade['action']} "
|
||||
f"{trade['shares']}股 @ {trade['price']:.2f}元 - {trade['reason']}")
|
||||
|
||||
return accuracy_metrics, performance, results_df
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ 回测过程中出现错误: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return None, None, None
|
||||
|
||||
|
||||
def main():
|
||||
"""主函数"""
|
||||
# 配置参数
|
||||
BACKTEST_CONFIG = {
|
||||
"stock_code": "300418",
|
||||
"data_dir": r"D:\lianghuajiaoyi\Kronos\examples\data",
|
||||
"output_dir": r"D:\lianghuajiaoyi\Kronos\examples\historical_backtest",
|
||||
"initial_capital": 100000,
|
||||
"lookback_days": 60, # 使用60天历史数据
|
||||
"pred_days": 30, # 预测30天
|
||||
"threshold": 0.03 # 3%的交易阈值
|
||||
}
|
||||
|
||||
print("🤖 Kronos模型历史回测系统")
|
||||
print("=" * 50)
|
||||
print(f"回测股票: {BACKTEST_CONFIG['stock_code']}")
|
||||
print(f"回看天数: {BACKTEST_CONFIG['lookback_days']}天")
|
||||
print(f"预测天数: {BACKTEST_CONFIG['pred_days']}天")
|
||||
print(f"初始资金: {BACKTEST_CONFIG['initial_capital']:,.0f}元")
|
||||
print()
|
||||
|
||||
# 创建回测器并运行
|
||||
backtester = HistoricalBacktester(
|
||||
data_dir=BACKTEST_CONFIG["data_dir"],
|
||||
initial_capital=BACKTEST_CONFIG["initial_capital"]
|
||||
)
|
||||
|
||||
accuracy, performance, results = backtester.run_complete_backtest(
|
||||
stock_code=BACKTEST_CONFIG["stock_code"],
|
||||
output_dir=BACKTEST_CONFIG["output_dir"],
|
||||
lookback_days=BACKTEST_CONFIG["lookback_days"],
|
||||
pred_days=BACKTEST_CONFIG["pred_days"],
|
||||
threshold=BACKTEST_CONFIG["threshold"]
|
||||
)
|
||||
|
||||
if accuracy and performance:
|
||||
print(f"\n✅ {BACKTEST_CONFIG['stock_code']} 历史回测完成!")
|
||||
|
||||
# 简单结论
|
||||
if performance['超额收益'] > 0:
|
||||
print("🎉 结论: 模型策略跑赢了买入持有策略!")
|
||||
else:
|
||||
print("⚠️ 结论: 模型策略未能跑赢买入持有策略。")
|
||||
|
||||
print(f"📁 详细结果保存在: {BACKTEST_CONFIG['output_dir']}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,41 @@
|
||||
{
|
||||
"timestamp": "2025-09-30 22:45:36",
|
||||
"market_analysis": {
|
||||
"is_main_uptrend": false,
|
||||
"trend_strength": 0.5,
|
||||
"market_status": "未知",
|
||||
"support_level": null,
|
||||
"resistance_level": null
|
||||
},
|
||||
"sector_analysis": {
|
||||
"industry": "电机",
|
||||
"sector_momentum": 0.5,
|
||||
"sector_limit_up_count": 0,
|
||||
"is_sector_hot": false,
|
||||
"resonance_score": 0.35
|
||||
},
|
||||
"macro_analysis": {
|
||||
"us_rate_cycle": {
|
||||
"current_rate": 4.25,
|
||||
"trend": "降息周期",
|
||||
"expected_cuts_2025": 2,
|
||||
"expected_cuts_2026": 2,
|
||||
"impact_on_emerging_markets": "positive",
|
||||
"usd_index_support": 95.0
|
||||
},
|
||||
"domestic_policy": {
|
||||
"monetary_policy": "宽松",
|
||||
"fiscal_policy": "积极",
|
||||
"market_liquidity": "充足",
|
||||
"holiday_effect": {
|
||||
"effect": "节前震荡,节后上涨概率大",
|
||||
"historical_win_rate": 0.8,
|
||||
"expected_return": 0.0227,
|
||||
"period": "节后5个交易日"
|
||||
}
|
||||
},
|
||||
"global_liquidity_outlook": "改善",
|
||||
"overall_macro_score": 0.7
|
||||
},
|
||||
"adjustment_factor": 1.00776
|
||||
}
|
||||
Reference in New Issue
Block a user