Spark RDD求平均值:键值对聚合、精度与输出格式实践
“求平均值”这四个字放在单机脚本里就是一行代码的事但换成 Spark RDD 来做同一件事我第一次提交的结果是所有人的平均分全变成了整数82.6 显示成 8291.3 显示成 91任务日志干干净净没有任何报错状态全是 SUCCESS。问题出在哪出在分布式计算的类型推断、分区边界和聚合顺序这三件事会同时咬你一口——而这一关考的恰恰就是这些。RDD 编程里的“求平均值”从来不是“加起来除以个数”这么简单它是把一次全局归约拆成“局部先算、再跨节点合并、最后才做除法”的三段式动作。下面这些内容适合刚接触 RDD 键值对算子的人也适合跑通过一次但被评测判定卡住、想弄清底层逻辑的人。我会把数据类型、聚合算子的选择、输出格式、排查链路一条条摊开讲代码给 Python 和 Scala 双份全部可以照着抄。1. 平均值在 RDD 里为什么不是“加起来除以个数”1.1 单机思路搬到集群上会断在哪单机求平均的思维链条是把所有数读到一个列表里求和除以长度。这在集群上直接照搬会断在两个地方。第一数据不在一个地方你没法“把所有数读到一个列表”除非把全量数据 collect 到 Driver 端——数据量小的时候能跑通几千万行的时候 Driver 直接 OOM而且这样做等于放弃了整个集群的算力。第二即使你分两次 action 去算 sum 和 countSpark 也会把同一份数据扫两遍中间还夹带着一次完整的 shuffle 开销。所以真正的分布式做法是让每个分区先算自己那部分的“局部和”与“局部个数”再把所有分区的局部结果按 key 合并最后一步才做除法。注意这个顺序——除法必须放在最后因为除法不满足结合律你不可能先把每个分区的平均值算出来再求平均除非每个分区的数据量完全相等这在实际数据里几乎不可能。很多初学者写出rdd.map(...).mean()之类的代码能跑那是 Spark 内部帮你封装了同一套逻辑但这一关要的是你自己手写聚合过程。1.2 三个必须先想清楚的问题类型、边界、顺序在动手写代码之前我习惯先问自己三个问题这三个问题基本覆盖了这一关所有的坑。第一是类型。sc.textFile读进来的一切都是字符串你必须显式转成数值类型。Scala 里90.toInt / 2得到 4590.toDouble / 2得到 45.0Python 里90 / 2得到 45.0但如果你用了90 // 2或者对两个整型做整除精度就丢了。这个类型问题不是小事它直接决定你的结果是对是错而且不会报错。第二是边界。数据里可能有空行、有分隔符数量不对的脏行、有某个键只出现一次的情况、有分数为 0 或者为空的情况。空行不处理解析时就会抛数组越界某个键只有一条记录计数器是 1这没问题但如果你的过滤逻辑把某个键的记录全过滤掉了它就不该出现在结果里这一点要跟评测的期望输出对齐。第三是顺序。RDD 的map、filter是窄依赖不触发 shufflereduceByKey、groupByKey、aggregateByKey是宽依赖会触发 shuffle。shuffle 意味着数据要落盘、要跨节点传输、要按 key 重新分区。理解这一点你才能解释为什么同样的逻辑groupByKey写法跑 40 秒reduceByKey写法跑 8 秒。提示数据规模小的时候各种写法的耗时差异看不出来这时候不要凭“能跑通”就下结论要看执行计划里有没有ShuffleExchange以及 shuffle 的读写量。2. 把“求平均”翻译成分布式可执行的两步聚合2.1 数据先变成 (键, 数值) 的键值对RDD 的聚合算子全都作用在键值对 RDD 上也就是RDD[(K, V)]或者 Python 里的RDD of tuple。所以第一步永远是解析文本、构造键值对。假设数据文件长这样Tom,90 Tom,80 Jerry,70 Jerry,88 Jerry,95目标是求每个人的平均分。解析过程我通常会写成一条链每一步只做一件事pairs raw.filter(lambda line: line.strip() ! ) \ .map(lambda line: line.strip().split(,)) \ .filter(lambda arr: len(arr) 2) \ .map(lambda arr: (arr[0].strip(), float(arr[1].strip())))这里有几个细节值得说。filter放在split之前过滤空行是因为空行split之后得到[]长度是 1虽然也能被后面的长度校验拦下来但先过滤掉更省事。float()显式转换是关键少了这一步后面所有算术都会按字符串处理或者按整型处理。arr[0].strip()是为了对付分隔符旁边的空格比如Tom , 90这种写法很多平台给的测试数据里真的会有。2.2 为什么要同时携带“和”与“个数”键值对构造好之后核心动作是把(姓名, 分数)变成(姓名, (分数和, 记录数))。为什么要携带两个值因为平均值 总和 / 个数而这两个量各自都满足结合律可以分段计算再合并平均值本身不满足。这是整个问题的题眼。具体做法是先用mapValues把每条记录从(Tom, 90)变成(Tom, (90.0, 1))再用reduceByKey做合并sum_count pairs.mapValues(lambda score: (score, 1)) \ .reduceByKey(lambda a, b: (a[0] b[0], a[1] b[1]))reduceByKey的合并函数是(a, b) - a b的推广这里a和b都是(和, 个数)这样的二元组所以合并逻辑就是两个位置分别相加。这个合并函数会被调用两次场景一是在 map 端对同一分区内的相同 key 做预聚合二是在 reduce 端对来自不同分区的结果做最终合并。函数必须是可交换、可结合的加法满足这两个性质所以没问题。2.3 常见数据格式的解析细节不同来源的数据解析时的注意点不一样。我用过几类数据形态推荐解析方式容易踩的坑CSV逗号分隔line.split(,)字段里有逗号需要split(,, -1)或引号处理制表符分隔line.split(\t)复制粘贴时 tab 变成空格空格分隔line.split(\\s)多空格、行首空格导致第一个字段为空固定宽度按索引切片中文字段宽度计算不准注意split在 Scala 里接收的是正则表达式所以split(\\s)要写双反斜杠而split(,)这种单字符分隔符写得没问题但如果你写了split(|)在正则里|是“或”的意思会得到一堆空字符串正确写法是split(\\|)。这个坑我在处理日志数据时踩过一次排查了半小时才发现问题不在算子而在分隔符。解析完之后建议立刻对pairs做一次cache()因为后面你可能既要算全局平均、又要算分组平均、还要打印样本多次 action 会重复解析缓存能省掉这部分开销。缓存也要看情况数据量超过 executor 内存的话缓存反而会拖慢速度甚至 OOM这时候用MEMORY_AND_DISK级别。3. 四种写法的横向对比从 reduceByKey 到 groupByKey3.1 reduceByKey 二元组累加最短的写法前面那段mapValues reduceByKey已经是最短路径了加上最后一步相除就是完整版avg sum_count.mapValues(lambda v: round(v[0] / v[1], 2))这里round(v[0] / v[1], 2)里的除法是浮点除法因为v[0]已经是float结果自然是浮点。如果数据里的分数是整数并且你没有转floatv[0]和v[1]都是整型v[0] / v[1]在 Python 3 里得到浮点在 Python 2 里得到截断的整数在 Scala 里Int / Int也是截断的整数——这就是开头那个“平均分全变整数”的真实原因。mapValues相比map有个细节优势mapValues只改变 value不改变 key 和分区信息在很多算子链里能减少不必要的分区重算。虽然性能差异不总是明显但语义上更清晰代码也更容易读。3.2 aggregateByKey 的零值陷阱aggregateByKey的签名是三个参数零值、分区内合并函数、分区间合并函数。avg pairs.aggregateByKey( (0.0, 0), # zeroValue lambda acc, v: (acc[0] v, acc[1] 1), # 分区内 lambda a, b: (a[0] b[0], a[1] b[1]) # 分区间 ).mapValues(lambda v: round(v[0] / v[1], 2))零值这里是(0.0, 0)一个不可变元组。为什么强调不可变因为aggregateByKey的零值会在每个分区、每个 key 上被初始化一次如果零值是可变对象比如 Python 里的列表而且你在合并函数里对它做了原地修改acc.append(v)那么同一分区内多个 key 可能共享同一个对象引用结果就会串数据。这个坑非常隐蔽因为小数据量下不一定复现。实践建议是永远用不可变结构做零值要累加就返回新对象。另外零值的类型必须和合并函数的返回类型一致。我见过有人把零值写成0然后合并函数返回(sum, count)元组运行时报类型不匹配。类型系统在 Scala 里会直接编译失败Python 里则要等到运行期某个 task 报错才暴露。3.3 combineByKey 三个函数各管什么combineByKey是这三个里的“原始形态”理解它另外两个都是它的特例。三个函数职责分明avg pairs.combineByKey( lambda v: (v, 1), # createCombiner某个 key 在本分区第一次出现 lambda acc, v: (acc[0] v, acc[1] 1), # mergeValue同分区后续记录并入 lambda a, b: (a[0] b[0], a[1] b[1]) # mergeCombiners跨分区合并 ).mapValues(lambda v: round(v[0] / v[1], 2))createCombiner只在每个分区的每个 key 上被调用一次它的输入是单条记录的 value输出是累加器的初始形态。mergeValue处理本分区内该 key 的其余记录。mergeCombiners处理不同分区之间相同 key 的累加器合并——注意它的两个入参都是累加器不是原始值这是最容易写错的地方。aggregateByKey和combineByKey的区别在于零值aggregateByKey让你显式给出一个零值合并函数不区分“第一次”和“后续”统一按(acc, v)处理combineByKey不要零值用createCombiner处理第一次。数据里 key 的分布比较均匀、累加器类型一致时两者效果一样。3.4 groupByKey 能跑通但代价在哪最直观的写法是先把同一 key 的所有 value 聚成一个迭代器再在迭代器上算平均avg pairs.groupByKey().mapValues(lambda scores: round(sum(scores) / len(scores), 2))逻辑上好理解但代价很大groupByKey不做 map 端预聚合它会把所有(key, value)原封不动地按 key 重新分区、通过网络传输、落盘然后在 reduce 端堆成完整的序列。数据量一上来网络传输量和内存占用都是reduceByKey的数倍。更要命的是数据倾斜——某个 key 有几百万条记录reduce 端某个 task 就得把这百万条塞进内存很容易 OOM。四种写法的对比如下写法shuffle 数据量内存风险代码复杂度适用场景reduceByKey 元组小map 端预聚合低低首选绝大多数求平均场景aggregateByKey小低中累加器类型与 value 类型不同时combineByKey小低中高需要精细控制初始值时groupByKey大高最低需要保留全部原始值做多次计算时经验如果一道题只是求平均、求和、求最值groupByKey一律不用。只有在“同一 key 的完整明细后面还要复用多次”这种场景下它才有存在价值。4. 从本地跑通到集群提交的完整链路4.1 环境与版本对齐本地调试和集群提交最大的差别在于依赖和模式。PySpark 环境下pip install pyspark装好之后代码里SparkContext的创建方式随版本有变化Spark 2.x 之后SparkSession是统一入口通过spark.sparkContext拿到sc老版本代码里直接SparkContext(conf)也还能用但混用两套 API 容易出问题。我一般写from pyspark.sql import SparkSession spark SparkSession.builder \ .appName(rdd-average) \ .master(local[2]) \ .getOrCreate() sc spark.sparkContext sc.setLogLevel(WARN).master(local[2])表示本地起两个线程模拟两个分区能让你在小数据量下也看到多分区合并的行为。setLogLevel(WARN)是为了把 INFO 级别的日志压掉不然屏幕上全是任务调度信息真正有用的输出会被淹没。Scala 侧的依赖如果用 Maven 管理spark-core的版本必须和集群里的 Spark 版本严格对齐2.4.x的代码配3.x的集群多数情况能跑但涉及序列化细节时会出现莫名其妙的ClassNotFoundException。我吃过一次亏本地用3.2.0编译集群是2.4.8mapValues返回类型推断有差异编译期就报了不兼容。4.2 完整代码Python 与 Scala 双版本Python 完整版from pyspark.sql import SparkSession def main(): spark SparkSession.builder.appName(rdd-average).getOrCreate() sc spark.sparkContext sc.setLogLevel(WARN) raw sc.textFile(data/score.txt) pairs raw.filter(lambda line: line.strip() ! ) \ .map(lambda line: line.strip().split(,)) \ .filter(lambda arr: len(arr) 2) \ .map(lambda arr: (arr[0].strip(), float(arr[1].strip()))) # 第一种reduceByKey result pairs.mapValues(lambda s: (s, 1)) \ .reduceByKey(lambda a, b: (a[0] b[0], a[1] b[1])) \ .mapValues(lambda v: round(v[0] / v[1], 2)) \ .sortByKey() for name, value in result.collect(): print(%s\t%s % (name, value)) sc.stop() if __name__ __main__: main()Scala 完整版用combineByKey顺便演示精度控制import org.apache.spark.{SparkConf, SparkContext} object AverageScore { def main(args: Array[String]): Unit { val conf new SparkConf().setAppName(rdd-average) val sc new SparkContext(conf) sc.setLogLevel(WARN) val pairs sc.textFile(data/score.txt) .map(_.trim) .filter(_.nonEmpty) .map { line val arr line.split(,) (arr(0).trim, arr(1).trim.toDouble) } val avg pairs.combineByKey( (v: Double) (v, 1), (acc: (Double, Int), v: Double) (acc._1 v, acc._2 1), (a: (Double, Int), b: (Double, Int)) (a._1 b._1, a._2 b._2) ) .mapValues { case (sum, cnt) BigDecimal(sum / cnt).setScale(2, BigDecimal.RoundingMode.HALF_UP).toDouble } avg.sortByKey().collect().foreach(println) sc.stop() } }Scala 里那个BigDecimal(...).setScale(2, HALF_UP)是有意写的。因为Math.round和BigDecimal默认的舍入行为在.5边界上处理不一样如果评测数据里有刚好落在x.xx5上的分数用不同的舍入方式会得到不同结果这种失败最难查。4.3 输出格式与精度控制这一关被卡住的人很大一部分不是算法写错而是输出格式跟期望不匹配。评测通常是把你的输出和标准答案逐字符比对多一个空格、少一个星号、小数点后位数不一致全算错。几个高频格式要求和对策要求姓名 平均分空格分隔时用%s %s % (name, value)别用print(name, value)因为后者在某些环境下会带上多余分隔。要求保留两位小数时round()在 Python 3 里采用的是“四舍六入五取偶”round(2.675, 2)得到的是2.67而不是2.68这是浮点表示误差导致的。要严格控制用%.2f % value或Decimal配合ROUND_HALF_UP。字符串格式化的%.2f底层走的是另一套舍入和round结果可能不同务必先确认评测期望哪一种。要求结果写到文件时用saveAsTextFile(result)它输出的是目录目录下part-00000才是数据文件。如果平台读取时只认单个文件加一步.coalesce(1)再保存。但coalesce(1)会让所有数据汇到一个 task数据量大时是性能杀手只在小结果集上用。值本身是元组时saveAsTextFile会写成(Tom,85.0)这种带括号的形式去掉括号需要map(lambda kv: %s,%s % kv)提前转换。4.4 提交、验证与结果自查本地跑通之后提交集群spark-submit的几个关键参数spark-submit \ --master yarn \ --deploy-mode client \ --executor-memory 2g \ --executor-cores 2 \ --num-executors 4 \ --conf spark.default.parallelism100 \ average.py--num-executors只在 YARN 模式下有效local模式下加了也没用。spark.default.parallelism决定了 shuffle 之后的分区数默认值跟集群配置有关在本地模式里shuffle 分区数默认是local[N]里的 N而集群模式里常常是 200。分区数太多会产生大量小文件每个分区一个输出文件几百个小文件对后续读取非常不友好。自查阶段我固定做三件事一是拿 5 条数据手动算一遍平均分跟程序输出对比二是检查结果条数是否等于预期的人数少了说明有数据被过滤掉了三是看有没有某个值明显异常比如 0 或者负数那通常是解析出了问题。5. 评测判定失败时的排查链路5.1 先分清是“结果错”还是“格式错”排查的第一刀必须切在这里。做法很简单把期望输出和你的输出并排贴出来逐行逐字符看。如果数值完全一样、只是格式不同那问题在输出环节跟算法无关如果数值本身不同才需要往计算逻辑里查。我遇到的真实案例是本地print的结果和期望完全一致提交后就是不过。最后发现平台是读取saveAsTextFile生成的目录而我打印用的是print保存用的却是原始 RDD两份数据不一致。这种“调试通道和产出通道不是同一条”的问题很多人第一次都会忽略。5.2 整数除法与除零报错的定位方法如果输出全是整数或者出现java.lang.ArithmeticException: / by zero按下面顺序查。第一步打印累加器的内容sum_count.take(5)看看 value 到底是(90.0, 1)这样的浮点元组还是(90, 1)这样的整型元组。是整型说明前面的float()转换没生效或者被 Python 2 的整除规则吃掉了。第二步检查count有没有可能是 0。什么情况下会出 0如果你用aggregateByKey并且后续做了过滤或者在 Driver 端算全局平均时对一个空 RDD 做了total / cnt那就会除零。对空数据要提前判断if cnt 0: return 0.0或者直接在过滤阶段把空值挡掉。第三步检查有没有null或None参与了运算。Python 里None 1.0会抛TypeErrorScala 里对应的报错是NullPointerException。数据里的空字段split之后是一个空字符串float()直接抛ValueError所以解析阶段的长度和内容校验必须做。5.3 shuffle 分区数与输出文件数量这是第二个高频的失败原因。集群默认spark.sql.shuffle.partitions或spark.default.parallelism是 200你的数据只有 10 个人reduceByKey之后会生成 200 个分区其中 190 个是空的。saveAsTextFile会为每个分区生成一个文件结果是 200 个文件其中part-00190之类的空文件也照样存在。对评测平台而言如果它只读第一个文件而第一个恰好是空的就会判定失败。对策是调整并行度result sum_count.repartition(2).mapValues(...)或者直接用coalesce。两者的差别值得记一下repartition会触发完整 shuffle把数据均匀打散coalesce只做分区合并不 shuffle速度快但可能不均匀。合并分区这种场景用coalesce(1)就够了。注意coalesce(1)把一个 RDD 压到单分区意味着后续的所有计算都在一个 task 上跑。求平均这种结果集很小的场景没问题但如果是在中间步骤上做这个操作整个作业的并行度就废了。5.4 闭包变量与序列化问题Scala 里如果合并函数引用了外部的类成员或者不可序列化的对象会报Task not serializable。原因是算子里的匿名函数会被闭包序列化后发到 executor函数里捕获的所有外部引用都得跟着序列化。我见过的典型写法是把一个HashMap定义在 Driver 端的类里面然后在map里引用它做查表——单机没问题一提交就报错。对策是把需要的配置抽成一个可序列化的常量或者用广播变量分发val bonusMap sc.broadcast(Map(A - 1.05, B - 1.0)) val adjusted pairs.map { case (k, v) (k, v * bonusMap.value.getOrElse(k, 1.0)) }Python 端也有类似问题闭包引用的对象会被 cloudpickle 序列化如果引用了打开的文件句柄、数据库连接或者线程锁同样会失败。原则很简单算子函数里只引用基本类型、不可变集合和广播变量。6. 这一关之后加权平均、方差与 TopN 的复用思路6.1 加权平均把权重塞进累加器课程成绩的加权平均是一个很自然的延伸每门课有学分总评 各科成绩×学分之和 / 学分之和。用同一套“双累加器”的思路只是累加器从(分数和, 个数)换成(加权和, 权重和)# pairs: (学生, (分数, 学分)) weighted pairs.mapValues(lambda x: (x[0] * x[1], x[1])) \ .reduceByKey(lambda a, b: (a[0] b[0], a[1] b[1])) \ .mapValues(lambda v: round(v[0] / v[1], 2))结构完全没变只是每个字段的含义换了。这正是掌握累加器思维的价值——一旦你把“平均值”抽象成“分子和 / 分母和”所有类似的指标都是同一个模式的变形。人均消费额、客单价、转化率、复购率本质都是这个双累加器的组合。6.2 同一套模式求方差方差需要三个量个数、和、平方和。用combineByKey的累加器变成三元组acc pairs.combineByKey( lambda v: (1, v, v * v), lambda a, v: (a[0] 1, a[1] v, a[2] v * v), lambda a, b: (a[0] b[0], a[1] b[1], a[2] b[2]) ) def variance(t): n, s, sq t mean s / n return sq / n - mean * mean var_rdd acc.mapValues(variance)这里有个数值稳定性的坑sq/n - mean²这个公式在数值很大、方差不大的时候会出现“大数相减”有效数字被吃掉结果可能变成负数。生产环境里更稳的做法是用 Welford 在线算法在合并函数里做增量更新。这一关的数据量下不会暴露但知道这个边界比不知道要好。6.3 性能上的几个经验值最后分享几个我在实际任务里摸出来的经验值都是可验证的数据量在几十万行以内本地local[4]跑就够没必要上集群提交和排队的时间比计算时间还长。shuffle 分区数设成 executor 核数的 2 到 3 倍比较均衡。4 核 4 executor 的场景16 到 48 个分区是很舒服的区间200 个分区做小数据完全是浪费。累加器类型选元组而不是自定义对象能省掉序列化的开销。Scala 里(Double, Int)用 Kryo 序列化比自定义 case class 便宜不少配置spark.serializer为KryoSerializer在大数据量下能省 20% 到 30% 的 shuffle 传输量。mapValues和map在这类聚合链里的性能差异不大但mapValues语义更准代码可读性更好选它。累加器sc.longAccumulator统计全局计数很方便但不要用它算业务结果。任务失败重试时累加器会被重复累加得到的数比真实值大用它做平均值除数会直接算错。它只适合做监控指标。还有一个我在这一关反复验证过的细节平均值这种问题一定先确认业务上到底要保留几位小数、按哪种规则舍入再动手写代码。我见过同一个数据集用round和用%.2f处理20 个学生里有 3 个结果不一样而且都恰好卡在.005这种边界上。与其事后猜测评测期望哪种不如在第一次提交前就把两种结果都打印出来对照一遍两分钟的事能省掉好几轮反复提交。