DataFrame编程
DataFrame 编程:创建 → 转换 → 输出,三步走
DataFrame 就是带 schema 的分布式表:
| 概念 | 类比 |
|---|---|
| DataFrame | 分布式表 |
| Row | 一行记录 |
| Column | 一列数据(有名称和类型) |
| Schema | 表结构(列名 + 类型) |
DataFrame 编程的核心流程就三步:
- 创建:从文件/表/RDD 创建 DataFrame
- 转换:用
select、filter、groupBy、join处理数据 - 输出:
show()、collect()、write输出结果
记住:DataFrame 的操作跟 SQL 思维一致——选列、过滤行、分组聚合、关联表。
创建 DataFrame:从文件、表、RDD 来
1.1 从文件读
val spark = SparkSession.builder().appName("MyApp").master("local[*]").getOrCreate()
import spark.implicits._
// 读 JSON
val df = spark.read.json("people.json")
// 读 CSV
val df = spark.read
.option("header", "true")
.option("inferSchema", "true")
.csv("people.csv")
// 读 Parquet(Spark 默认格式,推荐)
val df = spark.read.parquet("people.parquet")
1.2 从 RDD 转
// 方式1:RDD + toDF(最简单)
val rdd = sc.parallelize(Seq((1, "Alice"), (2, "Bob")))
val df = rdd.toDF("id", "name")
// 方式2:样例类(自动推断列名)
case class Person(id: Int, name: String)
val df = sc.parallelize(Seq(Person(1, "Alice"), Person(2, "Bob"))).toDF()
1.3 从 Hive 表读
val df = spark.table("hive_db.user_table")
查看数据:先看一眼长什么样
| 方法 | 作用 |
|---|---|
df.show() |
打印前 20 行(默认) |
df.show(10) |
打印前 10 行 |
df.printSchema() |
打印表结构(列名 + 类型) |
df.columns |
获取所有列名 |
df.count() |
总行数 |
df.describe().show() |
统计摘要(count、mean、stddev) |
df.printSchema()
// root
// |-- id: integer (nullable = false)
// |-- name: string (nullable = true)
// |-- age: integer (nullable = true)
转换操作:选列、过滤、聚合、关联
3.1 选列(SELECT)
// 选一列
df.select("name").show()
// 选多列
df.select("name", "age").show()
// 用 $ 语法做表达式
df.select($"name", $"age" + 1).show()
// 列重命名
df.select($"name".as("username")).show()
3.2 过滤行(WHERE / FILTER)
// 字符串表达式(类似 SQL)
df.filter("age > 18").show()
df.where("age > 18 AND gender = 'F'").show()
// 用 $ 语法
df.filter($"age" > 18).show()
3.3 排序(ORDER BY)
// 升序
df.sort("age").show()
// 降序
df.sort($"age".desc).show()
3.4 分组聚合(GROUP BY)
import org.apache.spark.sql.functions._
df.groupBy("gender")
.agg(
count("*").as("total"),
avg("age").as("avg_age"),
max("salary").as("max_salary")
)
.show()
3.5 表关联(JOIN)
val employees = spark.read.parquet("employees")
val depts = spark.read.parquet("departments")
// 内连接
employees.join(depts, "dept_id").show()
// 左连接
employees.join(depts, Seq("dept_id"), "left_outer").show()
// 多条件连接
employees.join(
depts,
employees("dept_id") === depts("id") && employees("year") === depts("year")
).show()
3.6 集合操作
df1.union(df2) // 并集(不去重)
df1.union(df2).distinct() // 并集(去重)
df1.intersect(df2) // 交集
df1.except(df2) // 差集(df1 有 df2 没有)
用 SQL 操作 DataFrame
把 DataFrame 注册成临时表,然后写 SQL:
df.createOrReplaceTempView("people")
spark.sql("""
SELECT gender, AVG(age) AS avg_age, COUNT(*) AS cnt
FROM people
WHERE age > 18
GROUP BY gender
HAVING cnt > 10
ORDER BY avg_age DESC
""").show()
临时表的生命周期:
| 方法 | 生命周期 |
|---|---|
createTempView |
当前 SparkSession |
createGlobalTempView |
跨 SparkSession(需加 global_temp. 前缀) |
输出:写文件、转 RDD、拿结果
5.1 写文件
// 写 Parquet(推荐)
df.write.parquet("output.parquet")
// 写 CSV
df.write
.option("header", "true")
.csv("output.csv")
// 写 JSON
df.write.json("output.json")
写入模式(.mode()):
| 模式 | 行为 |
|---|---|
append |
追加 |
overwrite |
覆盖 |
ignore |
存在就忽略 |
error |
存在就报错(默认) |
5.2 转 RDD
val rdd = df.rdd // RDD[Row]
rdd.foreach(row => {
val id = row.getInt(0)
val name = row.getString(1)
})
5.3 收集结果到 Driver
// 收集全部(小数据)
df.collect()
df.collectAsList()
// 取前 N 行
df.take(10)