【问题标题】:Predict() with sklearn LogisticRegression and RandomForest models always predict minority class (1)Predict() 与 sklearn LogisticRegression 和 RandomForest 模型总是预测少数类 (1)
【发布时间】:2019-01-28 19:43:40
【问题描述】:

我正在构建一个逻辑回归模型,以仅使用 150 个观察值的数据集来预测交易是否有效 (1) 或无效 (0)。我的数据在两个类之间分布如下:

  • 106 个观察结果为 0(无效)
  • 44 个观察结果为 1(有效)

我正在使用两个预测变量(均为数字)。尽管数据大多为 0,但我的分类器仅预测测试集中每个事务的 1,即使它们中的大多数应该为 0。分类器从不为任何观察输出 0。

这是我的全部代码:

# Logistic Regression
import numpy as np
import pandas as pd
from pandas import Series, DataFrame

import scipy
from scipy.stats import spearmanr
from pylab import rcParams
import seaborn as sb
import matplotlib.pyplot as plt
import sklearn
from sklearn.preprocessing import scale
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split
from sklearn import metrics
from sklearn import preprocessing

address = "dummy_csv-150.csv"
trades = pd.read_csv(address)
trades.columns=['location','app','el','rp','rule1','rule2','rule3','validity','transactions']
trades.head()

trade_data = trades.ix[:,(1,8)].values
trade_data_names = ['app','transactions']

# set dependent/response variable
y = trades.ix[:,7].values

# center around the data mean
X= scale(trade_data)

LogReg = LogisticRegression()

LogReg.fit(X,y)
print(LogReg.score(X,y))

y_pred = LogReg.predict(X)

from sklearn.metrics import classification_report

print(classification_report(y,y_pred)) 

log_prediction = LogReg.predict_log_proba(
    [
       [2, 14],[3,1], [1, 503],[1, 122],[1, 101],[1, 610],[1, 2120],[3, 85],[3, 91],[2, 167],[2, 553],[2, 144]
    ])
prediction = LogReg.predict([[2, 14],[3,1], [1, 503],[1, 122],[1, 101],[1, 610],[1, 2120],[3, 85],[3, 91],[2, 167],[2, 553],[2, 144]])

我的模型定义为:

LogReg = LogisticRegression()  
LogReg.fit(X,y)

X 看起来像这样:

X = array([[1, 345],
       [1, 222],
       [1, 500],
       [2, 120]]....)

对于每个观察,Y 只是 0 或 1。

归一化 传递给模型的 X 是这样的:

