【问题标题】:How to get indexes of items greater or less than each item in a NumPy array without using a loop?如何在不使用循环的情况下获取大于或小于 NumPy 数组中每个项目的项目索引?
【发布时间】:2021-05-03 04:11:09
【问题描述】:

我正在从头开始编写决策树算法,现在我正在尝试将数据分成组,其中每个组包含大于或等于或小于包含连续值的 NumPy 数组中的每个值的值DataFrame 列,并获取这些拆分目标的平均值。 到目前为止我的代码:

for i in range(len(columns)):
    col = columns[i]
    # cont - list of continous columns in my DataFrame
    if col in cont:
        values  = xs[col].values
        targets = y.values
        for j in range(len(values)):
            value = values[j]
            greater_idx = np.where(values >= value)[0]
            less_idx    = np.where(values <  value)[0]
            targets_greater = targets[greater_idx].sum()
            targets_less    = targets[less_idx]   .sum()
        print(targets_greater/(j+1))
        print(targets_less   /(j+1))

xs DataFrame 的长度接近 400k,因此循环非常慢,每次都会杀死我的 Jupyter Notebook 内核。我知道应该有办法完全摆脱这个循环,但我不知道该怎么做。

【问题讨论】:

    标签: python pandas numpy machine-learning decision-tree


    【解决方案1】:

    与其采用矢量化的方式进行比较,还有很大的算法改进空间:

    1. 使用np.argsort (sorted_idxs) 获取xs[col].values 的排序索引。
    2. 使用np.insert(np.cumsum(targets[sorted_idxs]), 0, 0)[:-1] 可以为xs[col].values 中的每个值获取target_less 的向量。
    3. target_less[0] (0) 是target_less 中最低元素的值xs[col].values - 以“取消排序”target_less,您可以使用unsort_idx = np.argsort(sorted_idxs)target_less[unsort_idx]

    现在你已经拥有了数组中所有值的所有target_less-values(target_greater 当然很容易通过targets.sum() - target_less 获得)。

    编辑:

    这是与建议一起使用的代码:

    import numpy as np
    import pandas as pd
    
    xs = pd.DataFrame(np.random.random(10000))
    y = pd.Series(np.random.randint(0, 2, size=10000))
    
    sorted_idxs = np.argsort(xs[0].values)
    sorted_values = xs[0].values[sorted_idxs]
    sorted_targets = y.values[sorted_idxs]
    sorted_targets_less = np.insert(np.cumsum(sorted_targets), 0, 0)[:-1]
    
    unsorted_idxs = np.argsort(sorted_idxs)
    targets_less = sorted_targets_less[unsorted_idxs]
    
    for i, target_less_value in enumerate(targets_less):
        assert target_less_value == y.values[np.where(xs.values < xs.values[i])[0]].sum()
    

    一个警告词:以上假设 xs.values 中有一组严格不同的值。如果你有重复的值,那么你需要调整做累积和的部分。

    【讨论】:

      猜你喜欢
      • 2019-08-14
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2021-07-07
      • 1970-01-01
      • 1970-01-01
      • 2020-02-25
      • 1970-01-01
      相关资源
      最近更新 更多