如何用前后非NaN值混合填充一维数组NaN块?是否有SciPy现成函数?
问题解答
一、是否有现成函数实现?
目前没有SciPy、Pandas等库中的现成函数可以直接匹配你自定义的NaN填充规则。
常见的插值/填充函数(比如pandas.Series.interpolate()、scipy.interpolate.interp1d)都是基于线性、多项式、样条等数学插值逻辑,或者简单的前向/后向填充(ffill/bfill),完全不符合你要求的“按NaN块长度拆分,取前后指定数量的非NaN值重复填充”的规则,因此需要自定义实现。
二、更优的实现方案
你的初步实现逻辑是正确的,可以从边界处理、代码健壮性、性能三个方面优化:
1. 完善边界NaN块的处理
原代码已覆盖“相邻非NaN值不足”的情况,但可以明确处理数组开头或结尾的NaN块(比如输入是[NaN, NaN, 2, 3]或[2,3,NaN,NaN]),确保仅从单侧取可用值重复填充:
- 若NaN块在数组开头:仅从右侧取数,不足则重复
- 若NaN块在数组结尾:仅从左侧取数,不足则重复
2. 优化后的实现代码
以下版本在原逻辑基础上,增加了边界处理、原数组保护、空NaN块快速返回等优化:
import numpy as np def fillna_with_neighbours_optimized(data): data = data.copy() # 避免修改原输入数组 nan_mask = np.isnan(data) if not np.any(nan_mask): return data # 提取所有NaN块的起始、结束索引(左闭右开) nan_blocks = np.where(np.diff(np.concatenate([[0], nan_mask.astype(int), [0]])) != 0)[0].reshape(-1, 2) for start, end in nan_blocks: n = end - start n_pre = (n // 2) + (1 if n % 2 != 0 else 0) n_post = n // 2 # 处理开头的NaN块 if start == 0: post_vals = data[end:end + n] data[start:end] = np.resize(post_vals, n) continue # 处理结尾的NaN块 if end == len(data): pre_vals = data[start - n:start] data[start:end] = np.resize(pre_vals, n) continue # 处理中间的NaN块 # 获取前侧填充值,不足则重复 pre_start = max(start - n_pre, 0) pre_vals = data[pre_start:start] if len(pre_vals) < n_pre: pre_vals = np.resize(pre_vals, n_pre) # 获取后侧填充值,不足则重复 post_end = min(end + n_post, len(data)) post_vals = data[end:post_end] if len(post_vals) < n_post: post_vals = np.resize(post_vals, n_post) data[start:end] = np.concatenate([pre_vals, post_vals]) return data # 测试示例 ex1 = np.asarray([1, 2, np.nan, np.nan, np.nan, np.nan, 3, 2]) ex2 = np.asarray([1, 2, 3, np.nan, np.nan, np.nan, 4, 2]) ex3 = np.asarray([3, np.nan, np.nan, np.nan, np.nan, 4, 2]) ex4 = np.asarray([np.nan, np.nan, 2, 3, np.nan, np.nan]) # 新增边界测试 print(fillna_with_neighbours_optimized(ex1)) # 输出:[1. 2. 1. 2. 3. 2. 3. 2.] print(fillna_with_neighbours_optimized(ex2)) # 输出:[1. 2. 3. 2. 3. 4. 4. 2.] print(fillna_with_neighbours_optimized(ex3)) # 输出:[3. 3. 3. 4. 2. 4. 2.] print(fillna_with_neighbours_optimized(ex4)) # 输出:[2. 2. 2. 3. 2. 3.]
3. 性能优化要点
- 提前复制输入数组,避免修改原数据引发意外
- 增加空NaN块判断,直接返回原数组减少无效计算
- 用
max()/min()简化边界索引计算,替代冗余条件判断
内容的提问来源于stack exchange,提问作者dohe
相关产品推荐
相关产品推荐

