Spark案例

wen java案例 3

本文目录导读:

Spark案例

  1. 案例一:WordCount(入门经典)
  2. 案例二:用户行为日志分析(ETL + 聚合)
  3. 案例三:实时流处理(Structured Streaming)
  4. 案例四:机器学习(协同过滤推荐)
  5. 案例五:图计算(PageRank 网页排名)
  6. 性能优化通用技巧(适用于所有案例)

下面我为你整理几个经典的 Apache Spark 案例,从基础到进阶,每个案例都包含场景描述核心代码(Scala/Python)和关键点解析


WordCount(入门经典)

📌 场景

统计文本文件中每个单词出现的次数,这是理解 Spark RDD 和函数式编程的基础。

💻 Python 代码(PySpark)

from pyspark import SparkContext
sc = SparkContext("local", "WordCount")
# 读取文件,每一行变成一个元素
lines = sc.textFile("hdfs:///data/input.txt")
# 扁平化:将每一行拆分成单词
words = lines.flatMap(lambda line: line.split(" "))
# 映射:每个单词变成 (word, 1)
pairs = words.map(lambda word: (word, 1))
# 聚合:相同 key 的 value 相加
counts = pairs.reduceByKey(lambda a, b: a + b)
# 按次数降序排序(可选)
sorted_counts = counts.sortBy(lambda x: x[1], ascending=False)
# 收集结果到 Driver 并打印
for word, count in sorted_counts.collect():
    print(f"{word}: {count}")

🔑 关键点

  • flatMapmap 的区别:flatMap 返回多个元素,map 一对一。
  • reduceByKey 会在分区内先做预聚合(Combiner),减少 Shuffle 数据量。
  • 惰性求值textFilemap 都是 Transformation,只有 collect() 触发计算。

用户行为日志分析(ETL + 聚合)

📌 场景

分析电商平台的用户点击流日志(格式:时间戳, 用户ID, 商品ID, 行为类型(click/buy/cart)),统计每小时各商品的购买量 Top10

💻 Scala 代码

import org.apache.spark.sql.SparkSession
import org.apache.spark.sql.functions._
val spark = SparkSession.builder()
  .appName("LogAnalysis")
  .master("yarn")
  .getOrCreate()
import spark.implicits._
// 1. 读取原始日志
val logs = spark.read.textFile("hdfs:///data/clickstream.log")
  .map { line =>
    val parts = line.split(",")
    (parts(0).toLong, parts(1), parts(2), parts(3)) // (时间戳, 用户ID, 商品ID, 行为)
  }.toDF("ts", "userId", "productId", "action")
// 2. ETL:清洗无效数据 + 添加小时字段
val cleaned = logs
  .filter($"productId".isNotNull && $"ts".isNotNull)
  .withColumn("hour", from_unixtime($"ts" / 1000, "yyyy-MM-dd HH:00"))
// 3. 过滤购买行为,按小时和商品分组统计
val result = cleaned
  .filter($"action" === "buy")
  .groupBy("hour", "productId")
  .agg(count("*").alias("buy_count"))
  .withColumn("rank", row_number().over(
    Window.partitionBy("hour").orderBy($"buy_count".desc)
  ))
  .filter($"rank" <= 10)
// 4. 结果写入 Hive 表
result.write.mode("overwrite").saveAsTable("dwd.product_hourly_top10")

🔑 关键点

  • SQL 窗口函数row_number().over(Window.partitionBy(...)) 实现分组 TopN。
  • DataFrame API 比 RDD 更高效,内置 Catalyst 优化器。
  • 时间处理from_unixtime 将毫秒时间戳转成可读格式。

实时流处理(Structured Streaming)

📌 场景

从 Kafka 读取实时订单流,统计每 5 分钟窗口内各区域的订单金额总和,并输出到 MySQL。

💻 PySpark 代码(Structured Streaming)

from pyspark.sql import SparkSession
from pyspark.sql.functions import from_json, col, window, sum
from pyspark.sql.types import StructType, StructField, StringType, DoubleType
spark = SparkSession.builder \
    .appName("RealtimeOrderAnalysis") \
    .getOrCreate()
# 定义 Kafka 数据 JSON 结构
schema = StructType([
    StructField("order_id", StringType()),
    StructField("region", StringType()),
    StructField("amount", DoubleType()),
    StructField("event_time", StringType())  # "2024-01-01 12:30:00"
])
# 1. 从 Kafka 读取流
df = spark.readStream \
    .format("kafka") \
    .option("kafka.bootstrap.servers", "node1:9092") \
    .option("subscribe", "order_topic") \
    .load() \
    .select(from_json(col("value").cast("string"), schema).alias("data")) \
    .select("data.*")
# 2. 事件时间 + 窗口聚合
windowed = df \
    .withWatermark("event_time", "10 minutes") \  # 允许 10 分钟延迟
    .groupBy(
        col("region"),
        window(col("event_time"), "5 minutes", "5 minutes")  # 5分钟窗口
    ) \
    .agg(sum("amount").alias("total_amount"))
# 3. 输出到 MySQL(foreachBatch 方式支持事务)
def write_to_mysql(batch_df, batch_id):
    batch_df.write \
        .format("jdbc") \
        .option("url", "jdbc:mysql://localhost:3306/realtime") \
        .option("driver", "com.mysql.jdbc.Driver") \
        .option("dbtable", "region_order_stats") \
        .option("user", "root") \
        .option("password", "123456") \
        .mode("append") \
        .save()
