如何在PySpark中逐元素比较两列并统计符合条件的行数?
问题
我想实现一个简单操作:在PySpark的多列DataFrame中,用存储列名的变量计算满足x_col < y_col的行数。我想用df.filter().count()的方式(在Pandas里这方法没问题),但PySpark里跑不通:
x_col = col(self.column_x) # self.column_x 是字符串 y_col = col(self.column_y) # 此处同理 tmp_df = df.filter(x_col < y_col) cnt = tmp_df.count()
之前把其中一列换成数字时能正常运行,所以我怀疑是不是缺了某种转换,没法让代码把它们当成数据比较,而是在比较列对象?
另外理想情况是能配合relation函数用:
if self.strict: relation = lambda x, y: x < y else: relation = lambda x, y: x <= y tmp_df = df.filter(relation(x_col, y_col))
不过现在只要能解决问题,用if判断区分<和<=也可以接受。
解决方案
核心问题排查
你的代码写法本身是对的,col()生成的Column对象支持直接用</<=比较生成过滤条件,问题大概率出在列的数据类型不兼容上:
- 如果其中一列是字符串类型、另一列是数值类型,直接比较会报错;
- 就算都是数值类型,要是存在
null值,也可能导致过滤逻辑不符合预期(PySpark中null参与比较会返回null,不会被计入结果)。
修复步骤
- 统一列数据类型
先检查两列的类型,把它们转换成相同的数值类型(比如DoubleType):
from pyspark.sql.types import DoubleType from pyspark.sql.functions import col x_col = col(self.column_x).cast(DoubleType()) y_col = col(self.column_y).cast(DoubleType())
- 处理null值(可选但推荐)
如果数据里有null,可以选择过滤掉null值再比较,或者给null设置默认值:
# 方式1:先过滤null值 filtered_df = df.filter(col(self.column_x).isNotNull() & col(self.column_y).isNotNull()) tmp_df = filtered_df.filter(x_col < y_col) # 方式2:给null设默认值(比如0) x_col = col(self.column_x).cast(DoubleType()).fillna(0) y_col = col(self.column_y).cast(DoubleType()).fillna(0) tmp_df = df.filter(x_col < y_col)
支持relation函数的写法
你的lambda写法其实是可行的,只要保证列类型统一,直接用就行:
from pyspark.sql.functions import col x_col = col(self.column_x).cast(DoubleType()) y_col = col(self.column_y).cast(DoubleType()) if self.strict: relation = lambda x, y: x < y else: relation = lambda x, y: x <= y cnt = df.filter(relation(x_col, y_col)).count()
或者更简洁的方式,直接用条件表达式生成过滤条件:
filter_condition = x_col < y_col if self.strict else x_col <= y_col cnt = df.filter(filter_condition).count()
内容的提问来源于stack exchange,提问作者Igor Agafonov
相关产品推荐
相关产品推荐

