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

如何为含Individual对象的Population类实现二维掩码索引并保留类型?

保留Individual类型的NumPy掩码切片与赋值方案

完整实现代码

import numpy as np

class Individual:
    def __init__(self, vector: np.ndarray) -> None:
        self.value = vector
        
    def __getitem__(self, index):
        return self.value[index]
    
    def __setitem__(self, index, value):
        self.value[index] = value

class Population:
    def __init__(self, individuals=np.array([])):
        self.individuals = individuals

    def __getitem__(self, index):
        result = self.individuals[index]
        # 处理单个元素的索引,返回Individual实例而非0维数组
        if isinstance(result, np.ndarray) and result.ndim == 0:
            return result.item()
        return result
    
    def __setitem__(self, index, value):
        if isinstance(index, np.ndarray) and index.ndim == 2:
            # 校验掩码与赋值对象的形状匹配
            if index.shape != value.shape:
                raise ValueError("Mask shape must match the shape of assigned value")
            
            # 遍历每个个体,用NumPy向量操作完成掩码赋值
            for ind_idx, (row_mask, source) in enumerate(zip(index, value)):
                # 自动适配Individual对象或纯数组类型的赋值源
                source_array = source.value if isinstance(source, Individual) else source
                self.individuals[ind_idx][row_mask] = source_array[row_mask]
        else:
            # 处理普通索引赋值(比如单元素、切片)
            self.individuals[index] = value

# 测试示例
if __name__ == "__main__":
    ind1 = Individual(np.array([1, 2]))
    ind2 = Individual(np.array([3, 4]))
    population = Population(np.array([ind1, ind2]))

    ind3 = Individual(np.array([10, 20]))
    ind4 = Individual(np.array([30, 40]))
    population2 = Population(np.array([ind3, ind4]))

    mask = np.array([[True, False], [False, True]])
    population[mask] = population2[mask]

    print(population[0].value)  # 输出: [10  2]
    print(population[1].value)  # 输出: [ 3 40]
    print(type(population[0]))  # 输出: <class '__main__.Individual'>

关键修改说明

  1. 给Individual添加__setitem__方法:
    允许直接通过individual[mask] = value的方式对其内部的value数组进行赋值,利用NumPy的向量操作提升效率。

  2. Population的__getitem__优化:
    处理单个元素索引场景,避免NumPy返回0维object数组,直接返回Individual实例,保证type(population[0])符合预期。

  3. Population的__setitem__掩码处理:

    • 专门识别二维掩码类型,将掩码操作拆解为对每个Individual的value数组的局部赋值
    • 自动适配赋值源是Individual对象还是纯NumPy数组,无需额外转换
    • 外层仅遍历Population中的个体数量,内部赋值用NumPy向量操作,兼顾效率与类结构保留

原理说明

直接用NumPy存储自定义对象时,二维掩码索引会触发NumPy对对象内部结构的展开,导致返回值变为纯数组而非Individual实例。通过在Population层面对掩码操作进行封装,我们将全局的二维掩码映射为每个个体的一维掩码操作,既保留了外层的Individual类结构,又充分利用了NumPy对数组元素的高效处理能力,避免了纯Python逐元素循环的性能损耗。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 09:10:39