一、开始
a. pyspark
i. dos 命令行下输入 pyspark
命令行下默认启用的是 yarn-client 模式,jar 包无法分发到各节点,使用 jar 包注册的 udf 可能会报错。
ii. 集成环境
spark-submit 配置:
PYSPARK_DRIVER_PYTHON=/opt/cloudera/parcels/Anaconda/bin/python3 pyspark --conf spark.port.maxRetries=1000
jupyter:
PYSPARK_DRIVER_PYTHON=/opt/cloudera/parcels/Anaconda/bin/jupyter-notebook PYSPARK_DRIVER_PYTHON_OPTS="--NotebookApp.open_browser=False --NotebookApp.ip='*' --NotebookApp.port=10229" pyspark --conf spark.port.maxRetries=1000
ipython:
PYSPARK_DRIVER_PYTHON=/opt/cloudera/parcels/Anaconda/bin/ipython3 pyspark --conf spark.port.maxRetries=1000
参数说明:
PYSPARK_PYTHON:用于指定启动 pyspark 使用的 python 版本,指定特定版本 python 时必须配置PYSPARK_DRIVER_PYTHON:用于驱动 pyspark 的 python,可以是 ipython 或者 jupyterPYSPARK_DRIVER_PYTHON_OPTS:为驱动 python 提供启动参数
补充:
# 命令
jupyter-notebook --NotebookApp.open_browser=False --NotebookApp.ip='*' --NotebookApp.port=10229
# 参数
PYSPARK_DRIVER_PYTHON=/opt/cloudera/parcels/Anaconda/bin/jupyter-notebook PYSPARK_DRIVER_PYTHON_OPTS="--NotebookApp.open_browser=False --NotebookApp.ip='*' --NotebookApp.port=10229"
# 极端例子
PYSPARK_DRIVER_PYTHON=ls PYSPARK_DRIVER_PYTHON_OPTS="-a" pyspark
iii. 引入依赖
启动时引入:
pyspark --jars /path/myjar1.jar,/path/myjar2.jar
启动后引入:
from pyspark import SparkContext
# 初始化 SparkSession
spark = SparkSession.builder.appName("lijp").enableHiveSupport().getOrCreate()
# 初始化 SparkContext
sc = spark.sparkContext
# 添加依赖的 JAR 包
sc.addPyFile("/path/to/hanlp.jar")
iv. 调用 jar 包
注意绑定时的包路径必须精确到类名、对象名,而不能是包名。
注意:个别方法的调用与 scala 中有差异,scala 中可能会因为隐式进行自动类型转换。
# 若 RuleEncryptUtil 是普通类,则需要实例化
ruleEncryptUtil = sc._jvm.RuleEncryptUtil()
# 若 RuleEncryptUtil 是静态类,则既可以实例化,也可以直接调用
ruleEncryptUtil = sc._jvm.RuleEncryptUtil() # 调用静态类的构造函数生成一个实例
ruleEncryptUtil = sc._jvm.RuleEncryptUtil # 调用静态类现有的实例
# 如果没有自动导入 jar 包,可以将类(构造函数)绑定给 sc._jvm
from py4j.java_gateway import java_import
java_import(sc._jvm, "com.fibodt.encrypt.RuleEncryptUtil")
ruleEncryptUtil = sc._jvm.RuleEncryptUtil() # 实例化
# 如果没有自动导入 jar 包,也可以将对象绑定给 sc._jvm
from py4j.java_gateway import java_import
java_import(sc._jvm, "com.fibodt.encrypt.RuleEncryptUtil")
ruleEncryptUtil = sc._jvm.RuleEncryptUtil # 这里的 RuleEncryptUtil 是一个实例
v. 自动加载内容
from pyspark.sql import SparkSession
spark = SparkSession.builder.appName("lijp").enableHiveSupport().getOrCreate()
from pyspark import SparkContext
sc = spark.sparkContext
vi. 显示日志级别
sc = pyspark.SparkContext()
sc.setLogLevel("ERROR")
vii. 定义 spark 和 sc
定义 spark:
spark = SparkSession.builder.appName("Word Count") \
.master("local[*]") \
.config("spark.driver.memory", "32g") \
.config("spark.executor.memory", "64g") \
.config("spark.executor.cores", "6") \
.getOrCreate()
定义 sc:
注意:由于 pyspark 没有隐式转换特性,因此没有
toDF接口,如需进行转化,需要使用map(lambda p: Row(字段名=p[索引]))。
sc = pyspark.SparkContext()
常用接口 col、when 等:
from pyspark.sql.functions import *
viii. 定义 fs
conf = spark._jvm.org.apache.hadoop.conf.Configuration()
fs = spark._jvm.org.apache.hadoop.fs.FileSystem.get(conf)
二、快速入门
a. 文件读取
textFile = sc.textFile("/Users/lijp/IdeaProjects/sc/src/main/java/sc/汽车品牌.csv")
# textFile 为 RDD 类型,具有 List 的很多相似操作,可以进行循环遍历,例如 map、foreach、filter 等
b. map 操作
对 rdd 中每行进行处理。
c. flatmap 操作
对 rdd 中每行进行展开处理。
d. collect 操作
将结果转换为 Array 类型。
e. cache 操作
将 rdd 和 dataset 保存在内存,被 session 持有。
三、RDD 编程指引
a. 创建 RDD 集合
可以将 rdd 看做是 spark 分布式环境下的 list。
data = [1, 2, 3, 4, 5]
distData = sc.parallelize(data, 5)
# distData 类型为 ParallelCollectionRDD,且分片数为 5
b. 读取文件
- 若读取本地文件,本地文件需要在所有节点上可以被访问到
- 所有读取文件的方法都支持在目录上、通配符、压缩包上运行:
sc.textFile("/my/directory")sc.textFile("/my/directory/*.txt")sc.textFile("/my/directory/*.gz")
- 控制返回文件数量,通常情况下返回文件为一个文件夹下的多个文件,可以使用
SparkContext.wholeTextFiles控制返回文件的个数,例如返回一个文件 SparkContext.sequenceFile[Int, String]SparkContext.hadoopRDDSparkContext.objectFile
c. RDD 操作
i. 转换(transform):生成新的 RDD
-
map:返回一个新的分布式数据集,该数据集是通过将源的每个元素传递给函数 func 形成的
-
mapValues:(标准形式元组可用)返回一个新的分布式数据集,同 map 相似,mapValues 在 (K,V) 对的数据集上调用,仅对 V 进行操作
-
filter:返回一个新的数据集,保留 func 返回 true 的元素
-
flatmap:对一个 array-like 的字段进行处理,先进行 map 生成一条记录,然后 flatten 为多条记录
-
mapPartitions:与 map 相似,但是分别在 RDD 的每个分区(块)上运行,因此 func 在类型 T 的 RDD 上运行时必须为
Iterator<T> => Iterator<U>类型 -
mapPartitionsWithIndex:与 mapPartitions 相似,但它还为 func 提供表示分区索引的整数值,因此当在类型 T 的 RDD 上运行时,func 必须为
(Int, Iterator<T>) => Iterator<U>类型 -
sample:使用给定的随机数发生器的种子进行抽样,共三个参数:
WithReplacement为true表示有放回抽样,原数据集大小不变- 为
false表示无放回抽样,原数据集在抽样后减少百分比 fraction表示抽样比例seed表示随机数种子,Long 型整数,例如12345L
-
sampleBy(colName: String, fractions: Map[String, Double], seed: Long):对数据集进行抽样,colName 是列名,fractions 是该列中各值的抽样概率,可以按照设定比例从各个值中抽取样本数据
-
union:返回一个新的数据集,其中包含源数据集中的元素的并集
-
intersection:返回一个新的 RDD,其中包含源数据集中的元素的交集
-
distinct:返回一个新的数据集,其中包含源数据集的不同元素
-
groupByKey:在 (K, V) 对的数据集上调用时,返回 (K, Iterable<V>) 对的数据集。
注意:如果要分组以便对每个键执行聚合(例如求和或平均值),则使用
reduceByKey或aggregateByKey将产生更好的性能。 注意:默认情况下,输出中的并行度取决于父 RDD 的分区数。您可以传递一个可选numPartitions参数来设置不同数量的任务。 -
reduceByKey:在 (K, V) 对的数据集上调用时,返回 (K, V) 对的数据集,其中每个键的值使用给定的 reduce 函数 func 进行汇总。与
groupByKey一样,reduce 任务的数量可以通过可选的第二个参数配置。 -
aggregateByKey:在 (K, V) 对的数据集上调用时,返回 (K, U) 对的数据集,其中每个键的值使用给定的 Combine 函数和中性的”零”值进行汇总。允许与输入值类型不同的聚合值类型,同时避免不必要的分配。与
groupByKey一样,reduce 任务的数量可以通过可选的第二个参数配置。 -
sortByKey:在由 K 实现 Ordered 的 (K, V) 对的数据集上调用时,返回 (K, V) 对的数据集,按布尔值
ascending指定按键以升序或降序排序。 -
join:在 (K, V) 和 (K, W) 类型的数据集上调用时,返回 (K, (V, W)) 对的数据集,其中每个键都有所有成对的元素。外连接通过
leftOuterJoin、rightOuterJoin和fullOuterJoin支持。join 之前最好确认 rdd 中元素的类型,防止出现 Any 类型,导致报错:
but class RDD is invariant in type T. You may wish to define T as +T instead. -
cogroup:在 (K, V) 和 (K, W) 类型的数据集上调用时,返回 (K, (Iterable<V>, Iterable<W>)) 元组的数据集。此操作也称为
groupWith。 -
cartesian:在类型 T 和 U 的数据集上调用时,返回 (T, U) 对(所有元素对)的数据集。
-
pipe:通过外壳命令(例如 Perl 或 bash 脚本)通过管道传输 RDD 的每个分区。将 RDD 元素写入进程的 stdin,并将输出到其 stdout 的行作为字符串的 RDD 返回。
-
coalesce:将 RDD 中的分区数减少到
numPartitions。筛选大型数据集后,对于更有效地运行操作很有用。 -
repartition:随机重排 RDD 中的数据以创建更多或更少的分区,并在整个分区之间保持平衡。这始终会拖曳网络上的所有数据。
repartition(1):重排 RDD 中的数据,合并为一个分区repartition(col("colName")):重排 RDD 中的数据,根据指定列的记录进行分区
-
repartitionAndSortWithinPartitions:根据给定的分区程序对 RDD 重新分区,并在每个结果分区中,按其键对记录进行排序。这比
repartition在每个分区内调用然后排序更为有效,因为它可以将排序推入洗牌机制。
ii. 行动(action):汇总所有结果返回驱动程序
-
reduce:使用函数 func(该函数接受两个参数并返回一个)来聚合数据集的元素。该函数应该是可交换的和关联的,以便可以并行正确地计算它。
-
collect:在驱动程序中将数据集的所有元素作为数组返回。这通常在返回足够小的数据子集的过滤器或其他操作之后很有用。
-
count:返回数据集中的元素数。
-
first:返回数据集的第一个元素(类似于
take(1))。 -
take:返回数据集的前 n 个元素的数组。
-
takeSample:返回一个数组,该数组包含数据集 num 个元素的随机样本(是否替换),可以选择预先指定随机数生成器种子。
-
takeOrdered:使用自然顺序或自定义比较器返回 RDD 的前 n 个元素。
-
saveAsTextFile:将数据集的元素以文本文件(或文本文件集)的形式写入本地文件系统、HDFS 或任何其他 Hadoop 支持的文件系统中的给定目录中。Spark 将在每个元素上调用
toString,以将其转换为文件中的一行文本。 -
saveAsSequenceFile:在本地文件系统、HDFS 或任何其他 Hadoop 支持的文件系统的给定路径中,将数据集的元素作为 Hadoop SequenceFile 写入。这在实现 Hadoop 的 Writable 接口的键/值对的 RDD 上可用。
-
saveAsObjectFile:使用 Java 序列化以简单的格式编写数据集的元素,然后可以使用
SparkContext.objectFile()加载。 -
countByKey:仅在类型 (K, V) 的 RDD 上可用。返回 (K, Int) 对的哈希图以及每个键的计数。
-
foreach:在数据集的每个元素上运行函数 func。通常这样做是出于副作用,例如更新累加器或与外部存储系统进行交互。
注意:在之外修改除累加器以外的变量
foreach()可能会导致不确定的行为。
iii. 缓存
- persist:可以根据参数进行不同级别的缓存:
MEMORY_ONLY、MEMORY_AND_DISK、MEMORY_ONLY_SER、MEMORY_AND_DISK_SER、DISK_ONLY - cache:默认缓存级别
MEMORY_ONLY - 缓存级别选择:
MEMORY_ONLY>MEMORY_ONLY_SER>MEMORY_AND_DISK - unpersist:释放缓存
iv. 打印部分记录
- collect:将全部记录汇总到一台机器上,可能会耗尽内存
- take:获取部分记录
v. 共享变量
广播变量: 在所有节点上创建一个只读变量,在使用时不应该调用函数中的指定变量值,而是直接使用指定广播变量,而且防止修改节点上的广播变量。
dataFrame 和变量都可以使用 broadcast 进行广播,但是 rdd 不可以。
broadcastVar = sc.broadcast([1, 2, 3])
broadcastDF = functions.broadcast(df)
累加器:
创建累加器:
accum = sc.accumulator(0) # 累加器的初值为 0
sc.parallelize([1, 2, 3, 4]).foreach(lambda x: accum.add(x))
print(accum.value)
构造累加器(继承 AccumulatorParam,实现 zero 和 addInPlace 方法):
from pyspark import AccumulatorParam
class IntAccumulatorParam(AccumulatorParam):
# zero 用于初始化,sc.accumulator(0, IntAccumulatorParam()) 会调用 zero,并给 initialValue 赋值为 0
def zero(self, initialValue):
return initialValue
# addInPlace 用于计算累计值,每次调用 accum.add 方法,会调用 addInPlace 进行处理
# v1 的类型与 initialValue 的类型一致,调用 addInPlace 时,会将 v1 赋值为 self.value
# v2 的类型与 rdd 中的元素一致,调用 addInPlace 时,将 v2 赋值为将要累加的新的值
def addInPlace(self, v1, v2):
v1 += v2
return v1
accum = sc.accumulator(0, IntAccumulatorParam())
sc.parallelize([1, 2, 3, 4]).foreach(lambda x: accum.add(x))
print(accum.value)
注意:
- 累加器务必在 action 算子中使用,当 action 算子被重复调用时,累加器仅计算一次
- 如果在 transform 算子中使用,可能因惰性求值、或重复调用,而导致结果错误
四、SparkSQL、DataSets、DataFrames
a. 读取文件
df = spark.read.json("examples/src/main/resources/people.json")
df = spark.read.csv("file:///D:/java_workspace/fun_test.csv")
支持的数据源:
spark.read.jdbcspark.read.jsonspark.read.orcspark.read.parquetspark.read.textFilespark.read.format("com.databricks.spark.avro").load("/raw_data/operator/8/dpi_result_fp/p_biz={e_17,e_20}")
使用 option 设置参数:
spark.read.option("header", "true").csv("/Users/lijp/IdeaProjects/testspark/src/main/scala/sample.txt")
b. 显示数据
df.show() # 如需禁用截断显示,指定 truncate=False
df.printSchema()
c. 选择数据
仅选择:
df.select("name").show()
选择并计算:
df.select("age", df["age"] + 1).show()
过滤:
df.filter("age > 18").show()
df.filter(f"age > {value}").show()
df.filter(df["age"] > 18).show()
补充——对整数类型过滤:
# 逻辑运算符:>, <, ==, !=
df.filter(df["num"] == 2)
df.filter(df["num"] > 2)
df.filter(df["num"] < 2)
# 或者
df.filter("num = 2")
df.filter("num > 2")
df.filter("num < 2")
# 传递参数过滤
ind = 2
df.filter(df["num"] == ind)
df.filter(df["num"] > ind)
df.filter(df["num"] < ind)
补充——对字符串过滤:
# equalTo
df.filter(df["id"].equalTo("a"))
# 传递参数过滤
str_val = "a"
df.filter(df["id"].equalTo(str_val))
# 当 dataframe 没有字段名时,可以用默认的字段名 [_1, _2, ...] 来进行判断
多条件判断: 逻辑连接符 &(并)、|(或)、~(非),多条件判断时,每个单独条件需要加小括号,~ 在括号外。
df.filter((df["num"] == 2) & (df["id"].equalTo("a")))
df.filter((df["num"] == 1) | (df["num"] == 3))
df.filter(~(df["num"] == 1))
df.where(df["num"] != 1)
df.where(~(df["num"] == 1))
df.where((df["num"] > 0) & ~(df["num"] == 1))
df.where(~isnull(df["null"])) # pyspark 环境下没有 col(colName).isNotNull,使用 ~isnull(df[colName])
d. na 处理
df.na.drop() # 丢弃含有 na 的行
df.na.drop(thresh=2) # 丢弃少于两个值的行
df.na.drop(how='all') # 丢弃全为 na 的行
df.na.drop(how='any') # 丢弃有一个 na 的行
df.na.drop(subset=['Sales']) # 仅针对 sales 列进行丢弃
df.na.fill(0) # 用 0 填充
df.na.fill(value="no label", subset=["label"]) # 仅针对 label 列进行填充
e. RDD-数据聚合操作
分组计数:
ds.select("tag_code", "rule").groupBy("tag_code").count().show()
分组后求最值、平均值、求和:
peopleDF.groupBy("address").max("age").show()
peopleDF.groupBy("address").avg("age").show()
peopleDF.groupBy("address").min("age").show()
peopleDF.groupBy("address").sum("age").show()
指定字段的数据类型(withColumn,指定类型不区分大小写):
peopleDF.withColumn("count", ds1.col("count").cast("double")).groupBy("address").max("age").show()
分组后求多个聚合值(使用 groupBy + agg):
peopleDF.groupBy("address").agg(count("age"), max("age"), min("age"), avg("age"), sum("age")).show()
分组聚合后取别名:
注意:pyspark 中聚合函数无法使用
as(),只能用alias()重命名。
peopleDF.groupBy("address").agg(count("age").alias("cnt"), avg("age").alias("avg")).show()
分组后行转列(pivot):
peopleDF.groupBy("address").pivot("name").avg("age").show()
peopleDF.groupBy("address").pivot("name").agg(countDistinct("IdCard").alias("uv")).orderBy(col("address")).show()
直接求 count、max、min(groupBy 中不传值):
peopleDF.groupBy().avg("age").show()
f. SQL 操作
注册临时表:
df.createOrReplaceTempView("people") # 临时表由当前 sparkSession 持有,Session 消失则临时表销毁
注册全局表:
df.createGlobalTempView("people") # 全局表由所有 session 共享,可以在新的 session 中继续使用
执行 SQL:
sqlDF = spark.sql("SELECT * FROM people")
Join 操作:
# 连接(支持 inner、left、right、all)
tag.join(stat1, tag["tag_code"] == stat1["tag"], "inner")
# 多次连接
tag.join(stat1, tag["tag_code"] == stat1["tag"], "left") \
.join(stat2, tag["tag_code"] == stat2["tag"], "left")
# 选择命名唯一的列、或者在计算过程中生成的列,可以使用 col
from pyspark.sql.functions import col
tag.select(col("tag"))
# 选择(新视图中存在多个重名列,需要根据 DF 名称进行选择,此时写入文件会报错)
tag.select(tag["tag"], stat1["tag"])
# 选择多列
tag.select(tag["first_category"], tag["second_category"], tag["tag"], tag["tag_code"])
tag.select(tag["*"])
# 选择并重命名
tag.select(tag["tag"].alias("tag1"), stat1["tag"].alias("tag2"))
# 排序
tag.sort("first_category")
# 多列排序
tag.sort("first_category", "second_category", "tag", "tag_code", "cover", "cover_ratio")
若使用 DataSet 进行 join 操作出现了重复列,可以使用列名对不同的 DataSet 进行索引。 结果显示:为防止较长的列名尾部出现省略号,可以使用
df.show(truncate=False)。
g. 创建 RDD
从文件创建: 调用 sc.textFile,按行解析为 rdd。
fileRdd = sc.textFile("/Users/lijp/test.txt")
从集合创建: 调用 sc.parallelize,按元素解析为 rdd。
arrayRdd = sc.parallelize([(1, 1), (2, 2)])
h. 创建 DataFrame
i. List + toDF
使用 List[Tuple] 包装每行记录,结合 toDF 接口,转化为 DataFrame。
df = spark.createDataFrame([
["ming", 20, 15552211521],
["hong", 19, 13287994007],
["zhi", 21, 15552211523]
]).toDF("name", "age", "phone")
ii. Dict
使用 dict 包装每行记录,转化为 DataFrame。
df = spark.createDataFrame({
"name": ["ming", "hong", "zhi"],
"age": [20, 19, 21],
"phone": [15552211521, 13287994007, 15552211523]
})
iii. Row
使用 Row 包装每行记录,转化为 DataFrame。
注意:pyspark 中的 Row 不支持 scala 中的
getAs[String]方法,需要用asDict获取。
df = spark.createDataFrame([
Row(name="ming", age=20, phone=15552211521),
Row(name="hong", age=19, phone=15552211521),
Row(name="zhi", age=21, phone=15552211523)
])
iv. pandas
使用 pandas 创建 df,转化为 DataFrame。
pandas_df = pd.DataFrame({'a': [1, 2], 'b': [3, 4], 'c': ['string1', 'string1']})
df = spark.createDataFrame(pandas_df)
v. RDD + StructType(推荐)
from pyspark.sql import Row
from pyspark.sql.types import *
testRdd = sc.parallelize([[1, 1], [2, 2]]).map(lambda line: Row(line[0], line[1]))
schema = StructType([
StructField("id", IntegerType(), nullable=True),
StructField("code", IntegerType(), nullable=True)
])
df = spark.createDataFrame(testRdd, schema)
vi. RDD + StructType(单元素)
单个元素构成一行记录,使用 Row()。
testRdd = sc.parallelize([Row(1), Row(2), Row(3)])
schema = StructType([StructField("id", IntegerType(), nullable=True)])
df = spark.createDataFrame(testRdd, schema)
vii. RDD + StructType(多元素 + Row 工厂)
Row(字段名1, 字段名2) 构造一个工厂对象。
pyspark 中没有 scala 环境下的
Row.fromSeq()接口。
Person = Row("id", "name", "score")
testRdd = sc.parallelize([
Person(*[1, "lijp", 99.0]),
Person(*[2, "zhangs", 85.5]),
Person(*[3, "lis", 60.0])
])
df = spark.createDataFrame(testRdd)
viii. list + tuple(pyspark 特性)
使用 list 和 tuple 封装数据,schema 只传入字段名即可,字段类型根据传入数据类型进行推断。
df = spark.createDataFrame([(1, "lijp"), (2, "sunj"), (3, "unknow")], ["id", "name"])
ix. List + map + class
class Test:
def __init__(self, Field1, Field2):
self.Field1 = Field1
self.Field2 = Field2
df = spark.createDataFrame(map(lambda x: Test(x[0], x[1]), [(1, 1), (2, 2), (3, 3)]))
x. RDD + map + class
class Test:
def __init__(self, Field1, Field2):
self.Field1 = Field1
self.Field2 = Field2
testRdd = sc.parallelize([(1, 1), (2, 2)])
df = spark.createDataFrame(testRdd.map(lambda line: Test(line[0], line[1])))
i. RDD 转化为 DataFrame
由 sc 转化为 DF:
class Line:
def __init__(self, uid, tag, time, freq):
self.uid = uid
self.tag = tag
self.time = time
self.freq = freq
tagDF = sc.textFile(path) \
.map(lambda line: line.split("|")) \
.map(lambda line: Line(line[0], line[1], line[2], line[3])) \
.toDF("uid", "tag", "time", "freq")
DF 按行操作:
对比 spark-shell,pyspark 下 DataFrame 没有 map 方法,需要转 rdd 处理。
根据索引操作:
rdd_ind = tagDF.rdd.map(lambda tag: "tagcode: " + str(tag[2])).take(1)
根据键操作:
rdd_key = tagDF.rdd.map(lambda tag: "tagcode: " + str(tag.asDict()["tag"])).take(1)
rdd 转 DF 需要通过 schema 参数显式指定 DataType:
# 如果直接使用 toDF(colName) 会报错 Can not infer schema for type
spark.createDataFrame(rdd_key, schema=StructType([
StructField("tagcode", StringType(), nullable=True)
]))
DataFrame 操作大全:
| 操作 | 说明 | 示例 |
|---|---|---|
show(10, False) | 显示十行数据,不省略末尾部分 | df.show(10, False) |
collect() | 收集所有结果数据,返回 Array | df.collect() |
collectAsList() | 收集所有结果数据,返回 List | df.collectAsList() |
describe() | 显示指定字段的描述统计信息 | df.describe("user").show() |
first() / head() / take() / takeAsList() | 获取头部数据 | df.first() |
where() | 传入 SQL 字符串筛选 | df.where("user=1 or type='助手1'").show() |
filter() | 传入 SQL 字符串筛选 | df.filter("user=1 or type='助手1'").show() |
like() | 模糊匹配(需与 where 一起用) | df.where(col("uid").like("%lijp%")).show() |
isin() | 判断值是否在列表中 | this.where(col("uid").isin(*list(mapped))) |
select() | 获取指定字段 | df.select("user", "type").show() |
selectExpr() | 传入 UDF、函数、as 别名 | df.selectExpr("user", "type as visittype", "to_date(visittime)").show() |
col() / apply() | 获取指定字段 | df.col("user") 等同于 df("user") |
limit() | 获取前 n 行记录(非 action 操作) | df.limit(10) |
orderBy() / sort() | 排序(desc 降序) | df.orderBy("visittime").show() / df.orderBy(df("visittime").desc).show() |
groupBy() | 数据分组 | df.groupBy("user").count().show() |
distinct() | 去重 | df.distinct() |
dropDuplicates() | 根据指定字段去重 | df.dropDuplicates(["type"]).show() |
drop() | 去除指定字段 | df.drop("type").show() |
agg() | 聚合操作 | df.groupBy("user").agg(max("id"), sum("user")).show() |
withColumn() | 添加/覆盖列 | df.withColumn("sex", df["user"] % 2).show() |
join() | 连接 | df.join(df2, ["id"]) / df.join(df2, df["id"] == df2["id"]) |
between() | 范围筛选(包含两端边界) | df.where(col("age").between(18, 30)) |
isin 在 pyspark 中的操作步骤:
# 1. collect,将对应 dataframe 作为 List[Row] 取回
t = that.select("uid").distinct().collect()
# 2. map,从每一行 Row 中使用 asDict.get("列名")
mapped = map(lambda row: row.asDict().get("uid"), t)
# 3. 列表解析元素逐个处理(类似 scala 的 :_*)
this.where(col("uid").isin(*list(mapped)))
添加常数列:
from pyspark.sql.functions import lit
newdf = df.withColumn("newcol", lit("myval"))
j. 聚合函数
i. 不带类型的 UDAF
继承 UserDefinedAggregateFunction,传入的是数值。
from pyspark.sql.types import *
# 实现自定义聚合函数(Scala 示例)
# object AverageUserDefinedAggregateFunction extends UserDefinedAggregateFunction {
# // 聚合函数的输入数据结构
# override def inputSchema: StructType = StructType(StructField("input", LongType) :: Nil)
# // 缓存区数据结构
# override def bufferSchema: StructType = StructType(StructField("sum", LongType) :: StructField("count", LongType) :: Nil)
# // 聚合函数返回值数据结构
# override def dataType: DataType = DoubleType
# // 聚合函数是否是幂等的,即相同输入是否总是能得到相同输出
# override def deterministic: Boolean = true
# // 初始化缓冲区
# override def initialize(buffer: MutableAggregationBuffer): Unit = {
# buffer(0) = 0L
# buffer(1) = 0L
# }
# // 给聚合函数传入一条新数据进行处理
# override def update(buffer: MutableAggregationBuffer, input: Row): Unit = {
# if (input.isNullAt(0)) return
# buffer(0) = buffer.getLong(0) + input.getLong(0)
# buffer(1) = buffer.getLong(1) + 1
# }
# // 合并聚合函数缓冲区
# override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
# buffer1(0) = buffer1.getLong(0) + buffer2.getLong(0)
# buffer1(1) = buffer1.getLong(1) + buffer2.getLong(1)
# }
# // 计算最终结果
# override def evaluate(buffer: Row): Any = buffer.getLong(0).toDouble / buffer.getLong(1)
# }
# 使用
# spark.udf.register("u_avg", AverageUserDefinedAggregateFunction)
# spark.sql("select count(1) as count, u_avg(age) as avg_age from v_user").show()
# spark.sql("select sex, count(1) as count, u_avg(age) as avg_age from v_user group by sex").show()
ii. 带有类型的 UDAF
继承 Aggregator,传入的是样例类。
# Scala 示例
# object AverageAggregator extends Aggregator[User, Average, Double] {
# // 初始化 buffer
# override def zero: Average = Average(0L, 0L)
# // 处理一条新的记录
# override def reduce(b: Average, a: User): Average = {
# b.sum += a.age
# b.count += 1L
# b
# }
# // 合并聚合 buffer
# override def merge(b1: Average, b2: Average): Average = {
# b1.sum += b2.sum
# b1.count += b2.count
# b1
# }
# // 减少中间数据传输
# override def finish(reduction: Average): Double = reduction.sum.toDouble / reduction.count
# override def bufferEncoder: Encoder[Average] = Encoders.product
# // 最终输出结果的类型
# override def outputEncoder: Encoder[Double] = Encoders.scalaDouble
# }
# case class Average(var sum: Long, var count: Long)
# case class User(id: Long, name: String, sex: String, age: Long)
# 使用
# val user = spark.read.json("data/user").as[User]
# user.select(AverageAggregator.toColumn.name("avg")).show()
k. 常用数据源加载、保存方法
i. 通用加载方法
# parquet
usersDF = spark.read.format("parquet").load("examples/src/main/resources/users.parquet")
# json
peopleDF = spark.read.format("json").load("examples/src/main/resources/people.json")
# csv
peopleDFCsv = spark.read.format("csv") \
.option("sep", ";") \
.option("inferSchema", "true") \
.option("header", "true") \
.load("examples/src/main/resources/people.csv")
# 直接运行 SQL
sqlDF = spark.sql("SELECT * FROM parquet.`examples/src/main/resources/users.parquet`")
ii. 保存方法
df.select("<sql语句>").write.format("<输出格式>").mode("<保存类型>").save("<保存路径>")
保存模式:
| 模式 | 说明 |
|---|---|
error(默认) | 保存时,若指定路径上数据文件已存在,则仅报错并退出 |
append | 追加模式,若指定路径上数据文件已存在,则将 DF 内容追加到数据文件末尾 |
overwrite | 覆写模式,若指定路径上数据文件已存在,则将 DF 内容覆写到数据文件 |
ignore | 忽略模式,若指定路径上数据文件已存在,则将 DF 内容丢弃,不对数据文件做任何改动(类似于 CREATE TABLE IF NOT EXISTS) |
df.select("name").write.format("parquet").save("file:///root/data/overwrite")
df.select("name").write.format("parquet").mode("append").save("file:///root/data/overwrite")
df.select("name").write.format("parquet").mode("overwrite").save("file:///root/data/overwrite")
df.select("name").write.format("parquet").mode("ignore").save("file:///root/data/overwrite")
补充——jupyter 保存 CSV 文件问题:
jupyter 下使用
write.csv()进行保存的时候,生成的文件里每行首末都会有一个双引号",但是 spark-submit 提交使用write.csv()的时候,没有双引号。衍生问题:jupyter 中如果使用
write.option("quote", "").csv()保存文件,那么每行首末会有特殊空字符,直接在文件里查看是看不到的,但是使用 vim 打开时会看到^@。如果每行首末都有^@的情况下,使用spark.read.csv打开这个文件会报空指针错误,但是使用spark.read.text就没有问题。解决方案:
- 使用 spark-submit 提交 py 文件,直接使用
write.csv()即可,生成的文件里没有双引号- 使用 spark-shell 也可以,但是必须是 scala 交互环境,直接使用
write.csv(),生成的文件里也没有双引号- 如果使用 pyspark 进入 python 交互命令行,使用
write.csv(),生成的文件里仍然会有双引号
iii. Option 参数信息
Parquet:
mergedDF = spark.read.option("mergeSchema", "true").parquet("data/test_table")
JSON:
mdf = spark.read.option("multiline", "true").json("multi.json")
mdf = spark.read.option("charset", "UTF-16BE").json("fileInUTF16.json")
CSV:
peopleDFCsv = spark.read.option("sep", "|").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("encoding", "UTF-8").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("quote", "'").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("escape", "\\").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("charToEscapeQuoteEscaping", "escape").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("comment", "").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("header", "true").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("ignoreLeadingWhiteSpace", "false").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("ignoreTrailingWhiteSpace", "false").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("nullValue", "").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("emptyValue", "").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("nanValue", "NaN").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("positiveInf", "Inf").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("negativeInf", "-Inf").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("dateFormat", "yyyy-MM-dd'T'HH:mm:ss.SSSXXX").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("maxColumns", "20480").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("maxCharsPerColumn", "-1").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("mode", "PERMISSIVE").csv("path/to/file.csv")
peopleDFCsv = spark.read.option("multiLine", "false").csv("path/to/file.csv")
补充:
- 写入 csv 文件时,有时候需要剔除字段两侧的引号,理论上讲可以使用
option("quote", "")将引号设置为空字符串,但是写入文件后发现引号部分为\00。 - 因此,需要剔除字段两侧的引号时,可以使用
option("quote", "\b"),\b的含义表示字符串的边界,实际上也可以用来标识字符串的边界,只是这种标识用的是空字符串,骗过了 spark 对 quote 传入参数的检查。 - 如果只是为了让空值不显示,可以使用
option("emptyValue", ""),即可让结果显示由...,"",...改变为...,,...。
Text:
df = spark.read.option("wholetext", "false").text("/path/to/spark/README.md") # 是否根据 \n 识别多行记录
l. 分桶、排序、分区
示例:
peopleDF.write.bucketBy(42, "name").sortBy("age").saveAsTable("fileout")
分桶: 根据值取哈希。
据说目前在使用
bucketBy时,必须和sortBy、saveAsTable一起使用!
usersDF.write.bucketBy(42, "name").saveAsTable("fileout")
排序: 根据值从大到小排序,会生成一个文件。
peopleDF.write.sortBy("age").saveAsTable("fileout")
分区: 根据值进行分组,会生成多个文件。
usersDF.write.partitionBy("favorite_color").format("parquet").save("fileout.parquet")
m. 占位符
# 空的 DataSet
spark.emptyDataSet
# 空的 DataFrame
spark.emptyDataFrame
# 空的 RDD
sc.emptyRDD
五、特殊操作
a. 函数操作生成新列
使用 spark 中的操作完成行转列(case when then else):
from pyspark.sql.functions import *
i. 聚合函数
| 函数 | 说明 |
|---|---|
avg | 平均值 |
collect_list | 聚合指定字段的值到 list |
collect_set | 聚合指定字段的值到 set |
corr | 计算两列的 Pearson 相关系数 |
count | 计数 |
countDistinct | 去重计数(SQL 中用法:select count(distinct class)) |
covar_pop | 总体协方差(population covariance) |
covar_samp | 样本协方差(sample covariance) |
first | 分组第一个元素 |
last | 分组最后一个元素 |
grouping | grouping_id |
kurtosis | 计算峰态(kurtosis)值 |
skewness | 计算偏度(skewness) |
max | 最大值 |
min | 最小值 |
mean | 平均值 |
stddev | 即 stddev_samp |
stddev_samp | 样本标准偏差(sample standard deviation) |
stddev_pop | 总体标准偏差(population standard deviation) |
sum | 求和 |
sumDistinct | 非重复值求和(SQL 中用法:select sum(distinct class)) |
var_pop | 总体方差(population variance) |
var_samp | 样本无偏方差(unbiased variance) |
variance | 即 var_samp |
ii. 集合函数
-
取值:如果一列为 array 类型,可以直接使用
(0)来进行取值 -
array_contains(column, value):检查 array 类型字段是否包含指定元素 -
array_position(column, value):返回 array 类型字段中第一个指定元素的索引,如果为 0 则表示无此元素 -
explode:展开 array 或 map 为多行 -
explode_outer:同 explode,但当 array 或 map 为空或 null 时,会展开为 null -
posexplode:同 explode,带位置索引 -
posexplode_outer:同 explode_outer,带位置索引 -
from_json:解析 JSON 字符串为 StructType 或 ArrayType,有多种参数形式 -
to_json:转为 json 字符串,支持 StructType、ArrayType of StructTypes、MapType 或 ArrayType of MapTypes -
get_json_object(column, path):获取指定 json 路径的 json 对象字符串select get_json_object('{"a":1,"b":2}','$.a'); -
json_tuple(column, fields):获取 json 中指定字段值select json_tuple('{"a":1,"b":2}','a','b'); -
map_keys:返回 map 的键组成的 array -
map_values:返回 map 的值组成的 array -
size:array 或 map 的长度,需要结合withColumn使用 -
sort_array(e: Column, asc: Boolean):将 array 中元素排序(自然排序),默认 asc
iii. 时间函数
| 函数 | 说明 |
|---|---|
add_months(startDate, numMonths) | 指定日期添加 n 月 |
date_add(start, days) | 指定日期之后 n 天,如 select date_add('2018-01-01', 3) |
date_sub(start, days) | 指定日期之前 n 天 |
datediff(end, start) | 两日期间隔天数 |
current_date() | 当前日期,日期格式为 yyyy-MM-dd |
current_timestamp() | 当前时间戳,TimestampType 类型 |
date_format(dateExpr, format) | 日期格式化 |
dayofmonth(e) | 日期在一月中的天数,支持 date/timestamp/string |
dayofyear(e) | 日期在一年中的天数,支持 date/timestamp/string |
weekofyear(e) | 日期在一年中的周数,支持 date/timestamp/string |
from_unixtime(ut, f) | 时间戳转字符串格式 |
from_utc_timestamp(ts, tz) | 时间戳转指定时区时间戳 |
to_utc_timestamp(ts, tz) | 指定时区时间戳转 UTC 时间戳 |
hour(e) | 提取小时值 |
minute(e) | 提取分钟值 |
month(e) | 提取月份值 |
quarter(e) | 提取季度 |
second(e) | 提取秒 |
year(e) | 提取年 |
last_day(e) | 指定日期的月末日期 |
months_between(date1, date2) | 计算两日期差几个月 |
next_day(date, dayOfWeek) | 计算指定日期之后的下一个周一、二…,dayOfWeek 区分大小写,只接受 "Mon", "Tue", "Wed", "Thu", "Fri", "Sat", "Sun" |
to_date(e, format) | 字段类型转为 DateType |
trunc(date, format) | 日期截断 |
unix_timestamp(s, p) | 指定格式的时间字符串转时间戳 |
unix_timestamp(s) | 同上,默认格式为 yyyy-MM-dd HH:mm:ss |
unix_timestamp() | 当前时间戳(秒),底层实现为 unix_timestamp(current_timestamp(), yyyy-MM-dd HH:mm:ss) |
window(timeColumn, windowDuration, slideDuration, startTime) | 时间窗口函数,将指定时间(TimestampType)划分到窗口 |
to_date 与 date_format 的区别:
to_date只负责解析字符串,但是不能进行格式转换。例如to_date(lit("20200101"), "yyyyMMdd")可以运行,但是to_date(lit("20200101"), "yyyy-MM-dd")会报错date_format只负责格式转换,但是不能进行类型转化。例如date_format(to_date("20200101", "yyyyMMdd"), "yyyy-MM-dd"),date_format(lit("20200101"), "yyyyMMdd")返回 null
数学函数:
| 函数 | 说明 |
|---|---|
cos, sin, tan | 计算角度的余弦、正弦、正切 |
sinh, tanh, cosh | 计算双曲正弦、正切、余弦 |
acos, asin, atan, atan2 | 计算余弦/正弦值对应的角度 |
bin | 将 long 类型转为对应二进制数值的字符串,如 bin("12") 返回 "1100" |
bround | 舍入,使用 Decimal 的 HALF_EVEN 模式(v>0.5 向上舍入,v<0.5 向下舍入,v=0.5 向最近的偶数舍入) |
round(e, scale) | HALF_UP 模式舍入到 scale 位小数(v>=0.5 向上舍入,v<0.5 向下舍入,即四舍五入) |
ceil | 向上舍入 |
floor | 向下舍入 |
conv(num, fromBase, toBase) | 转换数值(字符串)的进制 |
log(base, a) | \( \log_{base}(a) \) |
log(a) | \( \log_e(a) \) |
log10(a) | \( \log_{10}(a) \) |
log2(a) | \( \log_2(a) \) |
log1p(a) | \( \log_e(a+1) \) |
pmod(dividend, divisor) | 返回 dividend mod divisor 的正值 |
pow(l, r) | \( r^l \)(支持 Column、Double 参数) |
radians(e) | 角度转弧度 |
rint(e) | 返回最接近参数的 double 值且等于数学整数 |
shiftLeft(e, numBits) | 向左位移 |
shiftRight(e, numBits) | 向右位移 |
shiftRightUnsigned(e, numBits) | 向右位移(无符号位) |
signum(e) | 返回数值正负符号 |
sqrt(e) | 平方根 |
hex(column) | 转十六进制 |
unhex(column) | 逆转十六进制 |
iv. 混杂(misc)函数
| 函数 | 说明 |
|---|---|
crc32(e) | 计算 CRC32,返回 bigint |
hash(cols*) | 计算 hash code,返回 int |
md5(e) | 计算 MD5 摘要,返回 32 位 16 进制字符串 |
sha1(e) | 计算 SHA-1 摘要,返回 40 位 16 进制字符串 |
sha2(e, numBits) | 计算 SHA 摘要,返回 numBits 位 16 进制字符串。numBits 支持 224, 256, 384, 512 |
v. 其他非聚合函数
| 函数 | 说明 |
|---|---|
abs(e) | 绝对值 |
array(cols*) | 多列合并为 array,cols 必须为同类型 |
map(cols*) | 将多列组织为 map,输入列必须为 (key, value) 形式,各列的 key/value 分别为同一类型 |
bitwiseNOT(e) | 按位取反 |
broadcast[T](df) | 将 df 变量广播,用于实现 broadcast join,如 left.join(broadcast(right), "joinKey") |
coalesce(e*) | 返回第一个非空值 |
col(colName) | 返回 colName 对应的 Column |
column(colName) | col 函数的别名 |
expr(expr) | 解析 expr 表达式,将返回值存于 Column,并返回这个 Column |
greatest(exprs*) | 返回多列中的最大值,跳过 Null |
least(exprs*) | 返回多列中的最小值,跳过 Null |
input_file_name() | 返回当前任务的文件名 |
isnan(e) | 检查是否 NaN(非数值) |
isnull(e) | 检查是否为 Null |
lit(literal) | 将字面量(literal)创建一个 Column |
typedLit[T](literal) | 将字面量创建一个 Column,literal 支持 scala types(如 List, Seq, Map) |
monotonically_increasing_id() | 返回单调递增唯一 ID,但不同分区的 ID 不连续。ID 为 64 位整型 |
nanvl(col1, col2) | col1 为 NaN 则返回 col2 |
negate(e) | 负数,同 df.select(-df("amount")) |
not(e) | 取反,同 df.filter(!df("isActive")) |
rand() | 随机数 [0.0, 1.0] |
rand(seed) | 随机数 [0.0, 1.0],使用 seed 种子 |
randn() | 随机数,从正态分布取 |
randn(seed) | 同上,使用 seed 种子 |
spark_partition_id() | 返回 partition ID |
struct(cols*) | 多列组合成新的 struct column |
when(condition, value) | 当 condition 为 true 返回 value |
when 示例:
people.select(
when(people("gender") == "male", 0)
.when(people("gender") == "female", 1)
.otherwise(2)
)
# 如果没有 otherwise 且 condition 全部没命中,则返回 null
vi. 排序函数
| 函数 | 说明 |
|---|---|
asc(columnName) | 正序 |
asc_nulls_first(columnName) | 正序,null 排最前 |
asc_nulls_last(columnName) | 正序,null 排最后 |
desc(columnName) | 倒序 |
desc_nulls_first(columnName) | 倒序,null 排最前 |
desc_nulls_last(columnName) | 倒序,null 排最后 |
# 排序示例
df.sort(asc("dept"), desc("age"))
# 或者
df.orderBy(asc(df["dept"]), desc(df["age"]))
vii. 字符串函数
| 函数 | 说明 |
|---|---|
ascii(e) | 计算第一个字符的 ascii 码 |
base64(e) | base64 转码 |
unbase64(e) | base64 解码 |
concat(exprs*) | 连接多列字符串 |
concat_ws(sep, exprs*) | 使用 sep 作为分隔符连接多列字符串 |
decode(value, charset) | 解码 |
encode(value, charset) | 转码,charset 支持 'US-ASCII', 'ISO-8859-1', 'UTF-8', 'UTF-16BE', 'UTF-16LE', 'UTF-16' |
format_number(x, d) | 格式化 '#,###,###.##' 形式的字符串 |
format_string(format, arguments*) | 将 arguments 按 format 格式化,格式为 printf-style |
initcap(e) | 单词首字母大写 |
lower(e) | 转小写 |
upper(e) | 转大写 |
instr(str, substring) | substring 在 str 中第一次出现的位置 |
length(e) | 字符串长度 |
levenshtein(l, r) | 计算两个字符串之间的编辑距离(Levenshtein distance) |
locate(substr, str) | substring 在 str 中第一次出现的位置,位置编号从 1 开始,0 表示未找到 |
locate(substr, str, pos) | 同上,但从 pos 位置后查找 |
lpad(str, len, pad) | 字符串左填充,用 pad 字符填充 str 的字符串至 len 长度 |
rpad(str, len, pad) | 字符串右填充 |
ltrim(e) | 剪掉左边的空格、空白字符 |
rtrim(e) | 剪掉右边的空格、空白字符 |
ltrim(e, trimString) | 剪掉左边的指定字符 |
rtrim(e, trimString) | 剪掉右边的指定字符 |
trim(e, trimString) | 剪掉左右两边的指定字符 |
trim(e) | 剪掉左右两边的空格、空白字符 |
regexp_extract(e, exp, groupIdx) | 正则提取匹配的组 |
regexp_replace(e, pattern, replacement) | 正则替换匹配的部分 |
repeat(str, n) | 将 str 重复 n 次返回 |
reverse(str) | 将 str 反转 |
soundex(e) | 计算桑迪克斯代码(soundex code),用于按英语发音来索引姓名 |
split(str, pattern) | 用 pattern 分割 str,若想按索引使用返回的 Array 结果,需要使用 getItem(index) |
substring(str, pos, len) | 在 str 上截取从 pos 位置开始长度为 len 的子字符串 |
substring_index(str, delim, count) | 按分隔符 delim 截取子串(count 为正从左数,为负从右数,区分大小写) |
translate(src, matchingString, replaceString) | 把 src 中的 matchingString 全换成 replaceString |
此外,
col(colName)也支持一些有用的函数:isNotNull、startswith、endswith、substr等。
viii. 特殊类型字段操作
Array<Struct<uid:Int, rating:Float>> 类型字段操作(参考 20200909 相关内容)。
ix. UDF 函数(user-defined function)
1) SparkSQL 环境下注册并使用 UDF:
from pyspark.sql import *
df = spark.createDataFrame([("id1", 1), ("id2", 4), ("id3", 5)], ["id", "value"])
spark.udf.register("simpleUDF", lambda v: v * v)
df.registerTempTable("df")
spark.sql("select * simpleUDF(sum) from df").show(truncate=False)
2) functions 下的 udf 注册 UDF 函数:
# 使用 udf 装饰器,实际上是当调用 test 前,将 test 句柄传入 udf 装饰器,处理为 udf 函数
# 使用 udf 装饰器时,不支持 lambda 函数定义
@udf(returnType=IntegerType())
def test(x):
return x * x
df.select(test(col("value")).alias("test")).show(truncate=False)
# 使用 udf 接口
test = lambda x: x * x
test_udf = udf(test, IntegerType())
df.select(test_udf(col("value")).alias("test")).show(truncate=False)
3) 不支持序列化的对象在 pyspark 中不能和 UDF 配合使用:
Py4j 报错,scala 中可能可以。
解决:将不支持序列化的对象的处理结果,缓存为字典,将字典代入 UDF 进行数据处理。
# 假设引用了三方库中的 RuleEncryptUtil 静态类,该类不支持序列化,但是可以通过 Py4j 正常调用
RuleEncryptUtil = spark._jvm.com.fibodt.encrypt.RuleEncryptUtil
RuleEncryptUtil.b("http://www.baidu.com") # 可以正常调用,但是注册 udf 会报错
from pyspark.sql.functions import *
from pyspark.sql.types import *
test = spark.createDataFrame([['http://www.baidu.com'], ['']], ['url'])
def encrypt(dataFrame, host_col):
rows = dataFrame.select(host_col).distinct().collect()
rows_map = {row.asDict()['url']: RuleEncryptUtil.b(row.asDict()['url']) for row in rows}
encrypt_fp = udf(lambda url: rows_map.get(url, None), StringType())
return encrypt_fp(col(host_col))
4) pyspark 不支持 GenericRowWithSchema 类型:
如果需要返回指定类型记录,需要指定
returnType参数udf(returnType=schema)。
from pyspark.sql.types import *
from pyspark.sql.functions import *
# schema
struct_schema = StructType([
StructField("scenario", StringType(), True),
StructField("reach_type", StringType(), True),
StructField("product", StringType(), True),
StructField("campaign_dt", StringType(), True),
StructField("carrier", StringType(), True),
StructField("label_intention", StringType(), True),
StructField("label_sms_send", StringType(), True),
StructField("label_click", StringType(), True),
StructField("label_register", StringType(), True),
StructField("label_apply", StringType(), True),
StructField("label_approve", StringType(), True),
StructField("label_withdraw", StringType(), True),
StructField("credit_amount", StringType(), True),
StructField("credit_interest", StringType(), True),
StructField("approve_high", StringType(), True),
StructField("approve_mid", StringType(), True),
StructField("approve_low", StringType(), True),
StructField("label_approve_level", StringType(), True)
])
from pyspark.sql.functions import udf
@udf(returnType=struct_schema)
def shb_field_udf(row):
scenario = 'LOAN'
reach_type = 'SMS'
product = 'SHB'
campaign_dt = row.asDict().get('dt')
operator = str(row.asDict().get('operator'))
carrier = {'0': 'telecom', '1': 'mobile', '2': 'unicom'}.get(operator)
status = row.asDict().get('status')
if status == 'login':
label_intention = 1
label_sms_send = 1
label_click = 1
label_register = 0
label_apply = 0
label_approve = 0
label_withdraw = 0
elif status == 'regist':
label_intention = 1
label_sms_send = 1
label_click = 1
label_register = 1
label_apply = 0
label_approve = 0
label_withdraw = 0
elif status == 'finish':
label_intention = 1
label_sms_send = 1
label_click = 1
label_register = 1
label_apply = 1
label_approve = 0
label_withdraw = 0
elif status == 'approve':
label_intention = 1
label_sms_send = 1
label_click = 1
label_register = 1
label_apply = 1
label_approve = 1
label_withdraw = -1
else:
label_intention = -1
label_sms_send = -1
label_click = -1
label_register = -1
label_apply = -1
label_approve = -1
label_withdraw = -1
credit_amount = -1
credit_interest = -1
extend_field = json.loads(row.asDict()['extend_field']).get('acct_level')
if extend_field == 'A':
approve_high = 1
approve_mid = 0
approve_low = 0
label_approve_level = 10
elif extend_field == 'B':
approve_high = 0
approve_mid = 1
approve_low = 0
label_approve_level = 5
elif extend_field == 'C':
approve_high = 0
approve_mid = 0
approve_low = 1
label_approve_level = 1
else:
approve_high = 0
approve_mid = 0
approve_low = 0
label_approve_level = 0
return {
"scenario": scenario,
"reach_type": reach_type,
"product": product,
"campaign_dt": campaign_dt,
"operator": operator,
"carrier": carrier,
"label_intention": label_intention,
"label_sms_send": label_sms_send,
"label_click": label_click,
"label_register": label_register,
"label_apply": label_apply,
"label_approve": label_approve,
"label_withdraw": label_withdraw,
"credit_amount": credit_amount,
"credit_interest": credit_interest,
"approve_high": approve_high,
"approve_mid": approve_mid,
"approve_low": approve_low,
"label_approve_level": label_approve_level
}
uid_property = spark.table("data_test.e_jy_rules_mapping_uid_property")
df_status = spark.table("jxt_credit.rp_data_status").where(col("dt") == "2023-12-03") \
.where(col("product") == "SHENG_BEI") \
.drop(col("operator")) \
.join(uid_property, col("uuid") == uid_property["uid"]) \
.select("uuid", struct("*").alias("original_record"))
df_status.persist()
df_status.withColumn("field", shb_field_udf(col("original_record"))).show(truncate=False)
x. 窗口函数(排名分析函数)
配合分析函数使用:
from pyspark.sql.window import Window
# 基本用法
[排名分析函数, 聚合分析函数].over(Window.partitionBy(colName).orderBy(colName)).alias(newColName)
1) 支持的聚合函数:
| 函数 | 说明 |
|---|---|
cume_dist() | 窗口分区中值的累积分布 |
currentRow() | 返回表示窗口分区中当前行的特殊帧边界 |
row_number() | 行号,返回数据项在分组中的排名,排名相等时名次仍然按照单调递增序列生成数字 1,2,3,4 |
rank() | 排名,返回数据项在分组中的排名,排名相等会在名次中留下空位 1,2,2,4 |
dense_rank() | 排名,返回数据项在分组中的排名,排名相等在名次中不会留下空位 1,2,2,3 |
percent_rank() | 返回窗口分区中行的相对排名(即百分比) |
lag(e, offset, defaultValue) | 返回当前行向前偏移的行,使用时需结合 orderBy |
lead(e, offset, defaultValue) | 返回当前行向后偏移的行,使用时需结合 orderBy |
ntile(n) | 返回有序窗口分区中的分组 id(从 1 到 n) |
first(e, ignoreNulls) | 返回窗口分区中的第一个值,若指定 ignoreNull 参数为 true,则忽略 null 值 |
last(e, ignoreNull) | 返回窗口分区中的最后一个值,若指定 ignoreNull 参数为 true,则忽略 null 值 |
unboundedPreceding() | 返回表示窗口分区中前一行的特殊帧边界 |
unboundedFollowing() | 返回表示窗口分区中最后一行的特殊帧边界 |
count(e) | 返回窗口分区中的指定记录的数量 |
countDistinct(e) | 当前不支持该操作,使用 size(collect_set("colName1").over(Window.partitionBy("colName"))) 作为替代方案 |
当 partitionBy 分组之后,组内存在重复记录,orderBy 后重复记录会在一起,此时如果需要针对 id 列生成唯一排序值时,必须使用
dense_rank(),否则一个 id 会对应多个排序值。
2) 开窗函数 over:
-
[分析函数].over():对记录进行分组统计,每行每列都可以返回统计值 -
Window.partitionBy():对记录进行分组 -
Window.partitionBy(colName).orderBy(colName):对记录进行分组,在组内进行排序 -
Window.partitionBy(colName).rowsBetween(start, end):对记录进行分组,仅对介于 start、end 之间的行进行聚合操作- start 和 end 可以是:正数(1 代表当前行的下一行)、负数(-1 代表当前行的上一行)、0(当前行,等效于
Window.currentRow)、Window.unboundedPreceding(当前行之前的所有行)、Window.unboundedFollowing(当前行之后的所有行)、Window.currentRow
- start 和 end 可以是:正数(1 代表当前行的下一行)、负数(-1 代表当前行的上一行)、0(当前行,等效于
-
Window.orderBy(colName).rangeBetween(start, end):对排序的列(必须为数值类型)进行聚合操作,仅对介于 start、end 之间的行进行聚合操作- 参数含义同
rowsBetween
- 参数含义同
3) 特殊分组函数 rollup() 与 cube():
-
df.rollup(a, b, c):从左向右,从整体到局部,依次分组- 首先会对 (a, b, c) 进行 group by
- 然后再对 (a, b) 进行 group by
- 其后再对 (a) 进行 group by
- 最后对全表进行汇总操作
-
df.cube(a, b, c):从整体到局部,遍历所有组合进行分组- 则首先会对 (a, b, c) 进行 group by
- 然后依次是 (a, b),(a, c),(a),(b, c),(b),(c)
- 最后对全表进行汇总操作
4) 例子——按照累积和进行分段,每段求和值为 50:
原始数据 test.txt:
score|value|id
10|100|1
20|200|2
30|300|3
40|400|4
50|500|5
10|100|6
20|200|7
30|300|8
40|400|9
50|500|10
10|100|11
20|200|12
30|300|13
40|400|14
50|500|15
from pyspark.sql import Window
test = spark.read.option("header", "true").option("sep", "|").csv("test.txt")
group = test.withColumn("id", col("id").cast("int")) \
.withColumn("group",
ceil(sum(col("score")).over(
Window.orderBy("id").rowsBetween(Window.unboundedPreceding, Window.currentRow)
) / 50)
)
# Python 和 scala 中有细微区别:
# (1) 聚合函数中必需使用 col(),例如 sum(col("score")),否则会报错
# (2) Window.orderBy() 中不能使用 col().desc,但可以使用 col() 和字符串,
# 如需倒序排序,需要使用 desc(),如:desc("id")、asc("id")
group.show(False)
groupLine = group.groupBy("group") \
.agg(concat_ws(" or ", collect_set("id")).alias("groupLine"))
groupLine.show(False)
b. 对新的列进行处理
i. 导入必须的隐式转换
在 scala 中需要 import spark.implicits._,否则使用 $ 选择列时会报错。
ii. 在 where、filter 中使用列值筛选
在 pyspark 中直接使用函数即可:
import pyspark.sql.functions._
# 注意:如果是用 $"new_tag" 会报错没有这一列
s.withColumn("new_tag", lit("test")).where(col("new_tag") == "test").show()
c. 对重名列进行处理
df1.join(df2, df1("id") == df2("id"), "inner")
d. 对列进行重命名
df.withColumnRenamed("id", "id_other")
六、自定义 MLlib 类型,兼容 Pipeline
a. 自定义 Transform
from pyspark.sql import SparkSession
from pyspark.ml import Transformer
from pyspark.ml.param import Param, Params
from pyspark.ml.util import DefaultParamsReadable, DefaultParamsWritable
from pyspark.sql import DataFrame
from pyspark.sql.functions import *
# 创建 Spark 会话
spark = SparkSession.builder.appName("example").getOrCreate()
class CustomTransformer(Transformer, DefaultParamsReadable, DefaultParamsWritable):
def __init__(self, *args, **kwargs):
# 继承父类,需要调用父类的构造方法
super().__init__()
# 通过 Param 为自定义 Transformer 添加参数
self.inputCol = Param(self, "inputCol", "输入列的描述")
self.outputCol = Param(self, "outputCol", "输出列的描述")
self.customParam = Param(self, "customParam", "自定义参数的描述")
if len(args) != 0:
inputCol, outputCol, customParam = args
# 将传入参数赋值给自定 Transformer 作为属性值
self._set(inputCol=inputCol)
self._set(outputCol=outputCol)
self._set(customParam=customParam)
elif len(kwargs) != 0:
inputCol, outputCol, customParam = kwargs.get("inputCol"), kwargs.get("outputCol"), kwargs.get("customParam")
self._set(inputCol=inputCol)
self._set(outputCol=outputCol)
self._set(customParam=customParam)
else:
# 通过 _setDefault 为自定义 Transformer 添加参数默认值
# 当构造函数参数列表为空时,需要使用这种方式将成员属性初始化为默认值
self._setDefault(inputCol="defaultInputCol")
self._setDefault(outputCol="defaultOutputCol")
self._setDefault(customParam=1.0)
# _set 私有方法设置属性的值
def setInputCol(self, value):
return self._set(inputCol=value)
# getOrDefault 获取属性的值
def getInputCol(self):
return self.getOrDefault(self.inputCol)
def setOutputCol(self, value):
return self._set(outputCol=value)
def getOutputCol(self):
return self.getOrDefault(self.outputCol)
def setCustomParam(self, value):
return self._set(customParam=value)
def getCustomParam(self):
return self.getOrDefault(self.customParam)
# _transform 私有方法中实现处理逻辑
def _transform(self, df: DataFrame) -> DataFrame:
# 在这里实现你的自定义转换逻辑
custom_udf = udf(lambda x: x * self.getCustomParam())
return df.withColumn(self.getOutputCol(), custom_udf(col(self.getInputCol())))
# 创建示例数据
data = [(1, 2.0), (2, 3.0), (3, 4.0)]
columns = ["id", "value"]
df = spark.createDataFrame(data, columns)
# 创建并使用自定义 Transformer
custom_transformer = CustomTransformer(inputCol="value", outputCol="result", customParam=2.0)
# 也可以使用默认参数
# default_transformer = CustomTransformer()
result_df = custom_transformer.transform(df)
# 显示结果
result_df.show()
# 保存模型
# 为了支持 write 方法,需要继承 DefaultParamsWritable
model_path = "custom_transformer_model"
custom_transformer.write().overwrite().save(model_path)
# 加载模型
# 为了支持 read() 方法,需要继承 DefaultParamsReadable
loaded_transformer = CustomTransformer.read().load(model_path)
# 使用加载的 Transformer 进行转换
loaded_result_df = loaded_transformer.transform(df)
# 显示加载后的结果
loaded_result_df.show()
b. 自定义 Estimator
自定义 Estimator 需要两部分:
- EstimatorModel:用于模型加载(load)、模型预测(transform)
- Estimator:用于模型保存(save)、模型训练(fit)
from pyspark.ml import Estimator, Model
from pyspark.ml.param import Param, Params
from pyspark.ml.pipeline import DefaultParamsReadable, DefaultParamsWritable
from pyspark.sql import DataFrame
from pyspark.sql.functions import *
class CustomEstimatorModel(Model, DefaultParamsReadable, DefaultParamsWritable):
# 注意:需要对参数进行校验,需要支持无参数构造函数,
# 否则使用 CustomEstimatorModel 进行模型加载时,__init__() 会出问题
def __init__(self, *args, **kwargs):
super().__init__()
self.inputCol = Param(self, "inputCol", "输入列")
self.outputCol = Param(self, "outputCol", "输出列")
self.customParam = Param(self, "customParam", "初始化模型参数")
self.trainedParam = Param(self, "trainedParam", "训练后的模型参数")
# 注意:CustomEstimatorModel 负责模型的加载与预测,因此构造函数中需要带有训练得到的参数 trainedParam
if len(args) > 0:
inputCol, outputCol, customParam = args
self._set(inputCol=inputCol)
self._set(outputCol=outputCol)
self._set(customParam=customParam)
elif len(kwargs) > 0:
inputCol, outputCol, customParam = kwargs.get("inputCol"), kwargs.get("outputCol"), kwargs.get("customParam")
self._set(inputCol=inputCol)
self._set(outputCol=outputCol)
self._set(customParam=customParam)
else:
# 部分参数可以不传入参数,使用 _setDefault 初始化为默认值,或者不进行初始化,使用参数自己的 setter 方法进行设置
self._setDefault(inputCol="defaultInputCol")
self._setDefault(outputCol="defaultOutputCol")
self._setDefault(customParam=1.0)
self._setDefault(trainedParam=0.0)
def setInputCol(self, value):
return self._set(inputCol=value)
def getInputCol(self):
return self.getOrDefault(self.inputCol)
def setOutputCol(self, value):
return self._set(outputCol=value)
def getOutputCol(self):
return self.getOrDefault(self.outputCol)
def setCustomParam(self, value):
return self._set(customParam=value)
def getCustomParam(self):
return self.getOrDefault(self.customParam)
def setTrainedParam(self, value):
return self._set(trainedParam=value)
def getTrainedParam(self):
return self.getOrDefault(self.trainedParam)
def _transform(self, df: DataFrame) -> DataFrame:
custom_udf = udf(lambda x: x * self.getTrainedParam())
return df.withColumn(self.getOutputCol(), custom_udf(col(self.getInputCol())))
def transform(self, df: DataFrame) -> DataFrame:
return self._transform(df)
# CustomEstimatorModel 主要负责模型加载与预测,因此 load 和 transform 为必需函数,save 函数可以省略
def save(self, path):
self.write().save(path)
# 模型加载时,实际上会调用这里的 load 方法
@classmethod
def load(cls, path):
return cls.read().load(path)
class CustomEstimator(Estimator, DefaultParamsReadable, DefaultParamsWritable):
def __init__(self, *args, **kwargs):
super().__init__()
self.inputCol = Param(self, "inputCol", "输入列")
self.outputCol = Param(self, "outputCol", "输出列")
self.customParam = Param(self, "customParam", "初始化模型参数")
# 注意:CustomEstimator 负责模型的训练和保存,因此这里的参数列表中没有 trainedParam
if len(args) > 0:
inputCol, outputCol, customParam = args
self._set(inputCol=inputCol)
self._set(outputCol=outputCol)
self._set(customParam=customParam)
elif len(kwargs) > 0:
inputCol, outputCol, customParam = kwargs.get("inputCol"), kwargs.get("outputCol"), kwargs.get("customParam")
self._set(inputCol=inputCol)
self._set(outputCol=outputCol)
self._set(customParam=customParam)
else:
self._setDefault(inputCol="defaultInputCol")
self._setDefault(outputCol="defaultOutputCol")
self._setDefault(customParam=1.0)
def setInputCol(self, value):
return self._set(inputCol=value)
def getInputCol(self):
return self.getOrDefault(self.inputCol)
def setOutputCol(self, value):
return self._set(outputCol=value)
def getOutputCol(self):
return self.getOrDefault(self.outputCol)
def setCustomParam(self, value):
return self._set(customParam=value)
def getCustomParam(self):
return self.getOrDefault(self.customParam)
def _fit(self, df):
# 计算逻辑,计算出 trainedParam
trainedParam = 2.0
return CustomEstimatorModel(self.getInputCol(), self.getOutputCol(), self.getCustomParam(), trainedParam)
def fit(self, df):
return self._fit(df)
# CustomEstimator 主要负责模型训练和保存,因此 fit 函数和 save 函数为必需,load 函数可以省略
def save(self, path):
return self.write().save(path)
# load 函数可以省略
@classmethod
def load(cls, path):
return cls.read().load(path)
# 创建示例数据
data = [('1', 1), ('1', 2), ('1', 3), ('1', 4), ('1', 5)]
columns = ['ins', 'in']
df = spark.createDataFrame(data, columns)
# 创建并使用自定义 Estimator
custom_estimator = CustomEstimator('in', 'out', 1.0)
# 也可以使用默认参数
# custom_estimator = CustomEstimator()
# 也可以使用键值传参
# custom_estimator = CustomEstimator(inputCol='in', outputCol='out', customParam=1.0)
# 训练模型
model = custom_estimator.fit(df)
# 显示结果
result_df = model.transform(df)
result_df.show()
# 保存模型
# 为了支持 write 方法,需要继承 DefaultParamsWritable
model_path = "custom_estimator_model"
model.write().overwrite().save(model_path)
# 加载模型
# 为了支持 read() 方法,需要继承 DefaultParamsReadable
loaded_model = CustomEstimatorModel.read().load(model_path)
# 使用加载的模型进行转换
loaded_result_df = loaded_model.transform(df)
# 显示加载后的结果
loaded_result_df.show()