迁移Registry类调用代码至Jupyter Notebook遇注册缺失错误
问题
将原本在类模块中运行的代码迁移至Jupyter Notebook测试时,已复制所有相关类,但仍收到“注册表中对象缺失”的错误,提示calculate_psnr未在metric注册表中找到。相关代码及错误信息如下:
init.py 代码
from basic.utils import METRIC_REGISTRY from copy import deepcopy # 注:原代码可能遗漏该导入,需补充 __all__ = ['calculate_psnr', 'calculate2'] def calculate_metric(data, opt): opt = deepcopy(opt) metric_type = opt.pop('type') metric = METRIC_REGISTRY.get(metric_type)(**data, **opt) return metric
Registry类代码
class Registry(): def __init__(self, name): self._name = name self._obj_map = {} def _do_register(self, name, obj, suffix=None): if isinstance(suffix, str): name = name + '_' + suffix assert (name not in self._obj_map), (f"An object named '{name}' was already registered " f"in '{self._name}' registry!") self._obj_map[name] = obj def register(self, obj=None, suffix=None): if obj is None: # 用作装饰器 def deco(func_or_class): name = func_or_class.__name__ self._do_register(name, func_or_class, suffix) return func_or_class return deco # 用作函数调用 name = obj.__name__ self._do_register(name, obj, suffix) def get(self, name, suffix='basicsr'): ret = self._obj_map.get(name) if ret is None: ret = self._obj_map.get(name + '_' + suffix) print(f'Name {name} is not found, use name: {name}_{suffix}!') if ret is None: raise KeyError(f"No object named '{name}' found in '{self._name}' registry!") return ret def __contains__(self, name): return name in self._obj_map def __iter__(self): return iter(self._obj_map.items()) def keys(self): return self._obj_map.keys() LOSS_REGISTRY = Registry('loss') METRIC_REGISTRY = Registry('metric')
实现示例(以calculate_ssim为例)
@METRIC_REGISTRY.register() def calculate_ssim(val, val2, **kwargs): # 具体实现逻辑 return somevalue
调用代码
calculate_metric(metric_data, opt_)
错误信息
-> 2582 self.metric_results[name] += calculate_metric(metric_data, opt_) 2584 if use_pbar: Input In [53], in calculate_metric(data, opt) 11 opt = deepcopy(opt) 12 metric_type = opt.pop('type') ---> 13 metric = METRIC_REGISTRY.get(metric_type)(**data, **opt) 14 return metric Input In [54], in Registry.get(self, name, suffix) 1870 print(f'Name {name} is not found, use name: {name}_{suffix}!') 1871 if ret is None: -> 1872 raise KeyError(f"No object named '{name}' found in '{self._name}' registry!") 1873 return ret KeyError: "No object named 'calculate_psnr' found in 'metric' registry!"
解决方案
一、修复注册表缺失问题
核心是确保calculate_psnr函数被正确注册到METRIC_REGISTRY中:
- 定义并注册
calculate_psnr:在Notebook中编写该函数并使用装饰器注册,示例:
@METRIC_REGISTRY.register() def calculate_psnr(val, val2, **kwargs): # 标准PSNR计算逻辑 import math mse = ((val - val2)**2).mean() if mse == 0: return float('inf') return 20 * math.log10(255.0 / math.sqrt(mse))
- 确保注册时机:必须在调用
calculate_metric之前执行上述注册代码,让函数被加入注册表。 - 调试验证:调用
calculate_metric前,打印注册表键确认calculate_psnr存在:
print(METRIC_REGISTRY.keys())
二、绕过Registry机制快速测试
如果仅需临时测试,可直接修改calculate_metric函数,用字典映射替代注册表:
def calculate_metric(data, opt): opt = deepcopy(opt) metric_type = opt.pop('type') # 直接建立函数映射表 metric_map = { 'calculate_psnr': calculate_psnr, 'calculate_ssim': calculate_ssim # 按需添加其他metric函数 } metric_func = metric_map.get(metric_type) if not metric_func: raise ValueError(f"Metric {metric_type} not found!") return metric_func(**data, **opt)
这种方式无需依赖Registry类,直接通过字典调用目标函数,适合快速验证逻辑。
内容的提问来源于stack exchange,提问作者user836026
相关产品推荐
相关产品推荐

