大家好,我是jobleap.cn的小九。
PySpark 是 Apache Spark 的 Python 接口,凭借分布式计算能力成为大数据处理的核心工具。本教程从环境搭建到综合实战,系统讲解 PySpark 高频 API 的用法,并通过串联案例让你掌握端到端的数据处理流程。

一、环境准备与核心入口

1. 安装 PySpark

# 安装核心包(建议搭配Hadoop依赖,适配本地/集群环境)
pip install pyspark[sql]

2. 初始化 SparkSession(核心入口)

SparkSession 是 PySpark 2.0+ 操作 DataFrame/Dataset 的统一入口,所有数据处理都围绕它展开。

from pyspark.sql import SparkSession
from pyspark.sql.functions import (
    col, lit, concat, when, avg, sum, count, countDistinct,
    row_number, rank, dense_rank, udf
)
from pyspark.sql.window import Window
from pyspark.sql.types import (
    StringType, IntegerType, FloatType, StructType, StructField
)

# 初始化 SparkSession
spark = SparkSession.builder \
    .appName("PySpark常用API实战") \
    .master("local[*]")  # 本地模式,*表示使用所有CPU核心(集群环境需改为yarn/spark://)
    .config("spark.sql.adaptive.enabled", "true")  # 自适应执行优化
    .getOrCreate()

# 关闭日志冗余输出(可选)
spark.sparkContext.setLogLevel("WARN")

二、数据创建:生成 DataFrame(核心数据结构)

DataFrame 是 PySpark 最常用的分布式数据结构(类似关系型数据库表),以下是 3 种高频创建方式:

1. 从Python集合创建(测试常用)

# 方式1:列表+列名
data = [("Alice", 25, "F", 85.5),
        ("Bob", 30, "M", 90.0),
        ("Charlie", 28, "M", 88.5),
        ("Diana", 22, "F", 92.0)]
columns = ["name", "age", "gender", "score"]
df = spark.createDataFrame(data, schema=columns)

# 方式2:指定结构化schema(更严谨,避免类型自动推断错误)
schema = StructType([
    StructField("name", StringType(), nullable=False),
    StructField("age", IntegerType(), nullable=True),
    StructField("gender", StringType(), nullable=True),
    StructField("score", FloatType(), nullable=True)
])
df = spark.createDataFrame(data, schema=schema)

# 查看数据基本信息
df.printSchema()  # 打印字段类型
df.show(2)  # 显示前2行(默认前20行)

2. 从外部文件读取(生产常用)

支持 CSV/JSON/Parquet/Excel 等格式,重点演示高频的 CSV 和 Parquet:

# 读取CSV(带表头、指定编码)
df_csv = spark.read \
    .option("header", "true")  # 第一行作为列名
    .option("encoding", "utf-8") \
    .option("sep", ",")  # 分隔符
    .option("nullValue", "NA")  # 把NA识别为空值
    .csv("data/student.csv")

# 读取Parquet(列式存储,效率更高,Spark默认格式)
df_parquet = spark.read.parquet("data/student.parquet")

# 读取JSON(单行JSON格式)
df_json = spark.read.json("data/student.json")

三、数据转换:核心操作 API

转换操作是 Lazy 执行(仅当行动操作触发时才计算),以下是最常用的转换 API:

1. 字段选择与重命名(select/withColumnRenamed)

# 选择指定字段
df_select = df.select("name", "age", col("score") * 1.1)  # 字段计算后保留匿名列

# 字段重命名(两种方式)
df_rename1 = df.select(col("name"), col("age"), col("score").alias("final_score"))
df_rename2 = df.withColumnRenamed("score", "final_score")

2. 新增字段(withColumn)

# 方式1:基于现有字段计算
df_add = df.withColumn("score_level", 
                       when(col("score") >= 90, "A")
                       .when(col("score") >= 85, "B")
                       .otherwise("C"))

# 方式2:添加常量字段
df_add_const = df.withColumn("class", lit("Class_1"))  # lit()表示常量

# 方式3:字符串拼接
df_add_concat = df.withColumn("name_gender", concat(col("name"), lit("_"), col("gender")))

3. 过滤与条件筛选(filter/where)

filterwhere 功能完全一致,仅语法风格不同:

# 单条件筛选
df_filter1 = df.filter(col("age") > 25)
df_filter2 = df.where("score >= 88")  # 支持SQL表达式

# 多条件筛选(&表示且,|表示或,~表示非,需用括号包裹)
df_filter_multi = df.filter(
    (col("gender") == "F") & (col("score") > 90)
)

4. 去重与排序(distinct/orderBy)

# 去重(全字段)
df_distinct = df.distinct()

# 按指定字段去重(保留第一条)
df_drop_dup = df.dropDuplicates(["name"])

# 排序(asc升序,desc降序)
df_sort = df.orderBy(col("age").asc(), col("score").desc())

# 限制返回行数(常与排序配合)
df_limit = df_sort.limit(3)  # 取排序后前3行

5. 类型转换(cast)

