如何为TensorFlow交叉列指定交叉项词汇表?
解决Spark中过滤低频/零频交叉特征的问题
这个需求在特征工程里太常见了——交叉特征能捕捉变量间的关联,但低频甚至零频的组合不仅带不来有效信息,还会增加特征维度、引入噪声。在Spark里我们可以分几步来实现,既保证准确性又兼顾分布式效率:
核心思路
先基于训练数据统计所有animal和food的交叉组合频率,筛选出符合阈值的有效组合;再用这些有效组合作为规则,生成过滤后的交叉特征列。这里一定要用训练数据的统计结果,绝对不能碰测试数据,避免数据泄露!
具体实现步骤(Python示例)
假设你的训练数据是df_train,包含animal和food两列:
1. 统计有效交叉组合
首先计算每个交叉组合的出现次数,过滤掉低于阈值的项(这里阈值设为2,刚好排除出现1次的puma x quinoa和零频的cat x quinoa):
from pyspark.sql import functions as F # 统计交叉项频率 cross_freq_df = df_train.groupBy("animal", "food").count() # 过滤出频率>=2的有效组合 valid_cross_df = cross_freq_df.filter(F.col("count") >= 2).select("animal", "food") # 转换成集合并广播——分布式场景下广播变量能大幅提升效率 valid_cross_set = set(valid_cross_df.rdd.map(lambda row: (row.animal, row.food)).collect()) broadcast_valid = spark.sparkContext.broadcast(valid_cross_set)
2. 定义UDF生成过滤后的交叉特征
写一个自定义函数,判断当前的animal和food组合是否在有效集合里,是就生成交叉字符串,否则返回统一的"other_other"(或者null,看你后续特征处理的需求):
from pyspark.sql.types import StringType def generate_valid_cross(animal, food): if (animal, food) in broadcast_valid.value: return f"{animal}_{food}" else: return "other_other" # 也可以返回None,后续用fillna处理 # 注册成Spark UDF cross_udf = F.udf(generate_valid_cross, StringType())
3. 生成最终交叉特征列
调用UDF给数据集添加交叉列:
df_with_filtered_cross = df_train.withColumn( "animal_food_cross", cross_udf(F.col("animal"), F.col("food")) )
进阶:封装成ML Pipeline组件
如果你的特征工程是基于Spark ML Pipeline做的,可以把上面的逻辑封装成自定义Transformer,这样能和其他组件(比如StringIndexer、OneHotEncoder)无缝集成:
from pyspark.ml import Transformer from pyspark.ml.param.shared import HasInputCols, HasOutputCol, Param, Params class CrossFeatureFilter(Transformer, HasInputCols, HasOutputCol): # 定义阈值参数 threshold = Param(Params._dummy(), "threshold", "Minimum frequency for valid cross feature") def __init__(self, inputCols=None, outputCol=None, threshold=2): super().__init__() self._setDefault(threshold=2) self.setInputCols(inputCols) self.setOutputCol(outputCol) self.setThreshold(threshold) def getThreshold(self): return self.getOrDefault(self.threshold) def _transform(self, df): col1, col2 = self.getInputCols() # 复用前面的统计逻辑 cross_freq_df = df.groupBy(col1, col2).count() valid_cross_df = cross_freq_df.filter(F.col("count") >= self.getThreshold()).select(col1, col2) valid_cross_set = set(valid_cross_df.rdd.map(lambda row: (row[col1], row[col2])).collect()) broadcast_valid = df.sparkSession.sparkContext.broadcast(valid_cross_set) def generate_cross(a, b): return f"{a}_{b}" if (a, b) in broadcast_valid.value else "other_other" cross_udf = F.udf(generate_cross, StringType()) return df.withColumn(self.getOutputCol(), cross_udf(F.col(col1), F.col(col2))) # 使用方式: cross_filter = CrossFeatureFilter( inputCols=["animal", "food"], outputCol="animal_food_cross", threshold=2 ) df_with_cross = cross_filter.transform(df_train)
关键注意点
- 数据泄露防范:必须用训练数据统计有效组合,测试数据要完全复用这个规则,哪怕测试数据里出现了训练时没见过的组合,也要归为"other"。
- 阈值调整:阈值不一定是2,你可以根据业务场景调整,比如如果是分类任务,也可以结合卡方检验来筛选有统计显著性的交叉项。
- 空值处理:如果选择返回null,后续要记得用
fillna替换成统一类别,否则后续的编码组件(比如OneHotEncoder)会报错。
内容的提问来源于stack exchange,提问作者MrCartoonology
相关产品推荐
相关产品推荐

