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

如何为NumPy ndarray添加自定义语义标识并实现合理索引?

问题描述

我正在编写一个Python 3类,使用NumPy的np.ndarray存储数据,同时希望为该类添加数据语义解析信息。

例如,假设ndarray的dtype为np.float32,存在一个"color"标识(实际为整数)修改浮点值的语义:若要相加"red"和"blue"标识的数组,需先将两者转换为"magenta"标识,结果的_color为"magenta"。实际场景中,标识的转换与运算结果的标识均由数学规则定义。

现有类实现如下:

import numpy as np

class MyClass:
    
    def __init__(self, data : np.ndarray, color : str):
        self._data = data
        self._color = color
    
    
    # Example: Adding red numbers and blue numbers produces magenta numbers
    def convert(self, other_color):
        if self._color == "red" and other_color == "blue":
            return MyClass(10*self._data, "magenta")
        elif self._color == "blue" and other_color == "red":
            return MyClass(self._data/10, "magenta")
    
    
    def __add__(self, other):
        if other._color == self._color:
            # If the colors match, then just add the data values
            return MyClass(self._data + other._data, self._color)
        else:
            # If the colors don't match, then convert to the output color before adding
            new_self = self.convert(other._color)
            new_other = other.convert(self._color)
            return new_self + new_other

当前问题在于_color信息与_data分离,无法定义合理的索引行为:

  • 若定义__getitem__返回self._data[i],会丢失_color信息;
  • 若返回MyClass(self._data[i], self._color),会生成含标量的对象,引发索引错误;
  • 若返回MyClass(self._data[i:i+1], self._color),会导致赋值等操作报错。

我曾考虑为不同标识设置不同dtype,但不知如何实现。理论上标识总数约10万,但单次脚本使用不超100个。同时我不想为每个数据元素存储标识(数组可达数十亿元素,全局一个标识即可)。

请问如何在保留该语义标识的同时,实现可用的类?

解决方案

1. 完善__getitem__,兼容标量与数组场景

针对索引返回的结果类型做判断,分别返回带标识的自定义标量或MyClass实例,同时实现__setitem__处理赋值逻辑:

import numpy as np

# 自定义标量类,保留color标识
class MyScalar:
    def __init__(self, value, color):
        self.value = value
        self.color = color
    
    # 按需实现标量运算方法
    def __add__(self, other):
        if isinstance(other, MyScalar) and other.color == self.color:
            return MyScalar(self.value + other.value, self.color)
        elif isinstance(other, MyClass):
            return other + self
        else:
            return MyScalar(self.value + other, self.color)

class MyClass:
    
    def __init__(self, data : np.ndarray, color : str):
        self._data = data
        self._color = color
    
    def convert(self, other_color):
        if self._color == "red" and other_color == "blue":
            return MyClass(10*self._data, "magenta")
        elif self._color == "blue" and other_color == "red":
            return MyClass(self._data/10, "magenta")
        elif self._color == other_color:
            return self
        else:
            raise ValueError(f"No conversion rule from {self._color} to {other_color}")
    
    def __add__(self, other):
        if isinstance(other, MyScalar):
            other = MyClass(np.array([other.value]), other.color)
        if other._color == self._color:
            return MyClass(self._data + other._data, self._color)
        else:
            new_self = self.convert(other._color)
            new_other = other.convert(self._color)
            return new_self + new_other
    
    def __getitem__(self, idx):
        result = self._data[idx]
        if np.isscalar(result):
            return MyScalar(result, self._color)
        else:
            return MyClass(result, self._color)
    
    def __setitem__(self, idx, value):
        if isinstance(value, MyScalar):
            if value.color != self._color:
                raise ValueError(f"Cannot assign scalar with color {value.color} to array with color {self._color}")
            self._data[idx] = value.value
        elif isinstance(value, MyClass):
            if value._color != self._color:
                raise ValueError(f"Cannot assign array with color {value._color} to array with color {self._color}")
            self._data[idx] = value._data
        else:
            self._data[idx] = value

2. 自定义NumPy dtype绑定标识

利用NumPy的metadata特性(1.17+支持),把标识作为dtype的元数据,无需单独存储_color属性,同时避免每个元素存标识的开销:

import numpy as np

# 维护标识到自定义dtype的注册表
color_dtype_registry = {}

def get_color_dtype(color):
    if color not in color_dtype_registry:
        # 基于float32创建带color元数据的dtype
        base_dtype = np.dtype("float32")
        color_dtype = np.dtype(base_dtype, metadata={"color": color})
        color_dtype_registry[color] = color_dtype
    return color_dtype_registry[color]

