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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 10:20:41