为何针对多维numpy数组的函数缓存会抛出错误?
解决多维NumPy数组函数缓存的TypeError问题
问题描述
尝试为接收多维NumPy数组的函数实现缓存时,触发TypeError: unhashable type: 'numpy.ndarray'错误。原代码试图将数组转为元组来满足缓存的可哈希要求,但二维数组转换后仍包含不可哈希的子数组,导致缓存失效。
原代码
import numpy as np from functools import cache, wraps def np_cache(function): @cache def cached_wrapper(*args, **kwargs): args = [np.array(a) if isinstance(a, tuple) else a for a in args] kwargs = { k: np.array(v) if isinstance(v, tuple) else v for k, v in kwargs.items() } return function(*args, **kwargs) @wraps(function) def wrapper(*args, **kwargs): args = [tuple(a) if isinstance(a, np.ndarray) else a for a in args] kwargs = { k: tuple(v) if isinstance(v, np.ndarray) else v for k, v in kwargs.items() } return cached_wrapper(*args, **kwargs) wrapper.cache_info = cached_wrapper.cache_info wrapper.cache_clear = cached_wrapper.cache_clear return wrapper x = np.array([[1,2],[3,4]]) y = np.array([2, 4]) @np_cache def test2(x, shown = False): if shown: print(x) return x # test2(y,True) test2(x, True)
错误信息
(array([[1, 2], [3, 4]]), True) [(array([1, 2]), array([3, 4])), True] --------------------------------------------------------------------------- TypeError Traceback (most recent call last) <ipython-input-79-73b6dd94e3a9> in <cell line: 42>() 40 41 #test2(y,True) ---> 42 test2(x, True) 43 44 <ipython-input-79-73b6dd94e3a9> in wrapper(*args, **kwargs) 23 k: tuple(v) if isinstance(v, np.ndarray) else v for k, v in kwargs.items() 24 } ---> 25 return cached_wrapper(*args, **kwargs) 26 27 wrapper.cache_info = cached_wrapper.cache_info TypeError: unhashable type: 'numpy.ndarray'
问题根源
对二维NumPy数组使用tuple(a)时,仅会将数组的行作为元素生成元组(如(array([1,2]), array([3,4]))),元组内的元素仍是不可哈希的NumPy数组,因此无法被@cache正常处理。
修复方案
需要将多维数组完全转换为Python原生类型的嵌套元组,确保所有层级的元素都是可哈希的。利用tolist()方法将数组转为嵌套列表,再转为元组即可实现这一点。
修改后的代码
import numpy as np from functools import cache, wraps def np_cache(function): @cache def cached_wrapper(*args, **kwargs): # 将嵌套元组转回NumPy数组 args = [np.array(a) if isinstance(a, tuple) else a for a in args] kwargs = { k: np.array(v) if isinstance(v, tuple) else v for k, v in kwargs.items() } return function(*args, **kwargs) @wraps(function) def wrapper(*args, **kwargs): # 将多维NumPy数组转为全原生类型的嵌套元组 args = [tuple(a.tolist()) if isinstance(a, np.ndarray) else a for a in args] kwargs = { k: tuple(v.tolist()) if isinstance(v, np.ndarray) else v for k, v in kwargs.items() } return cached_wrapper(*args, **kwargs) wrapper.cache_info = cached_wrapper.cache_info wrapper.cache_clear = cached_wrapper.cache_clear return wrapper x = np.array([[1,2],[3,4]]) y = np.array([2, 4]) @np_cache def test2(x, shown=False): if shown: print(x) return x test2(y, True) test2(x)
关键说明
- 转换为可哈希类型:
a.tolist()将任意维度的NumPy数组转为Python原生嵌套列表,再通过tuple()转为嵌套元组,所有元素均为可哈希的原生数值类型(如int、float)。 - 还原数组类型:
cached_wrapper中使用np.array(a)将嵌套元组转回原NumPy数组,保证函数接收的参数类型与预期一致。 - 通用性:该方案支持一维、二维、三维等任意维度的NumPy数组,且不影响其他类型参数的缓存逻辑。
内容的提问来源于stack exchange,提问作者Offel21
相关产品推荐
相关产品推荐

