如何重写自定义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
相关产品推荐
相关产品推荐

