Numba jitclass的__getitem__如何区分slice并处理非slice键转换
Numba jitclass中实现slice与普通键的区分处理
我正在使用Numba的jitclass,希望在键不是slice类型时对其进行转换,同时保留slice功能。
问题: 如何实现这一需求?
补充背景: 我更倾向于编写tensor[coord]而非tensor[tensor_to_formalseries(coord, tensor.dim)],也更喜欢用简洁的tensor[:key]替代tensor.formal_series[:key]。
以下是三个在纯Python中可行,但在jitclass中无法运行的尝试:
尝试1:使用Python原生isinstance检查slice类型
import numpy as np from numba import njit from numba.experimental import jitclass from numba.core.types import int64, SliceType @njit def tensor_to_formalseries(coordinate: int, dim=2): key = coordinate * 2 # 自定义坐标转换逻辑 return key @jitclass(spec={"formal_series": int64[:], "dim": int64}) class Tensor1: def __init__(self, dim=2): self.dim = dim self.formal_series = np.arange(10) def __getitem__(self, key): if isinstance(key, (slice, SliceType)): print("key is a slice") return self.formal_series[key] else: print("key is not a slice") return self.formal_series[tensor_to_formalseries(key, self.dim)] tensor = Tensor1() print(tensor[:3]) print(tensor[2])
错误:
NumbaTypeError: isinstance()不支持类型为"slice<a:b>"的变量。
尝试2:通过索引结果类型判断
@jitclass(spec={"formal_series": int64[:], "dim": int64}) class Tensor2: def __init__(self, dim=2): self.dim = dim self.formal_series = np.arange(10) def __getitem__(self, key): if isinstance(self.formal_series[key], np.ndarray): print("key is a slice") return self.formal_series[key] else: print("key is not a slice") return self.formal_series[tensor_to_formalseries(key, self.dim)] tensor = Tensor2() print(tensor[:3]) print(tensor[2])
错误:
使用了不支持的NumPy函数'numpy.ndarray'或该函数的不支持用法。
尝试3:通过异常捕获区分
@jitclass(spec={"formal_series": int64[:], "dim": int64}) class Tensor3: def __init__(self, dim=2): self.dim = dim self.formal_series = np.arange(10) def __getitem__(self, key): try: len(self.formal_series[key]) print("key is a slice") return self.formal_series[key] except Exception: print("key is not a slice") return self.formal_series[tensor_to_formalseries(key, self.dim)] tensor = Tensor3() print(tensor[:3]) print(tensor[2])
错误:
调用tensor[:3]时:函数'mul'重载错误,参数为'(slice<a:b>, int64)'无匹配项。
调用tensor[2]时:函数'len'重载错误,参数为'(int64)'无匹配项。
可行解决方案
方案1:使用Numba专属类型检查函数
在Numba JIT环境中,需要用Numba提供的numba.core.types.isinstance来替代Python原生isinstance,它能正确识别Numba内部类型(如SliceType)。
import numpy as np from numba import njit from numba.experimental import jitclass from numba.core.types import int64, SliceType, isinstance as numba_isinstance @njit def tensor_to_formalseries(coordinate: int, dim=2): key = coordinate * 2 # 自定义坐标转换逻辑 return key @jitclass(spec={"formal_series": int64[:], "dim": int64}) class Tensor: def __init__(self, dim=2): self.dim = dim self.formal_series = np.arange(10) def __getitem__(self, key): if numba_isinstance(key, SliceType): return self.formal_series[key] else: converted_key = tensor_to_formalseries(key, self.dim) return self.formal_series[converted_key] # 测试 tensor = Tensor() print(tensor[:3]) # 输出: [0 1 2] print(tensor[2]) # 输出: 4
方案2:方法重载(更高效)
通过@overload_method为__getitem__分别实现int和slice类型参数的处理逻辑,类型明确时编译效率更高。
import numpy as np from numba import njit, overload_method from numba.experimental import jitclass from numba.core.types import int64, SliceType @njit def tensor_to_formalseries(coordinate: int, dim=2): key = coordinate * 2 # 自定义坐标转换逻辑 return key @jitclass(spec={"formal_series": int64[:], "dim": int64}) class Tensor: def __init__(self, dim=2): self.dim = dim self.formal_series = np.arange(10) def __getitem__(self, key): pass # 占位,实际逻辑由重载实现 # 重载处理slice类型键 @overload_method(Tensor, "__getitem__") def ol_getitem_slice(self, key: SliceType): def impl(self, key): return self.formal_series[key] return impl # 重载处理int类型键 @overload_method(Tensor, "__getitem__") def ol_getitem_int(self, key: int64): def impl(self, key): converted_key = tensor_to_formalseries(key, self.dim) return self.formal_series[converted_key] return impl # 测试 tensor = Tensor() print(tensor[:3]) # 输出: [0 1 2] print(tensor[2]) # 输出: 4
内容的提问来源于stack exchange,提问作者Louis-Amand
相关产品推荐
相关产品推荐

