Files
2025-10-27 19:45:33 +08:00

399 lines
12 KiB
Plaintext
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.
{
"cells": [
{
"cell_type": "markdown",
"id": "de76d792",
"metadata": {},
"source": [
"# LSTM外汇价格预测 - 精简版\n",
"\n",
"本Notebook包含LSTM模型训练的核心代码,用于预测EURUSD价格。\n",
"\n",
"## 主要步骤\n",
"1. 导入库和设置设备\n",
"2. 从MT5获取历史数据\n",
"3. 数据归一化\n",
"4. 构造训练/验证数据集\n",
"5. 定义LSTM模型\n",
"6. 训练模型(早停机制)\n",
"7. 可视化训练过程\n",
"8. 导出ONNX模型"
]
},
{
"cell_type": "markdown",
"id": "063db750",
"metadata": {},
"source": [
"## 1. 导入库和设置设备"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "fee98c7a",
"metadata": {},
"outputs": [],
"source": [
"import torch\n",
"import torch.nn as nn\n",
"import numpy as np\n",
"import pandas as pd\n",
"import MetaTrader5 as mt5\n",
"from sklearn.preprocessing import MinMaxScaler\n",
"from datetime import datetime, timedelta\n",
"import matplotlib.pyplot as plt\n",
"import os\n",
"\n",
"# 设置中文字体\n",
"plt.rcParams['font.sans-serif'] = ['SimHei']\n",
"plt.rcParams['axes.unicode_minus'] = False\n",
"\n",
"# 选择计算设备(GPU或CPU\n",
"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n",
"print(f\"使用设备: {device}\")"
]
},
{
"cell_type": "markdown",
"id": "19c6b9d0",
"metadata": {},
"source": [
"## 2. 从MT5获取EURUSD历史数据"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "eee75337",
"metadata": {},
"outputs": [],
"source": [
"# 连接MT5\n",
"mt5.initialize()\n",
"\n",
"# 设置时间范围(最近5年)\n",
"end_date = datetime.now()\n",
"start_date = end_date - timedelta(days=5*365)\n",
"\n",
"# 获取EURUSD H1数据\n",
"rates = mt5.copy_rates_range(\"EURUSD\", mt5.TIMEFRAME_H1, start_date, end_date)\n",
"mt5.shutdown()\n",
"\n",
"# 转换为DataFrame\n",
"df = pd.DataFrame(rates)\n",
"print(f\"获取到 {len(df):,} 条数据\")\n",
"df.head()"
]
},
{
"cell_type": "markdown",
"id": "4fda9d49",
"metadata": {},
"source": [
"## 3. 数据归一化"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c5dccb3f",
"metadata": {},
"outputs": [],
"source": [
"# 创建归一化器,将数据缩放到[0,1]范围\n",
"scaler = MinMaxScaler()\n",
"\n",
"# 对开、高、低、收、成交量进行归一化\n",
"df[['open', 'high', 'low', 'close', 'tick_volume']] = scaler.fit_transform(\n",
" df[['open', 'high', 'low', 'close', 'tick_volume']]\n",
")\n",
"\n",
"print(\"归一化完成\")\n",
"print(f\"最小值: {scaler.data_min_}\")\n",
"print(f\"最大值: {scaler.data_max_}\")"
]
},
{
"cell_type": "markdown",
"id": "15ddac0b",
"metadata": {},
"source": [
"## 4. 构造时间序列数据集"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "4887abce",
"metadata": {},
"outputs": [],
"source": [
"def create_dataset(data, lookback=10):\n",
" \"\"\"将时间序列转换为监督学习数据集\"\"\"\n",
" X, y = [], []\n",
" for i in range(len(data) - lookback - 1):\n",
" X.append(data[i:i+lookback]) # 输入:前10根K线\n",
" y.append(data[i+lookback, 3]) # 输出:第11根K线的收盘价\n",
" return np.array(X, dtype=np.float32), np.array(y, dtype=np.float32).reshape(-1, 1)\n",
"\n",
"# 提取特征\n",
"features = df[['open', 'high', 'low', 'close', 'tick_volume']].values\n",
"X, y = create_dataset(features)\n",
"\n",
"# 划分训练集(80%)和验证集(20%\n",
"train_size = int(len(X) * 0.8)\n",
"X_train = torch.from_numpy(X[:train_size]).to(device)\n",
"y_train = torch.from_numpy(y[:train_size]).to(device)\n",
"X_val = torch.from_numpy(X[train_size:]).to(device)\n",
"y_val = torch.from_numpy(y[train_size:]).to(device)\n",
"\n",
"print(f\"训练集: {X_train.shape}\")\n",
"print(f\"验证集: {X_val.shape}\")"
]
},
{
"cell_type": "markdown",
"id": "208c234b",
"metadata": {},
"source": [
"## 5. 定义LSTM模型"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c2916d2d",
"metadata": {},
"outputs": [],
"source": [
"class SimpleLSTM(nn.Module):\n",
" \"\"\"简单的LSTM神经网络\"\"\"\n",
" def __init__(self):\n",
" super().__init__()\n",
" # LSTM层:输入5个特征,隐藏层20个神经元\n",
" self.lstm = nn.LSTM(input_size=5, hidden_size=20, num_layers=1, batch_first=True)\n",
" # 全连接层:将20维输出映射到1维(预测值)\n",
" self.fc = nn.Linear(20, 1)\n",
"\n",
" def forward(self, x):\n",
" lstm_out, _ = self.lstm(x)\n",
" return self.fc(lstm_out[:, -1, :]) # 取最后一个时间步的输出\n",
"\n",
"# 创建模型\n",
"model = SimpleLSTM().to(device)\n",
"optimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n",
"criterion = nn.MSELoss()\n",
"\n",
"print(\"模型创建完成\")"
]
},
{
"cell_type": "markdown",
"id": "9e6f4eeb",
"metadata": {},
"source": [
"## 6. 训练模型(带早停机制)"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "04481e99",
"metadata": {},
"outputs": [],
"source": [
"# 训练参数\n",
"max_epochs = 500\n",
"patience = 20 # 验证Loss不降20轮则停止\n",
"\n",
"# 记录训练过程\n",
"train_losses = []\n",
"val_losses = []\n",
"best_val_loss = float('inf')\n",
"best_epoch = 0\n",
"patience_counter = 0\n",
"best_model_state = None\n",
"\n",
"print(\"开始训练...\")\n",
"for epoch in range(max_epochs):\n",
" # 训练阶段\n",
" model.train()\n",
" optimizer.zero_grad()\n",
" train_output = model(X_train)\n",
" train_loss = criterion(train_output, y_train)\n",
" train_loss.backward()\n",
" optimizer.step()\n",
" \n",
" # 验证阶段\n",
" model.eval()\n",
" with torch.no_grad():\n",
" val_output = model(X_val)\n",
" val_loss = criterion(val_output, y_val)\n",
" \n",
" # 记录Loss\n",
" train_losses.append(train_loss.item())\n",
" val_losses.append(val_loss.item())\n",
" \n",
" # 每10轮打印一次\n",
" if (epoch + 1) % 10 == 0:\n",
" print(f\"Epoch {epoch+1:3d} | 训练Loss: {train_loss.item():.6f} | 验证Loss: {val_loss.item():.6f}\")\n",
" \n",
" # 早停检查\n",
" if val_loss.item() < best_val_loss:\n",
" best_val_loss = val_loss.item()\n",
" best_epoch = epoch + 1\n",
" patience_counter = 0\n",
" best_model_state = {k: v.cpu().clone() for k, v in model.state_dict().items()}\n",
" else:\n",
" patience_counter += 1\n",
" \n",
" if patience_counter >= patience:\n",
" print(f\"\\n早停触发!最佳模型在Epoch {best_epoch}\")\n",
" break\n",
"\n",
"# 恢复最佳模型\n",
"if best_model_state:\n",
" model.load_state_dict(best_model_state)\n",
" model = model.to(device)\n",
" print(f\"已恢复最佳模型 (Epoch {best_epoch}, 验证Loss: {best_val_loss:.6f})\")"
]
},
{
"cell_type": "markdown",
"id": "1669bc75",
"metadata": {},
"source": [
"## 7. 可视化训练过程"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "cd1011f3",
"metadata": {},
"outputs": [],
"source": [
"# 绘制训练和验证Loss曲线\n",
"plt.figure(figsize=(12, 5))\n",
"\n",
"# 完整曲线\n",
"plt.subplot(1, 2, 1)\n",
"plt.plot(train_losses, label='训练Loss', color='blue', linewidth=2)\n",
"plt.plot(val_losses, label='验证Loss', color='orange', linewidth=2)\n",
"plt.axvline(x=best_epoch-1, color='red', linestyle='--', label=f'最佳模型 (Epoch {best_epoch})')\n",
"plt.xlabel('Epoch')\n",
"plt.ylabel('Loss')\n",
"plt.title('训练和验证损失曲线')\n",
"plt.legend()\n",
"plt.grid(True, alpha=0.3)\n",
"\n",
"# 后80%放大曲线\n",
"plt.subplot(1, 2, 2)\n",
"start_idx = int(len(train_losses) * 0.2)\n",
"plt.plot(range(start_idx, len(train_losses)), train_losses[start_idx:], label='训练Loss', color='blue', linewidth=2)\n",
"plt.plot(range(start_idx, len(val_losses)), val_losses[start_idx:], label='验证Loss', color='orange', linewidth=2)\n",
"plt.axvline(x=best_epoch-1, color='red', linestyle='--', label=f'最佳模型')\n",
"plt.xlabel('Epoch')\n",
"plt.ylabel('Loss')\n",
"plt.title('后期收敛细节')\n",
"plt.legend()\n",
"plt.grid(True, alpha=0.3)\n",
"\n",
"plt.tight_layout()\n",
"plt.savefig('training_loss.png', dpi=150)\n",
"plt.show()\n",
"\n",
"print(\"训练曲线已保存为 training_loss.png\")"
]
},
{
"cell_type": "markdown",
"id": "80b9c1c2",
"metadata": {},
"source": [
"## 8. 导出ONNX模型和归一化参数"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c5b82ee2",
"metadata": {},
"outputs": [],
"source": [
"# 导出ONNX模型\n",
"model.eval()\n",
"model_cpu = model.cpu()\n",
"dummy_input = torch.randn(1, 10, 5)\n",
"\n",
"torch.onnx.export(\n",
" model_cpu,\n",
" dummy_input,\n",
" \"lstm_model.onnx\",\n",
" export_params=True,\n",
" opset_version=11,\n",
" input_names=['input'],\n",
" output_names=['output']\n",
")\n",
"\n",
"print(\"✓ ONNX模型已导出: lstm_model.onnx\")\n",
"\n",
"# 保存归一化参数\n",
"np.save('scaler_params.npy', {'min': scaler.data_min_, 'max': scaler.data_max_})\n",
"print(\"✓ 归一化参数已保存: scaler_params.npy\")\n",
"\n",
"# 输出MQL5格式的归一化参数\n",
"print(\"\\n=\" * 60)\n",
"print(\"归一化参数 (复制到LSTM_EA.mq5):\")\n",
"print(\"=\" * 60)\n",
"print(f\"double data_min[5] = {{{', '.join([f'{x:.5f}' for x in scaler.data_min_])}}};\")\n",
"print(f\"double data_max[5] = {{{', '.join([f'{x:.5f}' for x in scaler.data_max_])}}};\")\n",
"print(\"=\" * 60)"
]
},
{
"cell_type": "markdown",
"id": "cb21d2d4",
"metadata": {},
"source": [
"## 9. 训练总结"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "b6ac5bc9",
"metadata": {},
"outputs": [],
"source": [
"print(\"=\" * 60)\n",
"print(\"训练总结\")\n",
"print(\"=\" * 60)\n",
"print(f\"总训练轮数: {len(train_losses)} epochs\")\n",
"print(f\"最佳模型: Epoch {best_epoch}\")\n",
"print(f\"最佳验证Loss: {best_val_loss:.6f}\")\n",
"print(f\"最终训练Loss: {train_losses[-1]:.6f}\")\n",
"print(f\"最终验证Loss: {val_losses[-1]:.6f}\")\n",
"\n",
"loss_diff = abs(train_losses[-1] - val_losses[-1])\n",
"print(f\"\\nLoss差异: {loss_diff:.6f}\")\n",
"if loss_diff < 0.0001:\n",
" print(\"✓ 模型状态: 良好\")\n",
"elif loss_diff < 0.001:\n",
" print(\"⚠ 模型状态: 尚可,有轻微过拟合\")\n",
"else:\n",
" print(\"✗ 模型状态: 可能过拟合\")\n",
"print(\"=\" * 60)"
]
}
],
"metadata": {
"language_info": {
"name": "python"
}
},
"nbformat": 4,
"nbformat_minor": 5
}