【问题标题】:Python, scatter plot for all independent variablesPython,所有自变量的散点图
【发布时间】:2020-05-04 19:18:53
【问题描述】:

我想创建一个方法,为我的数据集中的所有自变量生成散点图,但是 我有一个错误,我不知道为什么会出现这种情况

class DataAnalysis():
  def __init__(self, X_train, X_test):
    self.X_train = X_train # Train set
    self.X_test = X_test # Test set

  def multi_scatter(self,x_list, y):
    length = np.ceil(len(x_list)/3).astype(int)
    for x in range(0, length):
      fig, axs = plt.subplots(1,3, figsize = (20,10))
      fig.suptitle('Independent variables correlation with target')
      axs[0,0].scatter(self.X_train[x_list[x]], self.X_train[y])
      axs[0,0].set_title(x_list[x])
      axs[0,1].scatter(self.X_train[x_list[x+1]], self.X_train[y])
      axs[0,1].set_title(x_list[x+1])
      axs[0,2].scatter(self.X_train[x_list[x+2]], self.X_train[y])
      axs[0,2].set_title(x_list[x+2])
      x *= 3
      plt.show()

这是我得到的错误:

IndexError                                Traceback (most recent call last)
<ipython-input-34-e8ff51256833> in <module>()
----> 1 analyser.multi_scatter(x_list=train_columns,y=target)

<ipython-input-30-fb0defddaef8> in multi_scatter(self, x_list, y)
      9       fig, axs = plt.subplots(1,3, figsize = (20,10))
     10       fig.suptitle('Independent variables correlation with target')
---> 11       axs[0,0].scatter(ds_train['ExterQual'], ds_train['SalePrice'])
     12       axs[0,0].set_title(x_list[x])
     13       axs[0,1].scatter(self.X_train[x_list[x+1]], self.X_train[y])

IndexError: too many indices for array

提前感谢您的帮助

【问题讨论】:

  • axs[0]axs[1]axs[2] 而不是 axs[0,0]axs[0,1]axs[0,2]
  • 谢谢,解决了这个错误,但我希望每行有 3 个图,所以不会太长
  • 我不明白...你有一行 3 列。

标签: python dataframe matplotlib


【解决方案1】:

虽然问题显然已在 cmets 中解决,但另一种(更简单)的方法是使用更合适的绘图库。例如:

【讨论】:

  • 我正在考虑这个问题,但我不想绘制所有变量组合的矩阵我想要绘制所有独立变量(x 变量列表)和目标变量(y)的图。所以我的想法是创建一个循环,在每次迭代中在一行中绘制三个图,直到我的变量结束
  • @PawełMagdański 你可以告诉每个人要针对什么进行绘图。
【解决方案2】:

好的,我知道了,谢谢你的建议。还是有一些标签问题,但我会解决的

这是一个代码:

class DataAnalysis():
  def __init__(self, X_train, X_test):
    self.X_train = X_train # Train set
    self.X_test = X_test # Test set

  def multi_scatter(self,x_list, y):
    sns.set(style='whitegrid', rc={"grid.linewidth": 0.2})
    sns.set_context("paper", font_scale=2)  
    for x in range(0, len(x_list)):
      if x == 0 or x % 3:
        chart = sns.pairplot(data=self.X_train,
        y_vars=[y],
        x_vars=[x_list[x], x_list[x+1], x_list[x+2]],
        height = 10)
        plt.xticks(rotation = 45)
        plt.show()
      else:
        continue

【讨论】:

    猜你喜欢
    • 2023-03-24
    • 1970-01-01
    • 2020-06-27
    • 2021-11-12
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多