解决读取文件时字段类型推断错误的问题:

df_cast = df.withColumn("age", col("age").cast(IntegerType())) \
            .withColumn("score", col("score").cast(FloatType()))

6. 缺失值处理(na 相关 API)

数据清洗核心,处理空值/缺失值:

# 方式1:删除含空值的行(how="any"任意列空,"all"所有列空)
df_drop_na = df.na.drop(how="any", subset=["age", "score"])  # 仅检查age/score列

# 方式2:填充空值(按字段指定填充值)
df_fill_na = df.na.fill({
    "age": 0,
    "gender": "Unknown",
    "score": 80.0
})

# 方式3:替换指定值(把特定值转为空值/其他值)
df_replace_na = df.na.replace("M", "Male", subset=["gender"])

四、行动操作:触发计算的 API

行动操作是 Eager 执行(立即计算),返回结果到驱动节点,常用如下:

# 1. 显示数据(最常用)
df.show()  # 格式化显示
df.show(truncate=False)  # 不截断长字符串

# 2. 统计行数
count = df.count()
print(f"总行数:{count}")

# 3. 获取前N行数据(返回Python列表)
first_row = df.first()  # 第一行
top3_rows = df.take(3)  # 前3行
all_rows = df.collect()  # 所有行(慎用!数据量大时会撑爆驱动节点内存)

# 4. 数值型字段统计(count/mean/stddev/min/max)
df.describe(["age", "score"]).show()  # 基础统计
df.summary("min", "25%", "50%", "75%", "max").show()  # 分位数统计

# 5. 字段求和/均值(直接计算)
total_score = df.select(sum(col("score"))).first()[0]
avg_age = df.select(avg(col("age"))).first()[0]

五、聚合操作:分组统计 API

基于 groupBy 实现分组聚合,是数据分析核心:

# 1. 基础分组聚合
df_group = df.groupBy("gender") \
             .agg(
                 count("name").alias("person_count"),  # 计数
                 avg("age").alias("avg_age"),  # 平均年龄
                 sum("score").alias("total_score"),  # 总分
                 max("score").alias("max_score"),  # 最高分
                 countDistinct("age").alias("distinct_age_count")  # 去重计数
             )
df_group.show()

# 2. 分组后过滤(having 效果)
df_group_having = df.groupBy("gender") \
                    .agg(avg("score").alias("avg_score")) \
                    .filter(col("avg_score") > 88)  # 过滤分组结果

六、窗口函数:复杂排名/累计统计 API

窗口函数(Window)用于行与行之间的计算(如排名、累计求和),是 PySpark 高级操作的核心:

# 1. 定义窗口规则(按gender分区,按score降序排序)
window_spec = Window.partitionBy("gender").orderBy(col("score").desc())

# 2. 常用窗口函数
df_window = df.withColumn("row_num", row_number().over(window_spec))  # 行号(不重复)
df_window = df_window.withColumn("rank_num", rank().over(window_spec))  # 排名(有并列,跳号)
df_window = df_window.withColumn("dense_rank_num", dense_rank().over(window_spec))  # 密集排名(有并列,不跳号)

# 3. 累计求和(滑动窗口)
window_spec_sum = Window.partitionBy("gender").orderBy("age").rowsBetween(Window.unboundedPreceding, Window.currentRow)
df_window = df_window.withColumn("cumulative_score", sum("score").over(window_spec_sum))

df_window.show()

七、连接操作:多表关联 API

类似 SQL 的 JOIN,支持内连接/左连接/右连接/全连接:

# 准备关联表(成绩表 + 班级表)
class_data = [("Alice", "Class_A"), ("Bob", "Class_B"), ("Eve", "Class_C")]
class_df = spark.createDataFrame(class_data, ["name", "class"])

# 1. 内连接(仅保留匹配行)
inner_join = df.join(class_df, on="name", how="inner")

# 2. 左连接(保留左表所有行,右表无匹配则为空)
left_join = df.join(class_df, on="name", how="left")

# 3. 右连接(保留右表所有行,左表无匹配则为空)
right_join = df.join(class_df, on="name", how="right")

# 4. 全连接(保留所有行,无匹配则为空)
full_join = df.join(class_df, on="name", how="full")

# 5. 多字段关联
# df.join(other_df, on=["name", "age"], how="inner")

left_join.show()

八、性能优化:缓存与广播变量

1. 缓存(cache/persist)

重复使用的 DataFrame 缓存到内存/磁盘,避免重复计算:

# 缓存(默认MEMORY_ONLY,内存不足则溢出)
df.cache()  # 等价于 df.persist()

# 自定义存储级别(MEMORY_AND_DISK:内存不足时写入磁盘)
from pyspark import StorageLevel
df.persist(StorageLevel.MEMORY_AND_DISK)

# 用完后释放缓存(避免内存泄漏)
df.unpersist()

2. 广播变量(broadcast)

小表(<10GB)广播到所有Executor节点,避免Shuffle,提升JOIN性能:

