Files
my_ai_agent_celue_/app/data_loader.py
T

337 lines
9.4 KiB
Python
Raw Normal View History

2026-07-09 21:23:10 +08:00
"""
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