嵌套字典实现参数覆盖次数统计及值追踪问题
参数追踪字典的实现与问题修复
问题说明
需要实现一个参数追踪字典,每个嵌套键包含参数的最新值,以及每次更新时递增的计数器。现有代码运行时抛出KeyError,目标是得到如下结果:
{'threads': {'count': 2, 'value': '2'}, 'tolerance': {'count': 1, 'value': '1'}, 'timeout': {'count': 1, 'value': True}}
原代码如下:
options = {'threads' : '1', 'tolerance' : '1'} options_2 = {'threads' : '2', 'timeout': True} over_params = {} def overrides_tracker(options, over_params): print(options) for k, v in options.items(): over_params[k]['count'] += 1 over_params[k]['value'] = v return over_params over_params = overrides_tracker(options, over_params) over_params = overrides_tracker(options_2, over_params) print(over_params)
错误原因
第一次处理某个参数键时,over_params字典中还没有该键对应的嵌套字典结构,直接访问over_params[k]['count']会触发KeyError。
基础修复方案
在更新计数器前,先检查键是否存在,不存在则初始化嵌套结构:
options = {'threads' : '1', 'tolerance' : '1'} options_2 = {'threads' : '2', 'timeout': True} over_params = {} def overrides_tracker(options, over_params): for k, v in options.items(): # 初始化新键的嵌套结构 if k not in over_params: over_params[k] = {'count': 0, 'value': None} over_params[k]['count'] += 1 over_params[k]['value'] = v return over_params over_params = overrides_tracker(options, over_params) over_params = overrides_tracker(options_2, over_params) print(over_params)
运行这段代码就能得到预期结果。
更简洁的实现方式
方法1:用字典get方法简化初始化
不用显式判断键是否存在,借助get方法直接获取或初始化嵌套结构:
def overrides_tracker(options, over_params): for k, v in options.items(): # 存在则返回对应字典,不存在则返回初始结构 param = over_params.get(k, {'count': 0, 'value': None}) param['count'] += 1 param['value'] = v over_params[k] = param return over_params
方法2:使用collections.defaultdict自动初始化
借助defaultdict可以自动为新键生成指定的嵌套结构,代码更简洁:
from collections import defaultdict options = {'threads' : '1', 'tolerance' : '1'} options_2 = {'threads' : '2', 'timeout': True} # 每个新键自动生成初始嵌套字典 over_params = defaultdict(lambda: {'count': 0, 'value': None}) def overrides_tracker(options, over_params): for k, v in options.items(): over_params[k]['count'] += 1 over_params[k]['value'] = v return over_params overrides_tracker(options, over_params) overrides_tracker(options_2, over_params) # 可选:转换为普通字典输出 print(dict(over_params))
内容的提问来源于stack exchange,提问作者Chandler Cree
相关产品推荐
相关产品推荐

