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

如何重写自定义DataFrame类的__getitem__实现元素级运算与布尔索引?

自定义DataFrame类实现元素级运算与布尔索引需求

给定如下数据及自定义DataFrame类:

frame = {
    "a": ["X4E", "T3B", "F8D", "C7X"],
    "b": [7.0, 3.5, 8.0, 6.0],
    "c": [5, 3, 1, 10],
    "d": [False, False, True, False]
}

df = DataFrame(frame)

希望重写__getitem__方法,让该类支持以下操作:

  • 单条件布尔索引取列:res = df[(df["b"] + 5.0 > 10.0)]["a"](返回满足b+5.0>10.0的a列数据)
  • 多条件组合布尔索引取列:res = df[(df["b"] + 5.0 > 10.0) & (df["c"] > 3) & ~df["d"]]["a"]

已知__getitem__的基本用法,但不清楚如何实现类似Pandas的元素级数学运算,求实现思路。


实现思路

1. 封装列数据:用Column类承载元素级运算

不能直接用原生列表存储列数据,需要编写一个Column类封装每一列,所有元素级运算、布尔操作都在这个类中实现。核心是重载各类魔法方法:

  • 算术运算符:__add__、__sub__、__mul__等,实现元素级运算并返回新的Column实例
  • 比较运算符:__gt__、__lt__、__eq__等,返回存储布尔值的Column(即布尔掩码)
  • 位运算符:__and__、__or__、__invert__(对应&、|、~),实现多布尔条件的组合,返回布尔Column

示例Column类核心代码:

class Column:
    def __init__(self, data):
        self.data = data.copy()
    
    # 元素级加法
    def __add__(self, other):
        if isinstance(other, (int, float)):
            return Column([x + other for x in self.data])
        elif isinstance(other, Column):
            return Column([x + y for x, y in zip(self.data, other.data)])
        raise TypeError("不支持的运算类型")
    
    # 大于比较,返回布尔Column
    def __gt__(self, other):
        if isinstance(other, (int, float)):
            return Column([x > other for x in self.data])
        elif isinstance(other, Column):
            return Column([x > y for x, y in zip(self.data, other.data)])
        raise TypeError("不支持的比较类型")
    
    # 布尔掩码的位与运算
    def __and__(self, other):
        if isinstance(other, Column) and all(isinstance(x, bool) for x in self.data) and all(isinstance(x, bool) for x in other.data):
            return Column([x & y for x, y in zip(self.data, other.data)])
        raise TypeError("位与运算仅支持布尔类型Column")
    
    # 布尔掩码取反
    def __invert__(self):
        if all(isinstance(x, bool) for x in self.data):
            return Column([not x for x in self.data])
        raise TypeError("取反运算仅支持布尔类型Column")
    
    # 用掩码筛选列数据
    def __getitem__(self, mask):
        if isinstance(mask, Column) and all(isinstance(x, bool) for x in mask.data):
            return Column([self.data[i] for i, val in enumerate(mask.data) if val])
        elif isinstance(mask, list) and all(isinstance(x, bool) for x in mask):
            return Column([self.data[i] for i, val in enumerate(mask) if val])
        raise TypeError("掩码必须是布尔Column或布尔列表")
    
    # 转成原生列表方便输出
    def to_list(self):
        return self.data.copy()

2. 重构DataFrame的__getitem__方法

DataFrame内部用字典存储Column实例(而非原生列表),__getitem__需要处理两种场景:

  • 传入字符串(列名):返回对应的Column实例
  • 传入布尔Column(掩码):返回新的DataFrame实例,仅保留掩码为True的行

示例DataFrame类核心代码:

class DataFrame:
    def __init__(self, data_dict):
        # 将输入字典的每个值转为Column实例
        self.columns = {col: Column(data) for col, data in data_dict.items()}
        # 记录行数(假设所有列长度一致)
        self.row_count = len(next(iter(data_dict.values()))) if data_dict else 0
    
    def __getitem__(self, key):
        # 场景1:获取指定列,返回Column实例
        if isinstance(key, str):
            return self.columns[key]
        # 场景2:用布尔掩码筛选行,返回新DataFrame
        elif isinstance(key, Column) and all(isinstance(x, bool) for x in key.data):
            if len(key.data) != self.row_count:
                raise ValueError("掩码长度与数据行数不匹配")
            # 对每个列应用掩码,生成新数据
            new_data = {}
            for col_name, col in self.columns.items():
                new_col = col[key]
                new_data[col_name] = new_col.data
            return DataFrame(new_data)
        raise TypeError("不支持的键类型")

3. 功能验证

测试需求中的操作:

frame = {
    "a": ["X4E", "T3B", "F8D", "C7X"],
    "b": [7.0, 3.5, 8.0, 6.0],
    "c": [5, 3, 1, 10],
    "d": [False, False, True, False]
}

df = DataFrame(frame)

# 单条件筛选取a列
res1 = df[(df["b"] + 5.0 > 10.0)]["a"].to_list()
print(res1)  # 输出: ['X4E', 'F8D']

# 多条件组合筛选取a列
mask = (df["b"] + 5.0 > 10.0) & (df["c"] > 3) & ~df["d"]
res2 = df[mask]["a"].to_list()
print(res2)  # 输出: ['X4E']

4. 扩展优化方向

  • 补充更多运算符重载(减法、乘法、小于等于等),完善Column类的运算能力
  • 处理空数据、列长度不一致等边界情况
  • 改用numpy数组代替列表推导,提升大数据量下的运算性能
  • 实现Column和DataFrame的__repr__方法,优化打印格式

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 04:20:23