query = windowed.writeStream \
    .foreachBatch(write_to_mysql) \
    .outputMode("update") \
    .trigger(processingTime="1 minute") \
    .start()
query.awaitTermination()

🔑 关键点

  • 事件时间 vs 处理时间:用 withWatermark 处理乱序数据。
  • 窗口操作groupBy(window(...)) 实现滚动窗口聚合。
  • 输出模式update 模式只输出更新的结果。

机器学习(协同过滤推荐)

📌 场景

基于用户对电影的评分数据(MovieLens 数据集),使用 ALS 算法训练推荐模型,给指定用户推荐 Top5 电影。

💻 PySpark 代码(MLlib)

from pyspark.ml.recommendation import ALS
from pyspark.ml.evaluation import RegressionEvaluator
from pyspark.sql import SparkSession
spark = SparkSession.builder.appName("MovieRecommend").getOrCreate()
# 1. 加载数据 (userId, movieId, rating, timestamp)
ratings = spark.read.csv("ratings.csv", header=True, inferSchema=True) \
    .select("userId", "movieId", "rating")
# 2. 划分训练集和测试集
(training, test) = ratings.randomSplit([0.8, 0.2])
# 3. 训练 ALS 模型
als = ALS(
    userCol="userId",
    itemCol="movieId",
    ratingCol="rating",
    coldStartStrategy="drop",  # 忽略未知用户/商品
    maxIter=10,
    regParam=0.1
)
model = als.fit(training)
# 4. 评估模型(RMSE)
predictions = model.transform(test)
evaluator = RegressionEvaluator(metricName="rmse", labelCol="rating", predictionCol="prediction")
rmse = evaluator.evaluate(predictions)
print(f"Root-mean-square error = {rmse}")
# 5. 为用户 100 生成 Top5 推荐
user100 = spark.createDataFrame([(100,)], ["userId"])
recommendations = model.recommendForUserSubset(user100, 5)
recommendations.show(truncate=False)

🔑 关键点

  • ALS (交替最小二乘法):适合隐式/显式反馈的协同过滤。
  • 冷启动策略coldStartStrategy="drop" 避免 NaN 预测。
  • 评估指标:RMSE(均方根误差)衡量预测准确度。

图计算(PageRank 网页排名)

📌 场景

使用 GraphX 计算网页之间的 PageRank 值,找出影响力最大的网页节点。

💻 Scala 代码(GraphX)

import org.apache.spark.graphx._
import org.apache.spark.rdd.RDD
val spark = SparkSession.builder().appName("PageRank").getOrCreate()
val sc = spark.sparkContext
// 1. 定义图的顶点和边
val vertices: RDD[(VertexId, String)] = sc.parallelize(Array(
  (1L, "Wikipedia"), (2L, "Google"), (3L, "Baidu"),
  (4L, "Bing"), (5L, "Yahoo")
))
val edges: RDD[Edge[Double]] = sc.parallelize(Array(
  Edge(1L, 2L, 1.0),  // Wikipedia -> Google
  Edge(1L, 3L, 1.0),  // Wikipedia -> Baidu
  Edge(2L, 4L, 1.0),  // Google -> Bing
  Edge(3L, 4L, 1.0),  // Baidu -> Bing
  Edge(4L, 5L, 1.0)   // Bing -> Yahoo
))
val graph = Graph(vertices, edges)
// 2. 运行 PageRank 算法
val ranks = graph.pageRank(0.0001).vertices
// 3. 关联顶点名称,按排名降序排列
val result = ranks.join(vertices).sortBy(_._2._1, ascending = false)
result.collect().foreach { case (id, (rank, name)) =>
  println(s"$name: $rank")
}

🔑 关键点

  • GraphX 专用于图计算,底层基于 RDD。
  • pageRank(tol)tol 是收敛阈值,越小精度越高但耗时更长。
  • 迭代计算:PageRank 本质是迭代求解稳态概率分布。

性能优化通用技巧(适用于所有案例)

优化点 具体做法
避免 Shuffle 使用 mapPartitions 代替 map 频繁操作;合理设置分区数(repartition
使用 Broadcast Join 小表(<100MB)用 broadcast() 广播,避免 SortMergeJoin
压缩与序列化 设置 spark.sql.parquet.compression.codec=snappy;使用 Kryo 序列化
内存调优 spark.memory.fraction=0.8spark.memory.storageFraction=0.5
动态资源 开启 spark.dynamicAllocation.enabled=true 按需分配 Executor

案例 核心技能点 学习价值
WordCount RDD 算子、懒惰求值 入门基础
日志分析 DataFrame、SQL、窗口函数 离线 ETL 实战
流处理 Structured Streaming、Kafka 实时计算
推荐系统 MLlib、ALS 机器学习应用
PageRank GraphX 图计算

案例覆盖了 Spark RDDSQLStreamingMLlibGraphX 五大核心模块,建议你先动手跑通 WordCount,再逐步深入复杂场景,如果有具体某个案例想深入了解细节,可以告诉我!

上一篇Java批处理案例

下一篇Hadoop案例

抱歉,评论功能暂时关闭!