数据科学和机器学习(第 31 部分):利用 CatBoost AI 模型进行交易·进阶篇
◍ 先看清数据骨架再谈策略
把 EURUSD 的日线导进 pandas,第一件事不是画信号,而是用 info() 看字段存活率。上面这段输出里,DayofWeek、DayofYear、Month 三列都是 6999 个 non-null,说明在 6999 根日 K 里没有缺日期,时间序列是连续的。 dtypes 显示整体是 4 个 float64 加 4 个 object,内存占用 492.1+ KB。这种体量在 MT5 里用 CSV 缓存或直接用 MQL5 的 File 函数读都不吃力,但在 Python 侧做滚动统计会更顺手。 实操上,建议你开 MT5 导出最近一年日线,跑一遍同样的 info(),确认自己数据的 non-null 数和字段类型。若某列 non-null 明显少于总数,后续按星期几过滤信号时就会悄悄漏掉样本,回测结论可能偏乐观。
class="num">5 DayofWeek class="num">6999 non-null object class="num">6 DayofYear class="num">6999 non-null object class="num">7 Month class="num">6999 non-null object dtypes: float64(class="num">4), object(class="num">4) memory usage: class="num">492.1+ KB
「CatBoost 训练前的参数与拟合落地」
在跑 fit 之前,先把几个直接影响过拟合与算力的参数盯牢。iterations 是树的数量,给到 100 虽不算大,但配合 learning_rate=0.01 这种小步长,模型倾向更稳;depth=10 能抓复杂形态,却也更容易在外汇小时线这类噪声里过拟合。l2_leaf_reg=5 用 L2 惩罚大叶权重,border_count=64 决定分类特征切分精度,数字越高越吃 CPU。 cat_features 最好显式传索引列表,别全靠模型自动识别——贵金属品种标签若被误判为数值,Human-readable 的类别关系就丢了。early_stopping_rounds 注释掉了,实盘前建议打开,验证集多少轮不提升就停,能省时间也压过拟合。 下面这段是把参数塞进 sklearn 管道再 fit 的直跑写法。注意 catboost__eval_set 传了测试集,catboost__cat_features 传类别列名,管道前缀别拼错。 从日志看,第 3 轮 test Logloss 摸到 0.6931239 的 best 后再没被刷新,到第 99 轮 test 已滑到 0.6936898,说明早停若设 10 轮会在第 13 轮附近断掉,白跑 86 轮。外汇与贵金属杠杆高、回测甜区不等于实盘,参数只代表概率倾向。
params = dict( iterations=class="num">100, learning_rate=class="num">0.01, depth=class="num">10, l2_leaf_reg=class="num">5, bagging_temperature=class="num">1, border_count=class="num">64, # Number of splits for categorical features eval_metric=&class="macro">#x27;Logloss&class="macro">#x27;, random_seed=class="num">42, # Seed for reproducibility verbose=class="num">1, # Verbosity level # early_stopping_rounds=class="num">10 # Early stopping for validation ) pipe = Pipeline([ ("catboost", CatBoostClassifier(**params)) ]) # Fit the pipeline to the training data pipe.fit(X_train, y_train, catboost__eval_set=(X_test, y_test), catboost__cat_features=categorical_features) class="num">90: learn: class="num">0.6880592 test: class="num">0.6936112 best: class="num">0.6931239 (class="num">3) total: 523ms remaining: class="num">51.7ms class="num">91: learn: class="num">0.6880397 test: class="num">0.6936100 best: class="num">0.6931239 (class="num">3) total: 529ms remaining: 46ms class="num">92: learn: class="num">0.6880350 test: class="num">0.6936051 best: class="num">0.6931239 (class="num">3) total: 532ms remaining: 40ms class="num">93: learn: class="num">0.6880280 test: class="num">0.6936103 best: class="num">0.6931239 (class="num">3) total: 535ms remaining: class="num">34.1ms class="num">94: learn: class="num">0.6879448 test: class="num">0.6936110 best: class="num">0.6931239 (class="num">3) total: 541ms remaining: class="num">28.5ms class="num">95: learn: class="num">0.6878328 test: class="num">0.6936387 best: class="num">0.6931239 (class="num">3) total: 547ms remaining: class="num">22.8ms class="num">96: learn: class="num">0.6877888 test: class="num">0.6936473 best: class="num">0.6931239 (class="num">3) total: 553ms remaining: class="num">17.1ms class="num">97: learn: class="num">0.6877408 test: class="num">0.6936508 best: class="num">0.6931239 (class="num">3) total: 559ms remaining: class="num">11.4ms class="num">98: learn: class="num">0.6876611 test: class="num">0.6936708 best: class="num">0.6931239 (class="num">3) total: 565ms remaining: class="num">5.71ms class="num">99: learn: class="num">0.6876230 test: class="num">0.6936898 best: class="num">0.6931239 (class="num">3) total: 571ms remaining: 0us
CatBoost 早停后的模型裁剪点
CatBoost 在训练过程中若触发早停,会在日志里直接给出关键数值:某次回测中 bestTest 落在 0.6931239281,对应 bestIteration 为第 3 轮。 系统随后提示 Shrink model to first 4 iterations,意味着最终保留的是前 4 轮迭代得到的模型,而非跑满全部树。 对外汇与贵金属这类高波动品种做集成模型信号时,这种裁剪能压住过拟合,但样本外失效概率仍不低,实盘前务必在 MT5 策略测试器用分时段数据复跑确认。
bestTest = class="num">0.6931239281 bestIteration = class="num">3 Shrink model to first class="num">4 iterations.
◍ CatBoost 分类报告的真实水位
用 Sklearn 的 classification_report 直接拉训练集和测试集的明细,比只看一个准确率数字实在得多。下面这段脚本在管道训练完后分别跑两套预测并打印报告,能立刻看到 precision、recall 和 f1 的分布。 [CODE] # Make predicitons on training and testing sets y_train_pred = pipe.predict(X_train) y_test_pred = pipe.predict(X_test) # Training set evaluation print("Training Set Classification Report:") print(classification_report(y_train, y_train_pred)) # Testing set evaluation print("\nTesting Set Classification Report:") print(classification_report(y_test, y_test_pred)) [/CODE] 训练集整体 accuracy 0.54,测试集掉到 0.51;类别 1 的 recall 在测试集为 0.61,但 precision 仅 0.49,说明模型倾向于把样本判为 1,却伴随不少误报。删掉分类特征列表后,训练集准确率能爬到 60%,测试集却纹丝不动——过拟合的信号很明显。 再把 pipe 用 eval_set 挂上测试集跑 CatBoost,第 30 轮附近 best test 对数损失约 0.69305,之后 learn 继续下降到 0.684 左右而 test 徘徊在 0.6933,验证集收益已经吃尽。特征重要性图板显示,模型决策权重主要压在分类变量而非连续变量上,这点和删特征后训练集提升的现象对得上。外汇与贵金属行情受此类离散状态切换影响大,直接拿该口径建模实盘信号属高风险行为,仅适合作为特征筛选的参考。
# Make predicitons on training and testing sets y_train_pred = pipe.predict(X_train) y_test_pred = pipe.predict(X_test) # Training set evaluation print("Training Set Classification Report:") print(classification_report(y_train, y_train_pred)) # Testing set evaluation print("\nTesting Set Classification Report:") print(classification_report(y_test, y_test_pred))
「训练尾段与特征权重怎么读」
上面这段是 CatBoost 训练跑到第 95~99 轮时的日志切片。learn 损失从 0.6841427 缓降到 0.6838397,test 损失反而从 0.6933758 爬到 0.6934259,说明过拟合倾向在尾段已经显现;best test 停在 0.6930499562,对应 bestIteration = 30,模型随后被裁剪到前 31 轮。 训练集分类报告里,类别 0 的 precision 0.61、recall 0.53,类别 1 的 precision 0.59、recall 0.67,总 accuracy 0.60(support 6999)。这个精度在外汇或贵金属行情分类上只能算弱信号,实盘直接跟单风险很高,仅适合做过滤层。 日志之后是从 pipeline 里抠出 CatBoost 模型、拉特征重要度的代码。跑完能看到哪些字段对分类贡献大,进而决定下次训练砍掉哪些冗余特征、缩短迭代。
# Extract the trained CatBoostClassifier from the pipeline catboost_model = pipe.named_steps[&class="macro">#x27;catboost&class="macro">#x27;] # Get feature importances feature_importances = catboost_model.get_feature_importance() feature_im_df = pd.DataFrame({ "feature": X.columns, "importance": feature_importances }) feature_im_df = feature_im_df.sort_values(by="importance", ascending=False) plt.figure(figsize=(class="num">10, class="num">6)) sns.barplot(data = feature_im_df, x=&class="macro">#x27;importance&class="macro">#x27;, y=&class="macro">#x27;feature&class="macro">#x27;, palette="viridis") plt.title("CatBoost feature importance") plt.xlabel("Importance") plt.ylabel("feature") plt.show()