[[-1.67177659  0.14396503]
 [-1.67177659 -0.14538932]
 [-1.67177659  0.50859856]
 [-1.67177659 -0.3853417 ]
 [-1.67177659 -0.43239119]
 [-1.67177659  0.743846  ]
 [-1.67177659  4.32195953]
 [ 0.95657805 -0.46062089]
 [ 0.95657805 -0.45591594]
 [ 0.95657805 -0.37828428]
 [ 0.95657805 -0.52884264]
 [ 0.95657805 -0.20420118]
 [ 0.95657805 -0.63705646]
 [ 0.95657805 -0.65587626]
 [ 0.95657805 -0.66763863]
 [-0.35759927 -0.25125067]
 [-0.35759927  0.60975496]
 [-0.35759927 -0.33358727]
 [-0.35759927 -0.20420118]
 [-0.35759927  1.37195666]
 [-0.35759927  0.27805607]
 [-0.35759927  0.09456307]
 [-0.35759927  0.03810368]
 [-0.35759927 -0.41121892]
 [-0.35759927 -0.64411389]
 [-0.35759927 -0.69586832]
 [ 0.95657805 -0.57353966]
 [ 0.95657805 -0.57353966]
 [ 0.95657805 -0.53825254]
 [ 0.95657805 -0.53354759]
 [ 0.95657805 -0.52413769]
 [ 0.95657805 -0.57589213]
 [ 0.95657805  0.03810368]
 [ 0.95657805 -0.66293368]
 [ 0.95657805  2.86107294]
 [-1.67177659  0.14396503]
 [-1.67177659 -0.14538932]
 [-1.67177659  0.50859856]
 [-1.67177659 -0.3853417 ]
 [-1.67177659 -0.43239119]
 [-1.67177659  0.743846  ]
 [-1.67177659  4.32195953]
 [ 0.95657805 -0.46062089]
 [ 0.95657805 -0.45591594]
 [ 0.95657805 -0.37828428]
 [ 0.95657805 -0.52884264]
 [ 0.95657805 -0.20420118]
 [ 0.95657805 -0.63705646]
 [ 0.95657805 -0.65587626]
 [ 0.95657805 -0.66763863]
 [-0.35759927 -0.25125067]
 [-0.35759927  0.60975496]
 [-0.35759927 -0.33358727]
 [-0.35759927 -0.20420118]
 [-0.35759927  1.37195666]
 [-0.35759927  0.27805607]
 [-0.35759927  0.09456307]
 [-0.35759927  0.03810368]
 [-0.35759927 -0.41121892]
 [-0.35759927 -0.64411389]
 [-0.35759927 -0.69586832]
 [ 0.95657805 -0.57353966]
 [ 0.95657805 -0.57353966]
 [ 0.95657805 -0.53825254]
 [ 0.95657805 -0.53354759]
 [ 0.95657805 -0.52413769]
 [ 0.95657805 -0.57589213]
 [ 0.95657805  0.03810368]
 [ 0.95657805 -0.66293368]
 [ 0.95657805  2.86107294]
 [-1.67177659  0.14396503]
 [-1.67177659 -0.14538932]
 [-1.67177659  0.50859856]
 [-1.67177659 -0.3853417 ]
 [-1.67177659 -0.43239119]
 [-1.67177659  0.743846  ]
 [-1.67177659  4.32195953]
 [ 0.95657805 -0.46062089]
 [ 0.95657805 -0.45591594]
 [ 0.95657805 -0.37828428]
 [ 0.95657805 -0.52884264]
 [ 0.95657805 -0.20420118]
 [ 0.95657805 -0.63705646]
 [ 0.95657805 -0.65587626]
 [ 0.95657805 -0.66763863]
 [-0.35759927 -0.25125067]
 [-0.35759927  0.60975496]
 [-0.35759927 -0.33358727]
 [-0.35759927 -0.20420118]
 [-0.35759927  1.37195666]
 [-0.35759927  0.27805607]
 [-0.35759927  0.09456307]
 [-0.35759927  0.03810368]
 [-0.35759927 -0.41121892]
 [-0.35759927 -0.64411389]
 [-0.35759927 -0.69586832]
 [ 0.95657805 -0.57353966]
 [ 0.95657805 -0.57353966]
 [ 0.95657805 -0.53825254]
 [ 0.95657805 -0.53354759]
 [ 0.95657805 -0.52413769]
 [ 0.95657805 -0.57589213]
 [ 0.95657805  0.03810368]
 [ 0.95657805 -0.66293368]
 [ 0.95657805  2.86107294]
 [-1.67177659  0.14396503]
 [-1.67177659 -0.14538932]
 [-1.67177659  0.50859856]
 [-1.67177659 -0.3853417 ]
 [-1.67177659 -0.43239119]
 [-1.67177659  0.743846  ]
 [-1.67177659  4.32195953]
 [ 0.95657805 -0.46062089]
 [ 0.95657805 -0.45591594]
 [ 0.95657805 -0.37828428]
 [ 0.95657805 -0.52884264]
 [ 0.95657805 -0.20420118]
 [ 0.95657805 -0.63705646]
 [ 0.95657805 -0.65587626]
 [ 0.95657805 -0.66763863]
 [-0.35759927 -0.25125067]
 [-0.35759927  0.60975496]
 [-0.35759927 -0.33358727]
 [-0.35759927 -0.20420118]
 [-0.35759927  1.37195666]
 [-0.35759927  0.27805607]
 [-0.35759927  0.09456307]
 [-0.35759927  0.03810368]
 [-0.35759927 -0.41121892]
 [-0.35759927 -0.64411389]
 [-0.35759927 -0.69586832]
 [ 0.95657805 -0.57353966]
 [ 0.95657805 -0.57353966]
 [ 0.95657805 -0.53825254]
 [ 0.95657805 -0.53354759]
 [ 0.95657805 -0.52413769]
 [ 0.95657805 -0.57589213]
 [ 0.95657805  0.03810368]
 [ 0.95657805 -0.66293368]
 [ 0.95657805  2.86107294]
 [-0.35759927  0.60975496]
 [-0.35759927 -0.33358727]
 [-0.35759927 -0.20420118]
 [-0.35759927  1.37195666]
 [-0.35759927  0.27805607]
 [-0.35759927  0.09456307]
 [-0.35759927  0.03810368]]

Y 是:

[0 0 0 0 0 0 1 1 0 0 0 1 1 1 1 0 0 0 0 1 0 0 0 0 1 1 0 0 0 0 0 0 1 1 1 0 0
 0 0 0 0 1 1 0 0 0 1 1 1 1 0 0 0 0 1 0 0 0 0 1 1 0 0 0 0 0 0 1 1 1 0 0 0 0
 0 0 1 1 0 0 0 1 1 1 1 0 0 0 0 1 0 0 0 0 1 1 0 0 0 0 0 0 1 1 1 0 0 0 0 0 0
 1 1 0 0 0 1 1 1 1 0 0 0 0 1 0 0 0 0 1 1 0 0 0 0 0 0 1 1 1 0 0 0 1 0 0 0]

模型指标是:

             precision    recall  f1-score   support

          0       0.78      1.00      0.88        98
          1       1.00      0.43      0.60        49

avg / total       0.85      0.81      0.78       147

得分为 0.80

当我运行 model.predict_log_proba(test_data) 时,我得到如下所示的概率区间:

array([[ -1.10164032e+01,  -1.64301095e-05],
       [ -2.06326947e+00,  -1.35863187e-01],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00],
       [            -inf,   0.00000000e+00]])

我的测试集是,除了 2 之外的所有测试集都应该为 0,但它们都被归类为 1。这发生在每个测试集上,即使是那些具有模型训练值的测试集。

[2, 14],[3,1], [1, 503],[1, 122],[1, 101],[1, 610],[1, 2120],[3, 85],[3, 91],[2, 167],[2, 553],[2, 144]

我在这里发现了一个类似的问题:https://stats.stackexchange.com/questions/168929/logistic-regression-is-predicting-all-1-and-no-0 但在这个问题中,问题似乎是数据大多是 1,所以模型输出 1 是有道理的。我的情况正好相反,因为火车数据大多是 0,但由于某种原因,我的模型总是为所有内容输出 1,即使 1 相对较少。我还尝试了一个随机森林分类器来查看模型是否错误,但同样的事情发生了。也许这是我的数据,但我不知道它有什么问题,因为它符合所有假设。

可能出了什么问题?数据满足逻辑模型的所有假设(两个预测变量是独立的,输出是二进制的,没有丢失数据点)。任何建议表示赞赏。

【问题讨论】:

  • 我怀疑您的代码中存在错误,而不是统计问题。您的测试集和训练集中的标签计数是多少?
  • 您发布的模型指标与您所说的相反。根据那里的precision和recall值,可以观察到模型已经预测了至少30个条目为0,剩下的4或5个为1。那么这些指标是在训练数据上计算的还是在测试上计算的?
  • 您可以在这里发布数据和完整代码吗?
  • @Denziloe,谢谢。我的测试集中的标签总数为 12。 2 应该是 1,10 应该是 0,但分类器将它们都标记为 1。训练集总共有 150 个观测值,其中 106 个为 0,44 个为 1。训练将其中 98 个标记为 0,将 49 个标记为 1,这是合理的。当我对新数据运行 predict() 时,它只输出 1s...
  • @VivekKumar 谢谢,我已经更新了我的问题以包含代码和数据。

标签: python machine-learning scikit-learn logistic-regression


【解决方案1】:

您不会缩放您的test data。当您这样做时,您将正确缩放您的列车数据:

X= scale(trade_data)

培训您的型号后,您不与测试数据执行相同:

log_prediction = LogReg.predict_log_proba(
[
   [2, 14],[3,1], [1, 503],[1, 122],[1, 101],[1, 610],[1, 2120],[3, 85],[3, 91],[2, 167],[2, 553],[2, 144]
])

模型的系数是预期的标准化输入。您的测试数据未标准化。模型的任何正系数将乘以大量数字,因为您的数据不缩放可能将预测值驱动到全部为1。

一般规则是您在训练集上所做的任何转换都应该在您的测试集上完成。您还应该使用测试集在培训集上应用相同的转换。而不是:

X = scale(trade_data)

您应该从培训数据中创建一个缩放器:

scaler = StandardScaler().fit(trade_date)
X = scaler.transform(trade_data)
然后稍后将调光器应用于您的test data:
scaled_test = scaler.transform(test_x)

【讨论】:

  • aaah我看到了!这是完全有意义的。最后,我得到了测试集有意义的结果,大多数包括0和少数1s!非常感谢!
猜你喜欢
  • 2022-01-05
  • 1970-01-01
  • 2020-06-26
  • 2022-06-28
  • 1970-01-01
  • 2021-02-27
  • 2018-06-04
  • 2021-04-16
  • 1970-01-01
相关资源
最近更新 更多