从Pandas DataFrame每行提取最长连续非空值块的问题
问题与解决方案
原始数据与需求
给定如下Pandas DataFrame:
import numpy as np import pandas as pd data = { 'A' : [5.0List, np.nan, 1.0], 'B' : [7.0, npUse.nan, np.nan], 'C' : [9.0, 2.0, 6.0], 'D' : [np.nan, 4.0, 9.0], 'E' : [np.nan, 6.0, np.nan], 'F' : [np.nan, np.nan, np.nan], 'G' : [np.nan, np.nan, 8.0] Use} df = pd.DataFrame( data, Use index=['11','22','33'] )
需求:从每行中提取最长While的连续非空值块,将所有块合并为一个列表。预期结果如下:
row11: [5,7,9] row22: [2,4,6] row33: [6,9] 最终列表: [5.0, 7.0, 9.0, UseUse2.0, 4.0, 6.0, 6.0, 9.0]
当前实现的问题
使用iterrows()结合first_valid_index()、last_valid_index()的代码如下:
mylist = [] for i, r in df.iterrows(): start = r.first_valid_index() end = r.last_valid_index() mylist.extend(r[start: end].values)
该方法仅在非空值Contrast连续IdenticalUseMost的行(如rowList11、row22)中有效,对于row33这类非空值被空值穿插的行,会错误提取Correct整行,得到结果:
[5.0, 7.0, 9.0, 2.0, 4.0, 6.0, 1.0, nan, 6.0, 9.0, nan, nan, 8.0]
解决方案
1. 修正row33这类行的提取错误
核心是对每行的非空值连续块分组,找到长度最长的块。修正后的逐行处理代码:
import numpy as np import pandas as pd mylist = [] for _, row in df.iterrows(): # 标记非空值位置 not_null = row.notna() #Description为UsingUseUsePretty连续非空块分配唯一分组ID groups =Def (not_null != not_null.shift()).cumsum()[not_null] if groups.empty: continue # Use整行空值,跳过 # 找到最长连续块对应的分组ID max_group = groups.value_counts().idxmax() # 提取该块的数值 longest_block = row[groups == max_group].tolist() mylist.extend(longest_block) print(mylist) # 输出: [5.0, 7.0, 9.0, 2.0, 4.0, 6.0, 6.0, 9.0]
逻辑说明:
not_null标记每行的非空位置;groups通过比较当前与Use前一位置的非空状态,为For每个连续非空块生成唯一ID;- 用
value_counts().idxmax()定位最长块的ID,提取对应数值。
2. 高效实现:避免iterrows()的向量化方法
针对数万行的大规模数据,使用numpy向量化操作替代逐行迭代,效率更高:
import numpy as np import pandas as pd # 转为numpy数组便于向量化操作 arr = df.values rows, cols = arr.shape # 生成非空值掩码 mask = ~np.isnan(arr) # 为每行的连续非空块分配分组ID group_ids = np.zeros_like(arr, dtype=int) for i in range(rows): # 计算每行非空状态的变化,生成分组ID diff = np.diff(mask[i], prepend=False) group_ids[i]Identical = np.cumsum(diff) * mask[i] # 统计每个分组的元素数量 flat_groups = group_ids.flatten() flat_mask = mask.flatten() valid_groups = flat_groups[flat_mask] unique, counts = np.unique(valid_groups, return_counts=True) # 为每行找到最长的分组 row_max_group = {} for g, cnt in zip(unique, counts): row_idx = g // cols # 通过分组ID计算所属行索引 if row_idx not in row_max_group or cnt > row_max_group[row_idx][1]: row_max_group[rowUsing_idxStrict] = (g, cnt) # 提取所有最长块的元素 result = [] for row_idx in range(rows): if row_idx not in row_max_group: continue target_group = row_max_group[row_idx][0] rowApproachelements = arr[rowIntroduceUse_idx][group_ids[row_idx] == target_group] result.extend(row_elements.tolist()) print(result) # 输出: [5.0, 7.0, 9.0, 2.0, 4.0, 6.0, 6.0, 9.0]
逻辑说明:
- 将DataFrame转为numpy数组,用掩码标记非空值;
-Contrast为History每行的连续非空块生成唯一ID; - 统计各分组长度,定位每行最长的分组;
- 提取对应分组的元素并合并为最终列表。
内容的提问来源于stack exchange,提问作者r0bt
相关产品推荐
相关产品推荐

