【问题标题】:How to reproduce a RandomForestClassifier when random_state uses RandomState当 random_state 使用 RandomState 时如何重现 RandomForestClassifier
【发布时间】:2021-12-30 18:07:58
【问题描述】:

上下文:

我正在阅读 scikit-learn 的 Common Pitfalls 文档。我对控制随机性部分感到惊讶,因为在我的情况下,我一直使用random_state 和整数来读取估计器应该在分类器的初始化中使用np.random.RandomState 的实例,而不是在 cv分裂;在阅读了交叉验证结果的稳健性部分后,我想:“好的,我明白了”,但是,克隆部分中有一条警告说:

from sklearn import clone
from sklearn.ensemble import RandomForestClassifier
import numpy as np

rng = np.random.RandomState(0)
a = RandomForestClassifier(random_state=rng)
b = clone(a)

由于将 RandomState 实例传递给 a,a 和 b 不是严格意义上的克隆,而是统计意义上的克隆:a 和 b 仍然是不同的模型,即使调用 fit(X, y) on相同的数据。此外,a 和 b 会相互影响,因为它们共享相同的内部 RNG:调用 a.fit 将消耗 b 的 RNG,调用 b.fit 将消耗 a 的 RNG,因为它们是相同的。对于共享 random_state 参数的任何估计器,该位都是正确的;它并不特定于克隆。

如果传递一个整数,a 和 b 将是精确的克隆,它们不会相互影响 警告尽管 clone 很少在用户代码中使用,但它在 scikit-learn 代码库中被广泛调用:特别是,大多数接受非拟合估计器的元估计器在内部调用 clone(GridSearchCV、StackingClassifier、CalibratedClassifierCV 等)。

问题

如果我处于项目的开发阶段并试图确定哪种模型最适合使用 GridSearchCV,我如何才能获得模型的 random_state 值,以便在生产中使用它?

起初我认为让我们在网格中使用random_state,但交叉验证结果的鲁棒性部分写道:

传递实例会导致更稳健的 CV 结果,并使各种算法之间的比较更公平。它还有助于限制将估计器的 RNG 视为可以调整的超参数的诱惑。

【问题讨论】:

    标签: python scikit-learn cross-validation


    【解决方案1】:

    random_state 参数真正用于确定性可重复性,在开发更复杂的管道时特别有用。

    GridSearchCV 用于为您的学习过程寻找最佳设置。我强调程序,因为交叉验证方面是根据估计而不是实际的具体模型来获得统计结果。与机器学习中的许多技术一样,您的 RandomForest 分类器和交叉验证依赖熵/随机性来相当近似事物。 random_state 应该被视为超参数,否则你是指标在噪声上攀升。

    一旦您知道了在统计上产生超越随机机会的良好模型的最佳设置,您就想重新应用该过程来推导出您的生产模型。请注意,此模型的度量性能将限制在(但不等同于)网格搜索估计的范围内。为生产模型指定 random_state 是错误的。

    这是一个有效的方法来处理/不指定随机种子:

    # define your procedure as you wish. random_state is optional
    # good for testing, irrelevant for predicting.
    my_pipeline = RandomForestClassifier(random_state=42, ..)
    
    # search parameter space & then refit another model with your 
    # procedure on the discovered parameters.
    optimal = GridSearchCV(my_pipeline, params, refit=True, ...)
    optimal.fit(train_X, train_y)
    
    # get the new model trained with the best found parameters,
    # rather than best performing model of the cross-validation!!
    prod_model = optimal.best_estimator_  
    

    【讨论】:

    • 我刚看了你的解释,感觉很傻。谢谢@eliangius。
    【解决方案2】:

    警告源于使用随机数生成器对象而不是随机种子整数。

    当您使用生成器时,对它的后续调用将按随机顺序生成不同的种子。这意味着由于克隆共享相同的生成器序列,因此操作顺序、每个被调用的速率等都是隐式耦合的,因此会产生不同但在统计上等效的对象。使用随机整数种子是安全的。

    【讨论】:

      猜你喜欢
      • 2016-10-03
      • 2015-04-11
      • 2017-09-24
      • 2019-05-12
      • 1970-01-01
      • 2018-05-06
      • 2020-08-29
      • 1970-01-01
      • 2021-02-09
      相关资源
      最近更新 更多