过滤含不同数据类型值的字典遇ValueError,求解决方案
问题:过滤含数组的字典时触发ValueError错误
我尝试过滤包含不同数据类型值的字典,想要移除与'YALE'对应的记录,但运行代码时触发如下错误:
ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()
我的代码如下:
dataset = { 'timeseires': array([[ [ -5.653222, 7.39066 , 20.651941, 4.07861 ,-11.752331, -34.611312], [ -5.653222, 7.39066 , 20.651941, 4.07861 ,-11.752331, -34.611312] ]]), 'site': array(['YALE', 'KKI'], dtype='<U8') } dataset = data.tolist() def filter(pairs): key, value = pairs filter_key = 'site' if key == filter_key and value == 'YALE': return True else: return False final_dic = dict(filter(filter, dataset.items())) print(final_dic)
预期输出:
dataset = { 'timeseires': array([[ [ -5.653222, 7.39066 , 20.651941, 4.07861 ,-11.752331, -34.611312] ]]), 'site': array(['KKI'], dtype='<U8') }
问题分析
- 代码存在未定义变量:
data.tolist()中的data未声明,属于笔误,即使修正为dataset.tolist(),转成列表后也无法实现你的过滤需求。 - 核心逻辑错误:你当前的代码试图过滤字典的键值对,但实际需求是过滤字典内数组的元素。当
value是numpy数组时,value == 'YALE'会返回布尔数组,if判断无法直接解析数组的真值,因此触发报错。
正确实现代码
要达成预期输出,需要针对数组的索引构建过滤掩码,同步过滤timeseires和site数组:
import numpy as np dataset = { 'timeseires': np.array([[ [ -5.653222, 7.39066 , 20.651941, 4.07861 ,-11.752331, -34.611312], [ -5.653222, 7.39066 , 20.651941, 4.07861 ,-11.752331, -34.611312] ]]), 'site': np.array(['YALE', 'KKI'], dtype='<U8') } # 构建掩码:标记site中不等于'YALE'的元素位置 mask = dataset['site'] != 'YALE' # 过滤timeseires:对应第二个维度(样本维度)保留符合条件的元素 dataset['timeseires'] = dataset['timeseires'][:, mask, :] # 过滤site数组 dataset['site'] = dataset['site'][mask] print(dataset)
代码说明
- 布尔掩码
mask会生成[False, True],对应site数组中'YALE'(需要排除)和'KKI'(需要保留)的位置 timeseires数组的维度是(1, 2, 6),使用[:, mask, :]可以精准保留第二个维度中符合掩码条件的样本- 直接用掩码过滤
site数组,最终得到仅包含'KKI'的数组
内容的提问来源于stack exchange,提问作者Konstantin Simonov
相关产品推荐
相关产品推荐

