-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtemp.py
More file actions
210 lines (176 loc) · 7.24 KB
/
Copy pathtemp.py
File metadata and controls
210 lines (176 loc) · 7.24 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
import os
# 必须在导入 akshare 之前禁用代理
os.environ['HTTP_PROXY'] = ''
os.environ['HTTPS_PROXY'] = ''
os.environ['http_proxy'] = ''
os.environ['https_proxy'] = ''
import akshare as ak
import pandas as pd
from dataclasses import dataclass, field
from functools import lru_cache
from requests.exceptions import RequestException
import time
def retry_request(func, max_retries=3, delay=2):
"""带重试的请求包装器"""
for attempt in range(max_retries):
try:
return func()
except RequestException as e:
if attempt < max_retries - 1:
print(f"网络请求失败,{delay}秒后重试 ({attempt + 1}/{max_retries})...")
time.sleep(delay)
else:
raise e
@lru_cache(maxsize=1)
def get_stock_name_map() -> dict:
"""获取股票代码-名称映射表(带缓存)"""
df = retry_request(lambda: ak.stock_info_a_code_name())
return dict(zip(df['code'], df['name']))
def get_stock_name(symbol: str) -> str:
"""根据股票代码获取名称"""
name_map = get_stock_name_map()
return name_map.get(symbol, "未知")
@dataclass
class TradeConfig:
"""交易配置"""
initial_capital: float = 100000 # 初始本金 10万
buy_price: float = 3.8 # 低于此价格买入
sell_price: float = 4.2 # 高于此价格卖出
commission_rate: float = 0.00025 # 手续费万2.5
min_commission: float = 5 # 最低手续费5元
@dataclass
class StockConfig:
"""股票配置"""
symbol: str # 股票代码
start_date: str # 开始日期 YYYYMMDD
end_date: str # 结束日期 YYYYMMDD
name: str = field(default="") # 股票名称(可选,留空自动获取)
def __post_init__(self):
if not self.name:
self.name = get_stock_name(self.symbol)
def calc_commission(amount: float, config: TradeConfig) -> float:
"""计算交易手续费"""
return max(amount * config.commission_rate, config.min_commission)
def simulate_trading(df: pd.DataFrame, config: TradeConfig) -> dict:
"""
模拟交易
策略:低于buy_price全仓买入,高于sell_price全部卖出
"""
cash = config.initial_capital # 现金
shares = 0 # 持股数量
buy_count = 0 # 买入次数
sell_count = 0 # 卖出次数
total_commission = 0 # 总手续费
for _, row in df.iterrows():
date = row['日期']
close_price = row['收盘']
if shares == 0 and close_price < config.buy_price:
# 买入:全仓买入(按100股为一手)
max_shares = int(cash / (close_price * 100)) * 100
if max_shares > 0:
cost = max_shares * close_price
commission = calc_commission(cost, config)
if cash >= cost + commission:
shares = max_shares
cash -= cost + commission
total_commission += commission
buy_count += 1
print(f"[{date}] 买入 {shares} 股, 价格 {close_price:.2f}, "
f"花费 {cost:.2f}, 手续费 {commission:.2f}")
elif shares > 0 and close_price > config.sell_price:
# 卖出:全部卖出
revenue = shares * close_price
commission = calc_commission(revenue, config)
cash += revenue - commission
total_commission += commission
print(f"[{date}] 卖出 {shares} 股, 价格 {close_price:.2f}, "
f"收入 {revenue:.2f}, 手续费 {commission:.2f}")
shares = 0
sell_count += 1
# 计算最终资产
final_price = df.iloc[-1]['收盘']
final_asset = cash + shares * final_price
profit = final_asset - config.initial_capital
profit_rate = (profit / config.initial_capital) * 100
return {
'final_asset': final_asset,
'cash': cash,
'shares': shares,
'final_price': final_price,
'profit': profit,
'profit_rate': profit_rate,
'buy_count': buy_count,
'sell_count': sell_count,
'total_commission': total_commission
}
def run_backtest(stock_config: StockConfig, trade_config: TradeConfig):
"""运行单个股票的回测"""
print("=" * 60)
print(f"股票交易策略回测 - {stock_config.name}({stock_config.symbol})")
print(f"策略: 低于 {trade_config.buy_price} 买入, 高于 {trade_config.sell_price} 卖出")
print(f"本金: {trade_config.initial_capital:,.0f} 元")
print(f"手续费: {trade_config.commission_rate*100:.4f}%, 最低 {trade_config.min_commission} 元")
print("=" * 60)
# 获取股票数据(带重试)
print("\n正在获取股票数据...")
try:
df = retry_request(lambda: ak.stock_zh_a_hist(
symbol=stock_config.symbol,
period="daily",
start_date=stock_config.start_date,
end_date=stock_config.end_date,
adjust="" # 不复权
))
except RequestException as e:
print(f"获取股票数据失败: {e}")
print("提示:如果遇到代理问题,请取消代码第11-12行的注释以禁用代理")
return None
print(f"获取到 {len(df)} 条交易日数据")
print(f"日期范围: {df['日期'].iloc[0]} ~ {df['日期'].iloc[-1]}")
print(f"价格范围: {df['收盘'].min():.2f} ~ {df['收盘'].max():.2f}")
print()
# 模拟交易
print("交易记录:")
print("-" * 60)
result = simulate_trading(df, trade_config)
# 输出结果
print()
print("=" * 60)
print("回测结果")
print("=" * 60)
print(f"初始本金: {trade_config.initial_capital:>15,.2f} 元")
print(f"最终资产: {result['final_asset']:>15,.2f} 元")
print(f" - 现金: {result['cash']:>15,.2f} 元")
print(f" - 持股: {result['shares']:>15} 股")
print(f" - 持股市值: {result['shares'] * result['final_price']:>15,.2f} 元")
print(f"收益金额: {result['profit']:>15,.2f} 元")
print(f"收益率: {result['profit_rate']:>14.2f}%")
print(f"买入次数: {result['buy_count']:>15} 次")
print(f"卖出次数: {result['sell_count']:>15} 次")
print(f"总手续费: {result['total_commission']:>15,.2f} 元")
print("=" * 60)
return result
def main():
# ===== 股票配置(可添加多个股票进行测试)=====
stocks = [
StockConfig(symbol="000725", start_date="20250101", end_date="20251231"),
# 添加更多股票(名称会自动获取):
# StockConfig(symbol="000001", start_date="20240101", end_date="20241231"),
# StockConfig(symbol="000858", start_date="20240101", end_date="20241231"),
# 或者手动指定名称:
# StockConfig(symbol="600036", name="招商银行", start_date="20240101", end_date="20241231"),
]
# ===== 交易配置 =====
trade_config = TradeConfig(
initial_capital=100000, # 本金10万
buy_price=3.8, # 买入价格
sell_price=4.2, # 卖出价格
commission_rate=0.00025, # 万2.5
min_commission=5 # 最低5元
)
# 批量回测
for stock in stocks:
run_backtest(stock, trade_config)
print("\n")
if __name__ == "__main__":
main()