如何为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
相关产品推荐
相关产品推荐

