337 lines
9.4 KiB
Python
337 lines
9.4 KiB
Python
"""
|
||
CSV 数据加载器 — 支持 MT5 History Center 格式 (quant data manager 导出)
|
||
|
||
功能:
|
||
- 加载 M1 CSV (无表头 9 列格式)
|
||
- 多周期重采样 (M5/M15/M30/H1/H4/D1)
|
||
- 品种自动识别 (从文件名)
|
||
- spread 自动转 slippage (百分比)
|
||
- 时区处理 (UTC+2/UTC-3 等, 从文件名提取)
|
||
|
||
CSV 格式 (MT5 History Center 标准):
|
||
date,time,open,high,low,close,vol,vol_real,spread
|
||
2021.07.06,01:00,1791.37,1791.37,1790.55,1791.27,25850.0,25850.0,98
|
||
|
||
用法:
|
||
from app.data_loader import load_csv, find_csv_for_symbol, compute_slippage
|
||
|
||
# 自动查找 + 加载
|
||
path = find_csv_for_symbol("XAUUSD")
|
||
df = load_csv(path, timeframe="H1")
|
||
|
||
# 显式指定文件
|
||
df = load_csv("data/XAUUSD.csv", symbol="XAUUSD", timeframe="H1")
|
||
|
||
# 计算 slippage (注入 PyBacktestConfig)
|
||
slippage = compute_slippage(df["spread"], df["close"], "XAUUSD")
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
import re
|
||
from datetime import timedelta, timezone
|
||
from typing import Optional
|
||
|
||
import numpy as np
|
||
import pandas as pd
|
||
|
||
|
||
# ============================================================================
|
||
# 品种配置 — 点值 (最小变动单位)
|
||
# ============================================================================
|
||
|
||
SYMBOL_TICK_SIZE = {
|
||
"XAUUSD": 0.01, # 黄金: 0.01 美元/点
|
||
"XAGUSD": 0.001, # 白银
|
||
"EURUSD": 0.00001, # 5 位报价
|
||
"GBPUSD": 0.00001,
|
||
"AUDUSD": 0.00001,
|
||
"NZDUSD": 0.00001,
|
||
"USDCAD": 0.00001,
|
||
"USDCHF": 0.00001,
|
||
"EURJPY": 0.001, # 3 位报价 (JPY 货币对)
|
||
"USDJPY": 0.001,
|
||
"GBPJPY": 0.001,
|
||
"AUDJPY": 0.001,
|
||
"EURGBP": 0.00001,
|
||
"EURAUD": 0.00001,
|
||
"EURCHF": 0.00001,
|
||
"GBPCHF": 0.00001,
|
||
"CHFJPY": 0.001,
|
||
}
|
||
|
||
# MT5 周期 → pandas 频率
|
||
TIMEFRAME_MAP = {
|
||
"M1": "1min",
|
||
"M5": "5min",
|
||
"M15": "15min",
|
||
"M30": "30min",
|
||
"H1": "1h",
|
||
"H2": "2h",
|
||
"H4": "4h",
|
||
"H6": "6h",
|
||
"H8": "8h",
|
||
"H12": "12h",
|
||
"D1": "1D",
|
||
"W1": "1W",
|
||
"MN1": "1ME",
|
||
}
|
||
|
||
|
||
# ============================================================================
|
||
# 文件名解析
|
||
# ============================================================================
|
||
|
||
# 已知品种列表 (按长度降序匹配, 避免 GBPUSD 匹配 USD)
|
||
_KNOWN_SYMBOLS = sorted(SYMBOL_TICK_SIZE.keys(), key=len, reverse=True)
|
||
|
||
|
||
def extract_symbol_from_filename(filename: str) -> Optional[str]:
|
||
"""
|
||
从文件名提取品种代码
|
||
|
||
文件名示例:
|
||
2021.7.6-2026.07.03M1XAUUSD_TICK_UTCPlus02.csv → XAUUSD
|
||
2021.7.6-2026.07.03_M1USDJPY_TICK_UTCPlus02.csv → USDJPY
|
||
2021.7.6-2026.07.03EURUSD_M1_TICK_UTCPlus02.csv → EURUSD
|
||
"""
|
||
name = os.path.basename(filename).upper()
|
||
for sym in _KNOWN_SYMBOLS:
|
||
if sym in name:
|
||
return sym
|
||
return None
|
||
|
||
|
||
def extract_timezone_from_filename(filename: str) -> timezone:
|
||
"""
|
||
从文件名提取时区
|
||
|
||
文件名示例:
|
||
...UTCPlus02 → UTC+2 (MT5 标准服务器时间)
|
||
...UTCPlus03 → UTC+3 (夏令时)
|
||
...UTCMINUS05 → UTC-5
|
||
"""
|
||
name = os.path.basename(filename).upper()
|
||
m = re.search(r"UTCPLUS(\d+)", name)
|
||
if m:
|
||
return timezone(timedelta(hours=int(m.group(1))))
|
||
m = re.search(r"UTCMINUS(\d+)", name)
|
||
if m:
|
||
return timezone(timedelta(hours=-int(m.group(1))))
|
||
return timezone.utc
|
||
|
||
|
||
# ============================================================================
|
||
# CSV 加载
|
||
# ============================================================================
|
||
|
||
def load_csv(
|
||
file_path: str,
|
||
symbol: Optional[str] = None,
|
||
timeframe: str = "M1",
|
||
drop_weekend: bool = False,
|
||
) -> pd.DataFrame:
|
||
"""
|
||
加载 MT5 History Center 格式 CSV
|
||
|
||
参数:
|
||
file_path: CSV 文件路径
|
||
symbol: 品种代码 (None 时自动从文件名识别)
|
||
timeframe: 目标周期 (M1/M5/M15/M30/H1/H4/D1/W1/MN1)
|
||
drop_weekend: 是否过滤周末行 (默认 False, 因为外汇 CSV 已无周末数据)
|
||
|
||
返回:
|
||
DataFrame, 列: time, open, high, low, close, tick_volume, spread
|
||
time 列为 datetime64[ns, tz]
|
||
"""
|
||
if not os.path.exists(file_path):
|
||
raise FileNotFoundError(f"CSV 文件不存在: {file_path}")
|
||
|
||
# 自动识别品种
|
||
if symbol is None:
|
||
symbol = extract_symbol_from_filename(file_path)
|
||
# 解析时区
|
||
tz = extract_timezone_from_filename(file_path)
|
||
|
||
# 加载 CSV (无表头, 9 列格式)
|
||
df = pd.read_csv(
|
||
file_path,
|
||
header=None,
|
||
names=["date", "time", "open", "high", "low", "close",
|
||
"volume", "real_volume", "spread"],
|
||
dtype={
|
||
"open": np.float64, "high": np.float64,
|
||
"low": np.float64, "close": np.float64,
|
||
"volume": np.float64, "real_volume": np.float64,
|
||
"spread": np.int32,
|
||
},
|
||
)
|
||
|
||
# 合并 date + time → datetime
|
||
df["time"] = pd.to_datetime(
|
||
df["date"] + " " + df["time"],
|
||
format="%Y.%m.%d %H:%M",
|
||
utc=False,
|
||
)
|
||
df["time"] = df["time"].dt.tz_localize(tz)
|
||
|
||
# 清理列
|
||
df = df.drop(columns=["date", "real_volume"])
|
||
df = df.rename(columns={"volume": "tick_volume"})
|
||
|
||
# 排序 + 去重 + 建索引
|
||
df = df.sort_values("time").drop_duplicates(subset=["time"])
|
||
df = df.set_index("time")
|
||
|
||
# 可选: 过滤周末 (外汇 CSV 通常已无周末行, 此项保险用)
|
||
if drop_weekend:
|
||
df = df[df.index.dayofweek < 5]
|
||
|
||
# 重采样到目标周期
|
||
if timeframe != "M1":
|
||
df = _resample(df, timeframe)
|
||
|
||
return df.reset_index()
|
||
|
||
|
||
def _resample(df: pd.DataFrame, timeframe: str) -> pd.DataFrame:
|
||
"""重采样 M1 到更高周期"""
|
||
if timeframe == "M1":
|
||
return df
|
||
|
||
freq = TIMEFRAME_MAP.get(timeframe)
|
||
if not freq:
|
||
raise ValueError(
|
||
f"不支持的周期: {timeframe}, 可选: {', '.join(TIMEFRAME_MAP.keys())}"
|
||
)
|
||
|
||
resampled = df.resample(freq, label="left", closed="left").agg({
|
||
"open": "first",
|
||
"high": "max",
|
||
"low": "min",
|
||
"close": "last",
|
||
"tick_volume": "sum",
|
||
"spread": "mean",
|
||
}).dropna()
|
||
|
||
return resampled
|
||
|
||
|
||
# ============================================================================
|
||
# 点差 → slippage 转换
|
||
# ============================================================================
|
||
|
||
def compute_slippage(
|
||
spread_series: pd.Series,
|
||
close: pd.Series,
|
||
symbol: str,
|
||
fallback: float = 0.0005,
|
||
) -> float:
|
||
"""
|
||
根据平均 spread 和价格计算 slippage (百分比)
|
||
|
||
公式: slippage = (avg_spread × tick_size) / avg_price
|
||
|
||
参数:
|
||
spread_series: spread 列 (点数, 如 98 表示 98 个 tick)
|
||
close: 收盘价 (用于计算平均价格)
|
||
symbol: 品种代码 (决定 tick_size)
|
||
fallback: 无法识别品种时的默认值 (默认 0.05%)
|
||
|
||
返回:
|
||
slippage 百分比 (如 0.00055 表示 0.055%)
|
||
|
||
示例:
|
||
XAUUSD: avg_spread=98, tick_size=0.01, avg_price=1791
|
||
slippage = 98 × 0.01 / 1791 = 0.000547 ≈ 0.055%
|
||
|
||
EURUSD: avg_spread=43, tick_size=0.00001, avg_price=1.186
|
||
slippage = 43 × 0.00001 / 1.186 = 0.000363 ≈ 0.036%
|
||
"""
|
||
tick_size = SYMBOL_TICK_SIZE.get(symbol)
|
||
if tick_size is None:
|
||
return fallback
|
||
|
||
avg_spread = float(spread_series.mean())
|
||
avg_price = float(close.mean())
|
||
|
||
if avg_price == 0:
|
||
return fallback
|
||
|
||
spread_cost = avg_spread * tick_size
|
||
return spread_cost / avg_price
|
||
|
||
|
||
def get_tick_size(symbol: str) -> float:
|
||
"""获取品种的 tick_size (最小变动单位)"""
|
||
return SYMBOL_TICK_SIZE.get(symbol, 0.0001)
|
||
|
||
|
||
# ============================================================================
|
||
# 自动查找
|
||
# ============================================================================
|
||
|
||
def find_csv_for_symbol(
|
||
symbol: str,
|
||
data_dir: Optional[str] = None,
|
||
) -> Optional[str]:
|
||
"""
|
||
在 data/ 目录查找匹配品种的 CSV
|
||
|
||
匹配规则: 文件名 (大写) 包含品种代码 (大写)
|
||
如查找 XAUUSD, 匹配 "2021.7.6-2026.07.03M1XAUUSD_TICK_UTCPlus02.csv"
|
||
|
||
返回:
|
||
第一个匹配的文件路径, 未找到返回 None
|
||
"""
|
||
if data_dir is None:
|
||
data_dir = os.path.join(
|
||
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
|
||
"data",
|
||
)
|
||
|
||
if not os.path.exists(data_dir):
|
||
return None
|
||
|
||
symbol_upper = symbol.upper()
|
||
# 优先匹配 M1 文件 (原始数据, 精度最高)
|
||
files = [f for f in os.listdir(data_dir) if f.endswith(".csv")]
|
||
# 先找 M1 文件
|
||
for f in files:
|
||
if symbol_upper in f.upper() and "M1" in f.upper():
|
||
return os.path.join(data_dir, f)
|
||
# 再找任意匹配
|
||
for f in files:
|
||
if symbol_upper in f.upper():
|
||
return os.path.join(data_dir, f)
|
||
|
||
return None
|
||
|
||
|
||
def list_available_symbols(data_dir: Optional[str] = None) -> list:
|
||
"""
|
||
列出 data/ 目录下所有可用品种
|
||
|
||
返回:
|
||
[(symbol, filename), ...] 列表
|
||
"""
|
||
if data_dir is None:
|
||
data_dir = os.path.join(
|
||
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
|
||
"data",
|
||
)
|
||
|
||
if not os.path.exists(data_dir):
|
||
return []
|
||
|
||
result = []
|
||
for f in sorted(os.listdir(data_dir)):
|
||
if not f.endswith(".csv"):
|
||
continue
|
||
symbol = extract_symbol_from_filename(f)
|
||
if symbol:
|
||
result.append((symbol, f))
|
||
|
||
return result
|