本文目录导读:

- 案例一:WordCount(入门经典)
- 案例二:用户行为日志分析(ETL + 聚合)
- 案例三:实时流处理(Structured Streaming)
- 案例四:机器学习(协同过滤推荐)
- 案例五:图计算(PageRank 网页排名)
- 性能优化通用技巧(适用于所有案例)
下面我为你整理几个经典的 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}")
🔑 关键点
flatMap和map的区别:flatMap返回多个元素,map一对一。reduceByKey会在分区内先做预聚合(Combiner),减少 Shuffle 数据量。- 惰性求值:
textFile和map都是 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.8,spark.memory.storageFraction=0.5 |
| 动态资源 | 开启 spark.dynamicAllocation.enabled=true 按需分配 Executor |
| 案例 | 核心技能点 | 学习价值 |
|---|---|---|
| WordCount | RDD 算子、懒惰求值 | 入门基础 |
| 日志分析 | DataFrame、SQL、窗口函数 | 离线 ETL 实战 |
| 流处理 | Structured Streaming、Kafka | 实时计算 |
| 推荐系统 | MLlib、ALS | 机器学习应用 |
| PageRank | GraphX | 图计算 |
案例覆盖了 Spark RDD、SQL、Streaming、MLlib、GraphX 五大核心模块,建议你先动手跑通 WordCount,再逐步深入复杂场景,如果有具体某个案例想深入了解细节,可以告诉我!