Scala 2.12.x下Spark NLP Evaluation不可用,求Scala版NER评估方法
针对Scala 2.12.x + Spark 3.x的Spark NLP NER评估可行方案
方案1:手动实现NER评估逻辑
Spark NLP的NER输出和标注数据都能转化为Spark DataFrame格式,你可以手动计算精确率、召回率、F1值这些核心指标,步骤如下:
- 加载标注好的测试数据集,确保结构和模型输出对齐,比如包含
document、token、true_label列 - 运行NER模型得到预测结果,提取
token和pred_label列 - 给标注数据和预测数据添加索引列,保证文本token的顺序一致后关联两表
- 遍历每个token的真实标签和预测标签,统计TP(真正例)、FP(假正例)、FN(假负例)的数量
- 基于统计值计算指标:
- 精确率 = TP / (TP + FP)
- 召回率 = TP / (TP + FN)
- F1值 = 2 * (精确率 * 召回率) / (精确率 + 召回率)
示例代码片段:
import org.apache.spark.sql.functions._ // 假设标注数据df_label含token、true_label列,预测数据df_pred含token、pred_label列 val indexedLabel = df_label.withColumn("idx", monotonically_increasing_id()) val indexedPred = df_pred.withColumn("idx", monotonically_increasing_id()) val mergedDf = indexedLabel.join(indexedPred, Seq("idx", "token"), "inner") val metrics = mergedDf.select( when(col("true_label") === col("pred_label") && col("true_label") =!= "O", 1).otherwise(0).alias("tp"), when(col("true_label") === "O" && col("pred_label") =!= "O", 1).otherwise(0).alias("fp"), when(col("true_label") =!= "O" && col("pred_label") === "O", 1).otherwise(0).alias("fn") ).agg( sum("tp").alias("total_tp"), sum("fp").alias("total_fp"), sum("fn").alias("total_fn") ).collect()(0) val tp = metrics.getAs[Long]("total_tp") val fp = metrics.getAs[Long]("total_fp") val fn = metrics.getAs[Long]("total_fn") val precision = if (tp + fp > 0) tp.toDouble / (tp + fp) else 0.0 val recall = if (tp + fn > 0) tp.toDouble / (tp + fn) else 0.0 val f1 = if (precision + recall > 0) 2 * (precision * recall) / (precision + recall) else 0.0 println(s"Precision: $precision, Recall: $recall, F1: $f1")
方案2:利用Spark NLP的Annotation对比能力
Spark NLP的Annotation类型自带操作API,可以直接对比真实标注和模型预测的Annotation结果,步骤如下:
- 将标注数据转化为
Annotation格式的true_ner列,和模型输出的pred_ner列结构保持一致 - 自定义UDF对比两组Annotation,统计匹配情况
- 基于统计结果计算评估指标
示例代码片段:
import com.johnsnowlabs.nlp.Annotation import org.apache.spark.sql.functions._ // 假设数据集df含true_ner(真实标注的Annotation数组)和pred_ner(模型预测的Annotation数组)列 val compareUdf = udf((trueNer: Seq[Annotation], predNer: Seq[Annotation]) => { val trueEntities = trueNer.filter(_.result != "O").map(ann => (ann.begin, ann.end, ann.result)).toSet val predEntities = predNer.filter(_.result != "O").map(ann => (ann.begin, ann.end, ann.result)).toSet val tp = trueEntities.intersect(predEntities).size val fp = predEntities.diff(trueEntities).size val fn = trueEntities.diff(predEntities).size (tp, fp, fn) }) val metricsDf = df.withColumn("metrics", compareUdf(col("true_ner"), col("pred_ner"))) .agg( sum("metrics._1").alias("total_tp"), sum("metrics._2").alias("total_fp"), sum("metrics._3").alias("total_fn") ) // 后续计算指标逻辑同方案1
方案3:升级Spark NLP版本
检查Spark NLP的最新版本,后续发布的spark-nlp-eval模块可能已经支持Scala 2.12,替换sbt中的依赖版本尝试:
libraryDependencies += "com.johnsnowlabs.nlp" %% "spark-nlp-eval" % "最新版本号"
内容的提问来源于stack exchange,提问作者Faaiz
相关产品推荐
相关产品推荐

