You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

为何针对多维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)

关键说明

  1. 转换为可哈希类型:a.tolist()将任意维度的NumPy数组转为Python原生嵌套列表,再通过tuple()转为嵌套元组,所有元素均为可哈希的原生数值类型(如int、float)。
  2. 还原数组类型:cached_wrapper中使用np.array(a)将嵌套元组转回原NumPy数组,保证函数接收的参数类型与预期一致。
  3. 通用性:该方案支持一维、二维、三维等任意维度的NumPy数组,且不影响其他类型参数的缓存逻辑。

内容的提问来源于stack exchange,提问作者Offel21

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 10:47:03