DataFrame编程

DataFrame 编程:创建 → 转换 → 输出,三步走

DataFrame 就是带 schema 的分布式表:

概念 类比
DataFrame 分布式表
Row 一行记录
Column 一列数据(有名称和类型)
Schema 表结构(列名 + 类型)

DataFrame 编程的核心流程就三步:

  1. 创建:从文件/表/RDD 创建 DataFrame
  2. 转换:用 select、filter、groupBy、join 处理数据
  3. 输出: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)