Files

337 lines
9.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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