Python中基于指定列值移除数组重复行的实现方法
解决数组去重并保留第三列最大值的Pythonic实现
方法一:使用Pandas(简洁直观)
Pandas的groupby操作是向量化实现,无需显式循环,非常适合这类分组取极值的场景:
import pandas as pd import numpy as np # 定义原始数组 array = np.array([ [0.91, 0.33, 0.09], [0.52, 0.63, 0.05], [0.91, 0.33, 0.11], [0.52, 0.63, 0.07], [0.62, 0.41, 0.01], [0.36, 0.37, 0.01] ]) # 转换为DataFrame方便分组操作 df = pd.DataFrame(array, columns=["col1", "col2", "col3"]) # 按前两列分组,保留每组第三列的最大值所在行 result_df = df.groupby(["col1", "col2"], as_index=False).max() # 转回numpy数组 array2 = result_df.to_numpy() print(array2)
输出结果:
[[0.36 0.37 0.01] [0.52 0.63 0.07] [0.62 0.41 0.01] [0.91 0.33 0.11]]
方法二:使用NumPy(纯数值计算)
如果不想引入Pandas,也可以用NumPy的向量化操作实现,避免显式循环:
import numpy as np array = np.array([ [0.91, 0.33, 0.09], [0.52, 0.63, 0.05], [0.91, 0.33, 0.11], [0.52, 0.63, 0.07], [0.62, 0.41, 0.01], [0.36, 0.37, 0.01] ]) # 获取前两列的唯一值及每个行对应的组索引 unique_groups, group_indices = np.unique(array[:, :2], axis=0, return_inverse=True) # 对每个组,找到第三列最大值对应的原始行索引 max_row_indices = [] for group_id in range(len(unique_groups)): # 筛选当前组的所有行 group_mask = group_indices == group_id # 找到第三列最大值的位置 max_idx = np.argmax(array[group_mask, 2]) # 记录原始数组中的索引 max_row_indices.append(np.where(group_mask)[0][max_idx]) # 提取目标行 array2 = array[max_row_indices] print(array2)
输出结果与方法一一致。
内容的提问来源于stack exchange,提问作者Chelsea Zou
相关产品推荐
相关产品推荐

