You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何创建JSON配置文件优化RandomizedSearchCV参数?附代码片段

用JSON配置文件管理RandomizedSearchCV的参数优化

我来帮你把参数配置迁移到JSON文件里,这样后续调整参数范围会更灵活,不用每次改代码。下面是完整的实现步骤:


第一步:创建JSON配置文件

先新建一个rf_params_config.json文件,把你要优化的参数范围用JSON格式写进去。注意scipy的分布类型我们用字符串标记,后面在代码里再映射成对应的分布对象:

{
  "n_estimators": {"type": "randint", "low": 100, "high": 1000},
  "max_depth": {"type": "randint", "low": 3, "high": 20},
  "min_samples_split": {"type": "randint", "low": 2, "high": 10},
  "min_samples_leaf": {"type": "randint", "low": 1, "high": 5},
  "max_features": {"type": "choice", "values": ["sqrt", "log2", None]}
}

第二步:编写Python代码加载配置并运行RandomizedSearchCV

接下来修改你的代码,读取JSON配置,把配置里的参数转换成RandomizedSearchCV需要的格式,然后执行搜索:

import numpy as np
import json
from time import time
from scipy.stats import randint as sp_randint
from sklearn.model_selection import RandomizedSearchCV
from sklearn.datasets import load_digits
from sklearn.ensemble import RandomForestClassifier

# 加载JSON参数配置的工具函数
def load_param_config(config_path):
    with open(config_path, 'r') as f:
        config = json.load(f)
    
    param_dist = {}
    for param, settings in config.items():
        if settings['type'] == 'randint':
            param_dist[param] = sp_randint(settings['low'], settings['high'])
        elif settings['type'] == 'choice':
            param_dist[param] = settings['values']
        # 后续要加其他分布(比如uniform、norm)的话,在这里加对应的映射逻辑就行
    return param_dist

# 获取数据集
digits = load_digits()
X, y = digits.data, digits.target

# 初始化分类器
clf = RandomForestClassifier(random_state=42)

# 加载参数分布
param_dist = load_param_config('rf_params_config.json')

# 配置RandomizedSearchCV
random_search = RandomizedSearchCV(
    clf,
    param_distributions=param_dist,
    n_iter=50,  # 迭代次数,按需调整:次数越多越可能找到优参,但耗时更长
    cv=5,       # 交叉验证折数
    verbose=2,
    random_state=42,
    n_jobs=-1   # 用全部CPU核心加速
)

# 运行参数搜索
start = time()
random_search.fit(X, y)
print(f"RandomizedSearchCV 耗时 {time() - start:.2f} 秒")

# 输出最优结果
print("\n找到的最佳参数:")
print(random_search.best_params_)
print(f"\n最佳交叉验证得分:{random_search.best_score_:.4f}")

实用小提示

  • JSON里的参数名必须和RandomForestClassifier的参数完全对应,不然会抛出参数不匹配的错误
  • 如果需要测试连续型参数(比如学习率),可以在JSON里加"type": "uniform",然后在工具函数里映射scipy.stats.uniform
  • 可以根据机器性能调整n_iter和n_jobs,平衡搜索效果和耗时

内容的提问来源于stack exchange,提问作者gabboshow

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 08:26:27