class MyClass:
    def __init__(self, data : np.ndarray, color : str):
        self._data = data.astype(get_color_dtype(color))
    
    @property
    def color(self):
        return self._data.dtype.metadata["color"]
    
    def convert(self, target_color):
        current_color = self.color
        if current_color == target_color:
            return self
        if current_color == "red" and target_color == "blue":
            converted_data = 10 * self._data
        elif current_color == "blue" and target_color == "red":
            converted_data = self._data / 10
        elif current_color in ["red", "blue"] and target_color == "magenta":
            converted_data = 10 * self._data if current_color == "red" else self._data /10
        else:
            raise ValueError(f"No conversion rule from {current_color} to {target_color}")
        return MyClass(converted_data, target_color)
    
    def __add__(self, other):
        if not isinstance(other, MyClass):
            raise TypeError("Only MyClass instances can be added")
        if self.color == other.color:
            return MyClass(self._data + other._data, self.color)
        else:
            target_color = "magenta"
            self_converted = self.convert(target_color)
            other_converted = other.convert(target_color)
            return self_converted + other_converted
    
    def __getitem__(self, idx):
        result = self._data[idx]
        if np.isscalar(result):
            return MyScalar(result.item(), self.color)
        else:
            return MyClass(result, self.color)
    
    def __setitem__(self, idx, value):
        if isinstance(value, MyScalar):
            if value.color != self.color:
                raise ValueError("Color mismatch in assignment")
            self._data[idx] = value.value
        elif isinstance(value, MyClass):
            if value.color != self.color:
                raise ValueError("Color mismatch in assignment")
            self._data[idx] = value._data
        else:
            self._data[idx] = value

class MyScalar:
    def __init__(self, value, color):
        self.value = value
        self.color = color
    
    def __add__(self, other):
        if isinstance(other, MyScalar) and other.color == self.color:
            return MyScalar(self.value + other.value, self.color)
        elif isinstance(other, MyClass):
            return other + self
        else:
            return MyScalar(self.value + other, self.color)

3. 动态生成子类绑定标识

针对每个标识生成对应子类,用类类型体现语义信息,适合单次使用标识数量可控的场景:

import numpy as np

class MyClassMeta(type):
    _registry = {}
    
    def __new__(cls, name, bases, attrs):
        new_cls = super().__new__(cls, name, bases, attrs)
        if "color" in attrs:
            cls._registry[attrs["color"]] = new_cls
        return new_cls
    
    @classmethod
    def get_class(cls, color):
        if color not in cls._registry:
            # 动态生成对应color的子类
            subclass_name = f"{color.capitalize()}MyClass"
            subclass = cls(subclass_name, (MyBaseClass,), {"color": color})
        return cls._registry[color]

class MyBaseClass(metaclass=MyClassMeta):
    color = None
    
    def __init__(self, data : np.ndarray):
        self._data = data
    
    def convert(self, target_color):
        current_color = self.color
        if current_color == target_color:
            return self
        target_cls = MyClassMeta.get_class(target_color)
        if current_color == "red" and target_color == "blue":
            return target_cls(10 * self._data)
        elif current_color == "blue" and target_color == "red":
            return target_cls(self._data /10)
        elif current_color in ["red", "blue"] and target_color == "magenta":
            converted_data = 10 * self._data if current_color == "red" else self._data /10
            return target_cls(converted_data)
        else:
            raise ValueError(f"No conversion rule from {current_color} to {target_color}")
    
    def __add__(self, other):
        if not isinstance(other, MyBaseClass):
            raise TypeError("Only MyBaseClass instances can be added")
        if self.color == other.color:
            return type(self)(self._data + other._data)
        else:
            target_color = "magenta"
            self_converted = self.convert(target_color)
            other_converted = other.convert(target_color)
            return self_converted + other_converted
    
    def __getitem__(self, idx):
        result = self._data[idx]
        if np.isscalar(result):
            # 返回对应color的标量对象
            return type(f"{self.color.capitalize()}Scalar", (), {"value": result, "color": self.color})()
        else:
            return type(self)(result)
    
    def __setitem__(self, idx, value):
        if hasattr(value, "color") and value.color != self.color:
            raise ValueError("Color mismatch in assignment")
        self._data[idx] = value.value if hasattr(value, "value") else value

# 预定义常用子类,或动态生成
class RedMyClass(MyBaseClass):
    color = "red"

class BlueMyClass(MyBaseClass):
    color = "blue"

class MagentaMyClass(MyBaseClass):
    color = "magenta"

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 10:49:59