习题 12:文本分类完整案例(读取 CSV + JSON)

题目:读取 user_data.csv(用户文本特征)和 label_data.json(标签),join 后训练逻辑回归文本分类模型。

user_data.csv

userId,user_intro
user_1,医生,专注于心脏病学研究
user_2,技术爱好者,专注于人工智能和机器学习

label_data.json

[
  {
    "userId": "user_1",
    "label": "medical"
  },
  {
    "userId": "user_2",
    "label": "finance"
  }
]
import org.apache.spark.ml.Pipeline
import org.apache.spark.ml.classification.LogisticRegression
import org.apache.spark.ml.evaluation.MulticlassClassificationEvaluator
import org.apache.spark.ml.feature.{HashingTF, Tokenizer}
import org.apache.spark.sql.SparkSession

object SparkTextClassifier {
  def main(args: Array[String]): Unit = {
    val spark = SparkSession.builder()
      .master("local[*]")
      .appName("SparkTextClassifier")
      .getOrCreate()

    // 关键:导入隐式转换,使 $"col" 语法生效
    import spark.implicits._

    // 读取训练数据
    val trainUserDF = spark.read.option("header", "true").csv("user_data.csv")
    val trainLabelDF = spark.read.json("label_data.json")

    // 合并并一次性完成列选取和类型转换
    val trainDF = trainUserDF.join(trainLabelDF, "userId")
      .selectExpr(
        "user_intro as text",
        "cast(label as double) as label"
      )

    // 构建 Pipeline
    val tokenizer = new Tokenizer()
      .setInputCol("text")
      .setOutputCol("words")
    val hashingTF = new HashingTF()
      .setNumFeatures(1000)
      .setInputCol(tokenizer.getOutputCol)
      .setOutputCol("features")
    val lr = new LogisticRegression()
      .setMaxIter(10)
      .setRegParam(0.01)

    val pipeline = new Pipeline()
      .setStages(Array(tokenizer, hashingTF, lr))

    val model = pipeline.fit(trainDF)

    // 定义测试集:从训练集中随机拆分 20% 作为测试集
    val Array(trainingData, testData) = trainDF.randomSplit(Array(0.8, 0.2), seed = 42L)
    val predictions = model.transform(testData)

    predictions.select("text", "label", "probability", "prediction").show(false)

    val evaluator = new MulticlassClassificationEvaluator()
      .setLabelCol("label")
      .setPredictionCol("prediction")
      .setMetricName("accuracy")

    println(f"Test Error = ${1.0 - evaluator.evaluate(predictions)}%2.2f")
  }
}