Spark DataFrame逐行读取Path列Avro文件提取Accuracy方案
需求说明
需要处理DataFrame指定列,从每条记录对应的Avro文件提取指标,核心逻辑:
- 逐行读取
Path列存储的Avro文件访问路径 - 读取对应路径的Avro文件为DataFrame,提取Struct结构中存储的
accuracy指标 - 新增
Accuracy列存储提取到的精度值
等价逻辑:对
Path列每一行的路径值,都执行spark.read.format("com.databricks.spark.avro").load(avro_path)读取操作
数据结构示例
输入DataFrame结构
+----------+-----+--------------------------+ |timestamp |Model| Path | +----------+-----+--------------------------+ |11:02 |Vgg |projects/Vgg/results.avro | |18:31 |Dnet |projects/Dnet/results.avro| |15:54 |Rnet |projects/Rnet/results.avro| |12:19 |ViT |projects/ViT/results.avro | +----------+-----+--------------------------+
期望输出结构
+----------+-----+--------------------------+-----------+ |timestamp |Model| Path | Accuracy | +----------+-----+--------------------------+-----------+ |11:02 |Vgg |projects/Vgg/results.avro | 0.72 | |18:31 |Dnet |projects/Dnet/results.avro| 0.78 | |15:54 |Rnet |projects/Rnet/results.avro| 0.75 | |12:19 |ViT |projects/ViT/results.avro | 0.82 | +----------+-----+--------------------------+-----------+
已尝试的不可行方案
方案1:UDF实现
UDF属于分布式执行的算子,内部无法执行Spark驱动级的DataFrame读取操作,运行会抛出SparkException、NullPointerException异常,代码如下:
val get_auc: (String => String) = (avro_path: String) => { val auc_avro_file = spark.read.format("com.databricks.spark.avro").load(avro_path) val auc = auc_avro_file.select("metrics.Accuracy").first.toString auc } val auc_udf = udf(get_auc) val auc_path = models_df.withColumn("Accuracy", auc_udf(col("avro_path")))
方案2:input_file_name函数直接调用
input_file_name()仅能返回当前处理的主DataFrame所属的文件路径,无法读取Path列存储的外部目标Avro路径,不符合需求,错误返回结果如下:
+----------+-----+--------------------------+------------------------------------+ |timestamp |Model| Path | different_output_Path | +----------+-----+--------------------------+------------------------------------+ |11:02 |Vgg |projects/Vgg/results.avro |projects/models/all_model_runs.avro | |18:31 |Dnet |projects/Dnet/results.avro|projects/models/all_model_runs.avro | |15:54 |Rnet |projects/Rnet/results.avro|projects/models/all_model_runs.avro| |12:19 |ViT |projects/ViT/results.avro |projects/models/all_model_runs.avro | +----------+-----+-------------------------------------------------------------------+
可行实现方案
禁止在UDF、mapPartitions等分布式执行的算子内部嵌套Spark驱动端的读取操作,可根据数据量选择以下两种稳定实现:
方案1:批量读取后关联(性能最优,生产环境推荐)
核心逻辑是先把所有待读取的Avro路径去重后收集到驱动端,一次性读取所有目标Avro文件,通过input_file_name()标记每条指标对应的来源路径,再和原表做关联,避免逐行读取的性能损耗:
import org.apache.spark.sql.functions._ // 1. 提取所有去重后的待读取Avro路径 val avroPaths = models_df.select("Path").distinct().as[String].collect() // 2. 批量读取所有Avro文件,标记数据来源路径,提取Accuracy指标 val allAvroMetrics = spark.read .format("com.databricks.spark.avro") .load(avroPaths: _*) .withColumn("Path", input_file_name()) .select( col("Path"), col("metrics.Accuracy").as("Accuracy") ) // 3. 和原DataFrame关联得到最终结果 val resultDf = models_df.join(allAvroMetrics, Seq("Path"), "left")
该方案仅触发一次Avro读取作业,适合路径数量多、数据规模大的生产场景
方案2:驱动端循环逐行读取(仅适合小数据量测试场景)
如果待读取的Avro路径数量极少(个位数级别),可以直接在驱动端遍历每一行路径,读取指标后拼接结果:
import org.apache.spark.sql.functions.lit // 遍历原表每一行数据,逐个读取Avro指标后拼接 val resultDf = models_df.collect().foldLeft(models_df.sparkSession.emptyDataFrame) { (tempDf, row) => val avroPath = row.getAs[String]("Path") val timestamp = row.getAs[String]("timestamp") val model = row.getAs[String]("Model") // 读取单个Avro文件提取精度指标 val accuracy = spark.read .format("com.databricks.spark.avro") .load(avroPath) .select(col("metrics.Accuracy")) .first().getAs[Double](0) // 构造单行结果 val singleRowDf = Seq((timestamp, model, avroPath, accuracy)) .toDF("timestamp", "Model", "Path", "Accuracy") if (tempDf.isEmpty) singleRowDf else tempDf.union(singleRowDf) }
该方案会把原表全量数据收集到驱动端内存,路径量大时会触发驱动内存溢出,禁止在生产环境使用
内容的提问来源于stack exchange,提问作者Illustrious Imp
相关产品推荐
相关产品推荐