# 广播小表(class_df是小表)
broadcast_class = spark.sparkContext.broadcast(class_df.collect())

# 基于广播变量创建DataFrame
broadcast_df = spark.createDataFrame(broadcast_class.value)

# JOIN时自动使用广播(或显式指定broadcast())
from pyspark.sql.functions import broadcast
df.join(broadcast(class_df), on="name", how="left").show()

九、自定义函数(UDF)

内置函数满足不了需求时,自定义UDF处理数据:

# 1. 定义普通函数
def age_category(age):
    if age < 25:
        return "Young"
    elif age < 30:
        return "Middle"
    else:
        return "Old"

# 2. 注册UDF(指定返回类型)
age_udf = udf(age_category, StringType())

# 3. 使用UDF
df_udf = df.withColumn("age_category", age_udf(col("age")))
df_udf.show()

# (进阶)注册为SQL函数(可直接在SQL中调用)
spark.udf.register("age_category_udf", age_category, StringType())
df.createOrReplaceTempView("student")  # 创建临时视图
spark.sql("SELECT name, age, age_category_udf(age) FROM student").show()

十、数据写入:输出到外部存储

处理完成后写入文件/数据库,常用格式如下:

# 1. 写入CSV(覆盖已有文件、按gender分区)
df.write \
    .mode("overwrite")  # 模式:overwrite/append/ignore/error(默认)
    .option("header", "true") \
    .partitionBy("gender")  # 按字段分区存储(提升读取效率)
    .csv("output/student_csv")

# 2. 写入Parquet(推荐,压缩比高)
df.write \
    .mode("append") \
    .parquet("output/student_parquet")

# 3. 写入MySQL(需添加JDBC驱动)
df.write \
    .format("jdbc") \
    .option("url", "jdbc:mysql://localhost:3306/test") \
    .option("dbtable", "student") \
    .option("user", "root") \
    .option("password", "123456") \
    .option("driver", "com.mysql.cj.jdbc.Driver") \
    .mode("overwrite") \
    .save()

十一、综合实战:串联所有常用 API

需求:分析学生成绩数据,完成以下步骤:

  1. 读取CSV数据并清洗(处理空值、类型转换);
  2. 新增成绩等级、年龄分类字段;
  3. 按性别分组统计平均分、人数;
  4. 按性别排名(取各性别前2名);
  5. 关联班级表,补充班级信息;
  6. 缓存结果并写入Parquet文件。
# 步骤1:读取并清洗数据
student_df = spark.read \
    .option("header", "true") \
    .option("nullValue", "NA") \
    .csv("data/student_score.csv") \
    .withColumn("age", col("age").cast(IntegerType())) \
    .withColumn("score", col("score").cast(FloatType())) \
    .na.fill({"age": 0, "score": 80.0, "gender": "Unknown"})  # 填充空值

# 步骤2:新增字段
# 成绩等级
student_df = student_df.withColumn("score_level", 
                                   when(col("score") >= 90, "A")
                                   .when(col("score") >= 85, "B")
                                   .otherwise("C"))
# 年龄分类(UDF)
age_udf = udf(lambda x: "Young" if x <25 else "Middle" if x<30 else "Old", StringType())
student_df = student_df.withColumn("age_category", age_udf(col("age")))

# 步骤3:分组统计
gender_stats = student_df.groupBy("gender") \
                         .agg(
                             avg("score").alias("avg_score"),
                             count("name").alias("total_person")
                         )

# 步骤4:窗口排名(取各性别前2名)
window_spec = Window.partitionBy("gender").orderBy(col("score").desc())
student_ranked = student_df.withColumn("rank", row_number().over(window_spec)) \
                           .filter(col("rank") <= 2)

# 步骤5:关联班级表
class_df = spark.createDataFrame([("Alice", "Class_A"), ("Bob", "Class_B"), ("Diana", "Class_C")], ["name", "class"])
student_final = student_ranked.join(broadcast(class_df), on="name", how="left")

# 步骤6:缓存并写入文件
student_final.cache()
# 显示最终结果
student_final.show(truncate=False)
# 写入Parquet
student_final.write \
    .mode("overwrite") \
    .partitionBy("gender") \
    .parquet("output/student_final.parquet")

# 释放缓存
student_final.unpersist()

# 关闭SparkSession
spark.stop()

十二、核心注意事项

  1. Lazy执行:转换操作仅记录逻辑,行动操作才触发计算;
  2. 避免collect():大数据集使用collect()会将所有数据拉到驱动节点,易OOM;
  3. 分区优化:数据倾斜时需重分区(repartition)或合并分区(coalesce);
  4. UDF性能:尽量使用内置函数(Spark优化过),UDF需序列化,性能较低;
  5. 资源配置:集群环境需合理设置executor内存、核数,避免资源浪费。

本教程覆盖了PySpark 90%以上的常用API,重点在于理解DataFrame的分布式特性和API的组合使用。建议结合实际业务场景反复练习,掌握性能优化技巧,即可应对大部分大数据处理需求。

更多推荐