
1. 计数器里那个永远不动的 0累加器解决的到底是什么问题几年前我接过一个日志清洗的活需求很简单把一堆原始 JSON 行解析成结构化记录顺带统计有多少行解析失败。我当时写的是这样一段代码bad 0 def parse(line): global bad try: return json.loads(line) except Exception: bad 1 return None rdd sc.textFile(hdfs://.../raw/*.json) rdd.map(parse).filter(lambda x: x is not None).count() print(bad lines:, bad)跑完之后打印出来是bad lines: 0。可我明明知道这批数据里有脏行——我在本地用同样的解析逻辑跑过一遍至少三万条会挂。那一刻我的第一反应是Spark 是不是把异常吞了排查了半个小时日志才想明白bad 1这行代码确实执行了只不过是在 executor 的进程里执行的改的是 executor 本地那份bad变量的副本。driver 上那个bad从头到尾就没被碰过因为它根本没被序列化过去——Spark 序列化的是parse这个函数的闭包而global bad指向的是每个 Python 进程自己的模块全局命名空间。这就是PySpark 累加器存在的理由。它把在分布式任务里做计数/汇总这件事从闭包里的变量变成了一种有明确定义的分布式数据结构driver 端创建、序列化分发到每个 executor、任务在 executor 侧只能往里加、任务结束后增量回传、driver 端合并出最终结果。我特别喜欢用一个类比来解释它累加器像是一个只进不出的投票箱。每个 executor 都是投票站工作人员只能往箱子里投纸条不能打开箱子看现在有多少票所有投票站关门之后票箱统一运回总部开箱统计。这个不能打开看的设定不是缺陷是刻意设计——如果允许 executor 随时读取当前值就意味着每读一次都要跟全集群同步一次等于给每个任务加了一把全局锁分布式计算的性能优势会被直接抹平。所以累加器的能力边界很清晰它适合做监控指标、统计口径、调试探针不适合做需要精确对账的业务数据。这个判断会贯穿后面所有内容。适合读这篇文章的人包括几类正在写 ETL 清洗逻辑、需要在任务里埋统计点的数据工程师想知道自己那个数一直不对的累加器到底哪里出问题的人以及需要在 PySpark 里实现自定义聚合类型比如字典计数、样本收集的开发者。如果你只是用spark.sql()做 DataFrame 聚合那累加器基本跟你无关但只要你写 RDD 算子或者 Python UDF它迟早会找上你。2. 累加器的更新什么时候才算数四条语义规则很多人对累加器的困惑本质上不是 API 不会用而是不知道什么时候算数、算几次。这四条规则我建议你直接背下来能省掉大量排查时间。2.1 没有 action累加器永远是初始值Spark 是惰性求值的map、filter、flatMap这些 transformation 只是往血缘图lineage里记了一笔真正干活是在 action 被调用的时候。累加器的更新动作挂在任务执行上没有任务就没有更新。我见过有人这样写acc sc.accumulator(0) rdd.map(lambda x: (acc.add(1), x)[1]) print(acc.value) # 0因为没有任何 action这段代码里map返回的是一个新的 RDD 对象它被创建完就丢弃了一个任务都没跑起来。acc.value当然是 0。要让它有数必须补一个count()、collect()、foreach()、saveAsTextFile()之类的 action。这个坑在交互式调试的时候特别容易踩因为你在 notebook 里逐行敲很容易忘了最后那一下。2.2 action 里的更新是只算一次transformation 里是至少一次这是 Spark 官方文档里明确写过的语义差异也是最容易让人翻车的地方。对于在 action 内部发生的累加器更新比如rdd.foreach()里的那个函数Spark 保证每个任务的更新只被应用一次。但对于在 transformation 内部的更新比如rdd.map()里的那个函数如果任务或 stage 被重新执行更新可能会被应用多次。哪些情况会导致重新执行我给你列全任务失败重试某个 task 抛异常了Spark 会重跑它默认最多spark.task.maxFailures4次重跑时那次失败的累加器增量如果已经回传了就会重复计数。推测执行spark.speculationtrue时Spark 会为跑得慢的任务启动一个副本两个副本都会往累加器里加最后只有一个结果被采纳但累加器加了两次。缓存被驱逐导致血缘重算RDD 没有 persist被多个 action 消费每次 action 都会从头算一遍。stage 被多个 action 共享同一个 RDD 上连续做了count()和take(10)如果中间没 cache累加器就是双倍。所以那句老话是对的累加器在 transformation 里的语义是at-least-once在 action 里是exactly-once。如果你需要一个不重复的计数最稳的做法是把累加器更新放在 action 触发的算子里而不是顺手塞在中间的map里。2.3 executor 侧读 value 没有意义acc.value这个属性在 executor 里调用是个危险动作。它拿到的绝对不是全局值而是那个 executor 本地的一份副本甚至在某些版本里会直接抛异常。Spark 的设计意图很明确累加器是 write-only 的别读。我在代码 review 里见过这种写法def process(x): total.add(1) if total.value 1000: # 千万别这么干 return None return x写这段代码的人想让处理满 1000 条之后就跳过剩余数据但这个判断在分布式环境下完全没有意义——每个 executor 的total.value都是它自己那一小份而且这个值什么时候更新、能不能读到都是未定义行为。要读累加器永远只在 driver 端、在 job 结束之后读。2.4 任务被杀死时增量会丢这一条很少有人提但在生产上很致命。PySpark 的累加器增量是通过 Python worker 进程在任务处理过程中回传给 JVM 侧的。如果一个 task 因为内存溢出被 YARN/K8s 杀掉或者在回传之前 Python worker 崩了那么这次任务累积的更新就一起没了。这带来一个直接结论累加器的值只会偏大重复计算也可能会偏小丢失更新它不是一个可靠的账本。拿它做数据质量告警、做趋势观察都很好但如果你拿它对账比如处理了多少条、入库了多少条、必须严格相等迟早会出问题。真要精确对账用 DataFrame 聚合或者在任务结束后统一统计。3. PySpark 里开箱即用的累加器边界在哪里搞清楚语义之后回到最实际的问题PySpark 里能直接创建哪些累加器哪些必须自己写。3.1 默认支持的类型只有三种PySpark 的SparkContext.accumulator在做类型推断时只认三种 Python 类型对应三个内置的AccumulatorParam初始值示例推断出的类型背后的 Paramsc.accumulator(0)int整数加法sc.accumulator(0.0)float浮点加法sc.accumulator(0j)complex复数加法sc.accumulator([])list直接抛 TypeErrorsc.accumulator({})dict直接抛 TypeError如果你写sc.accumulator([])会拿到一句很不友好的报错TypeError: No default accumulator param for type class list。这句话我第一次看到的时候还以为是版本问题其实就是列表没有内置的加法参数实现。注意这里的设计逻辑Python 的list list是拼接如果 Spark 默认就用拼接来实现列表累加那每个任务回传一个完整列表、driver 端做 N 次拼接复杂度是 O(N×M)整个集群的数据会往 driver 汇聚这跟分布式计算的初衷背道而驰。所以 Spark 干脆不给默认实现逼你想清楚我到底要不要把数据收上来。3.2accumulator()的签名和参数顺序PySpark 侧的签名很朴素sc.accumulator(value, accum_paramNone)。第一个参数是初始值这个值在 driver 端创建时就用上了第二个参数是你自定义的AccumulatorParam实例不传就走上面那张表的类型推断。这里有个容易搞混的点初始值同时承担两个角色——driver 端它是起点值executor 端在反序列化时的起点值则是accum_param.zero()的返回值。很多人以为初始值会被复制到每个 executor其实不是。这个细节在写自定义累加器的时候非常关键后面第 4 节会展开。3.3 PySpark 没有具名累加器也没有longAccumulatorScala/Java 那边有sc.longAccumulator(myCounter)这种写法可以在 Spark UI 的 Accumulators 标签页里按名字区分。PySpark 长期只有sc.accumulator这一个入口既没有name参数也没有longAccumulator/doubleAccumulator这类快捷方法。你要是拿不准自己手上的版本最省事的办法是探一下print(hasattr(sc, longAccumulator)) # 大概率为 False因为没有名字Spark UI 里会显示一堆accumulator_0、accumulator_1编号还是全局递增的两个不同的 job 混在一起完全分不清谁是谁。我的做法是在自定义 Param 上挂一个label字段job 结束时把累加器的值连同 label 一起打进日志class LabeledParam(AccumulatorParam): def __init__(self, label, zero_value): self.label label self._zero zero_value def zero(self, value): return type(self._zero)() def addInPlace(self, v1, v2): return v1 v2然后在 driver 端统一打印。这样即使 Spark UI 里看不清日志里也能一眼对上。这是个很小的习惯但在排查线上问题时能救你半天时间。3.4add()和的等价关系累加器实例上只有两个动作可用acc.add(x)和acc x两者完全等价内部就是调add。需要注意的是在 lambda 里不太好用因为对闭包中的变量做增量赋值会触发 Python 的作用域规则。add()是更稳妥的选择acc sc.accumulator(0) rdd.foreach(lambda x: acc.add(1)) # 推荐 rdd.foreach(lambda x: acc.__iadd__(1)) # 不推荐虽然能跑4. 自定义累加器AccumulatorParam 的两个方法就是全部契约当你要统计的东西不是一个数字的时候比如按错误类型分类计数、收集异常样本、记录最大值和最小值就得自己实现AccumulatorParam。好消息是这个接口简单到只有两个方法。4.1zero()与addInPlace()的契约细节from pyspark import AccumulatorParam class DictParam(AccumulatorParam): def zero(self, value): return {} def addInPlace(self, v1, v2): for k, c in v2.items(): v1[k] v1.get(k, 0) c return v1用起来是这样counter sc.accumulator({}, DictParam()) rdd.foreach(lambda x: counter.add({classify(x): 1})) print(counter.value)看着简单但有两个契约必须遵守违反了就会出各种奇怪的错误。第一zero(value)必须返回一个全新的、独立的空值对象。我在早期代码里犯过这个错写了个def zero(self, value): return self._empty把空字典缓存在实例上。结果所有 executor 拿到的初始值是同一个对象的引用在同一个进程内多个累加器如果共用同一个 Param 实例就会互相串数据。正确做法是每次zero()都返回新的字面量{}或者[]。第二addInPlace(v1, v2)必须return合并后的值。这是最常见的翻车点。Python 的原地方法list.extend()、dict.update()、set.update()都返回None如果你顺手写成def addInPlace(self, v1, v2): v1.extend(v2) # 返回 None那累加器的内部值会在第一次更新后变成None第二次更新就会抛AttributeError: NoneType object has no attribute extend。这个错误只在有数据的时候才出现本地跑空数据集完全测不出来。记住不管方法名里有没有 InPlace都要显式 return。另外还有一个容易被忽略的点自定义 Param 类必须定义在模块顶层不能定义在函数内部。因为累加器实例要经过__reduce__序列化分发到 executor序列化时会连带 Param 实例一起打包定义在函数里的类会报AttributeError: Cant pickle local object。这个错误信息指向的是序列化但根因是类的作用域。4.2 一个可以直接抄的完整例子统计不合格记录并保留样本下面这个是我在好几个项目里反复用过的模式——统计脏数据条数同时保留前 N 条样本用于事后归因。关键在于limit参数限制样本数量避免 collection 类累加器把 driver 撑爆。from pyspark import AccumulatorParam class RejectParam(AccumulatorParam): def __init__(self, limit20): self.limit limit def zero(self, value): return {count: 0, samples: []} def addInPlace(self, v1, v2): v1[count] v2[count] room self.limit - len(v1[samples]) if room 0: v1[samples].extend(v2[samples][:room]) return v1 reject sc.accumulator({count: 0, samples: []}, RejectParam(limit20)) def validate(row): msg check_row(row) if msg: reject.add({count: 1, samples: [{id: row.get(id), msg: msg}]}) return None return row rdd.map(validate).filter(lambda x: x is not None).count() print(rejected:, reject.value[count]) for s in reject.value[samples]: print(s)这段代码有几个设计决策值得说一下。为什么样本上限是 20 而不是 1000因为每个 executor 都可能往累加器里塞样本最终 driver 端的样本数是所有任务贡献之和理论上限是limit但中间过程里每个任务的v2里可能带着几十条样本一起回传上限设大了通信开销会明显上升。20 条对定位问题已经足够了你要看具体某条数据用filter单独取出来看更合适。为什么用 dict 而不是两个独立的累加器因为count和samples是同一次判定的产物拆成两个累加器意味着每次判定要调用两次add()通信次数翻倍。合成一个复合值只需要一次回传。4.3 用defaultdict做分类计数时的两种写法如果要统计的分类很多用普通 dict 就要写一堆if k in d。用collections.defaultdict更顺手from collections import defaultdict class CounterParam(AccumulatorParam): def zero(self, value): return defaultdict(int) def addInPlace(self, v1, v2): for k, c in v2.items(): v1[k] c return v1defaultdict(int)是可 pickle 的所以序列化没问题。但要注意一点defaultdict在你读取一个不存在的 key 时会悄悄创建这个 key这在 driver 端遍历结果的时候会产生为什么结果里多了一个值为 0 的 key这种困惑。如果你对输出格式有洁癖合并完最后转成普通 dict 再返回def addInPlace(self, v1, v2): for k, c in v2.items(): v1[k] c return v1 # driver 端 result dict(reject.value)4.4 列表类型累加器为什么不该被当成数据收集器我见过有人这样用累加器all_rows sc.accumulator([], ListParam()) rdd.foreach(lambda x: all_rows.add([x])) # 然后在 driver 端遍历 all_rows.value 做后续处理这段代码能跑但它把分布式计算退化成了把所有数据搬到 driver。所有数据通过网络汇聚到单个进程数据量上到百万级的时候 driver 会直接 OOM。想取数据就用collect()、take(n)、sample()或者干脆write到存储再读。累加器的定位始终是指标不是数据通道。5. 三个可以直接搬进生产管线的累加器用法讲完原理和 API说几个我实际项目里用顺手了的场景你可以直接抄结构。5.1 数据质量看门狗让每条脏数据都能归因第 4.2 节那段代码就是最小实现这里补充一下怎么把它接到真实管线上。我的做法是在 pipeline 的入口处统一挂一个Metrics对象把一轮 job 需要的所有累加器收拢在一起这样 driver 端只在最后统一输出class Metrics: def __init__(self, sc, sample_limit20): self.raw_count sc.accumulator(0) self.reject sc.accumulator({count: 0, samples: []}, RejectParam(sample_limit)) self.by_type sc.accumulator({}, DictParam()) def report(self, logger): logger.info(raw%s rejected%s, self.raw_count.value, self.reject.value[count]) for t, c in self.by_type.value.items(): logger.info(reject_type[%s]%s, t, c) for s in self.reject.value[samples]: logger.info(sample: %s, s)这个类的价值在于把累加器的读取点压缩到一处。累加器在这个类里是只写不读的所有读取都发生在report()里而report()只在 job 结束后调用一次。这样就不会出现在任务里读 value这种问题了。5.2mapPartitions 累加器把回传次数从 N 降到 P这是我认为最有价值的一个优化技巧。累加器的每一次add()都要经过 Python worker 到 JVM 的通信层虽然不是每条都立即发送但高频调用确实会带来额外开销。如果你在逐行统计改成分区统计能有可见的提升def process_partition(rows): local_ok 0 local_bad 0 for row in rows: cleaned clean(row) if cleaned is None: local_bad 1 else: local_ok 1 yield cleaned stats.add({ok: local_ok, bad: local_bad}) # 每个分区只加一次 stats sc.accumulator({ok: 0, bad: 0}, DictParam()) rdd.mapPartitions(process_partition).count()这里绕了一点stats这个累加器必须在闭包里被引用到而process_partition定义在顶层所以得先定义累加器再定义函数或者把累加器作为参数传进去。我通常选择前者代码更短。回传次数从 RDD 行数降到分区数通常是一到两个数量级的差别在千万行级别的作业上很值得。顺带一个注意点mapPartitions返回的是生成器如果你在函数里yield之后才add那这次更新发生在分区数据全部消费完之后。如果下游有take(n)这类会提前终止的算子生成器可能没被完整消费累加器的更新就不会执行。这不是 bug是惰性求值的正常表现但会导致我用take试探性看几行累加器是空的这种困惑。5.3 在 UDF 里埋探针排查那个偶尔返回 None的函数Python UDF 最让人头疼的地方是它跑在 executor 里异常栈不一定完整回传到 driver你只能看到某行返回了 None。我一般会在 UDF 内部挂一个异常累加器errors sc.accumulator({count: 0, traces: []}, TraceParam(limit10)) udf(string) def normalize(v): try: return do_something(v) except Exception as e: errors.add({count: 1, traces: [f{type(e).__name__}: {e}]}) return NoneTraceParam就是RejectParam换个字段名。这样跑完一轮之后driver 端就能拿到异常发生了多少次、前 10 条异常信息是什么比你翻 executor 日志快得多。需要提醒的是在 DataFrame 的 UDF 里更新累加器是可行的——PySpark 会在 Python worker 处理完一批数据后统一回传——但不要在 UDF 里读.value原因和第 2.3 节说的一样。另外如果你用的是 Pandas UDFapplyInPandas那套一批数据的处理粒度更粗回传时机也随版本有差异建议先在测试环境验证一次再上生产。6. 结果对不上时的完整排查链路累加器出问题症状基本就那几种。我整理成一张表你可以按症状直接对号入座。症状最可能的原因快速验证方式永远是初始值没有触发 action或累加器没进闭包加一个count()再看值是预期的 2 倍整数倍RDD 没 cache被多个 action 重算rdd.cache()后重跑对比值比预期多一点不是整数倍推测执行、任务重试关掉spark.speculation重跑值比预期少任务被杀增量丢失看 executor 日志有无 OOM kill忽大忽小不稳定在 job 未完成时读 value确认读取点在 action 之后报NoneType相关错误自定义 Param 的addInPlace没 return检查返回值报Cant pickle local objectParam 类定义在函数内部移到模块顶层具体排查时我的顺序是这样的第一步确认读取时机。把acc.value那行挪到所有 action 之后在所有东西都跑完之后再读。这一步能解决一半以上的数值不对。第二步检查缓存。在同一个 RDD 上做两次 action如果中间没有cache()或persist()累加器必然翻倍。验证方法很简单加一行rdd.cache()再跑一遍如果数字变对了根因就找到了。注意cache()也是惰性的第一次 action 才会真正落盘。第三步把集群关掉。用local[1]加一个几百行的小数据集复现。如果本地单线程跑出来是准的、集群上是错的那基本可以确定是并发相关的问题——推测执行、重试、分区重算。如果本地也不准那就是代码本身的问题跟分布式无关这时候排查难度会低很多。第四步数 job 和 stage 的次数。打开 Spark UI看 Jobs 页面有几个 job、每个 job 几个 stage。如果发现同一个血缘被触发了两遍而你没有 cache那答案就出来了。这一步能帮你看清 Spark 实际执行了什么而不是你以为它执行了什么。第五步确认没有动态资源相关的重试。开启动态资源分配dynamic allocation之后executor 会被回收和重建正在跑的任务可能需要重试。这种场景下累加器偏大是常态别指望它精确。排查完之后你会得出一个很朴素的结论累加器适合做大概对的指标不适合做必须对的账。这个认知一旦建立你对它的容忍度就正常了。7. 这些场景我宁可不用累加器说了这么多累加器的用法现在说说什么情况下我会绕开它。需要精确结果的时候我会用rdd.aggregate()或者直接聚合。aggregate的返回值是随着 action 一起回到 driver 的走的是正常的任务结果通道不依赖累加器的回传机制语义上是可靠的。如果你只是想要一个总和或者计数比如这批数据的总金额是多少用aggregate或者DataFrame的agg(sum(...))都行没必要引入累加器。我做过一个简单的对照同样是统计一个 RDD 里满足条件的元素个数方案语义可靠性额外开销代码侵入性rdd.filter(pred).count()精确多一次 action需要重算无rdd.aggregate()精确无需要改变返回结构累加器 foreachat-least-once每条或每分区一次回传低可嵌入任意位置DataFrame聚合精确无需要转成 DataFrame从这张表能看出来累加器在语义可靠性这一栏是唯一一个不精确的。那它凭什么还有存在价值答案是代码侵入性低。aggregate要求你把整个 pipeline 的输出结构改造成值 累加器的组合而累加器可以在任意深度的嵌套逻辑里随手add一下完全不影响主流程的返回结构。在调试和监控这种顺手加一笔的场景下这个优势是决定性的。需要全量数据的时候我会用collect()或者写文件。前面说过了累加器不是数据通道。要把数据收到 driver用collect()或者toLocalIterator()后者更省内存一次只拉一个分区的数据。要落盘就用write然后离线分析。需要 driver 实时感知进度的时候我会用 SparkListener 或者进度日志。累加器是任务完成后才回传做不到实时。如果你要做一个进度条正确的方式是在每 N 条记录里打一行日志或者用 Spark 的 status tracker API 拿任务进度。累加器在长任务里看着很久不更新是因为任务还没结束——一个跑 40 分钟的大分区累加器在这 40 分钟里都是不动的。这个现象曾经让我一度以为作业卡死了后来才想明白。需要区分不同分区的贡献时我会在mapPartitions里带上分区 ID 自己组装。比如rdd.mapPartitionsWithIndex(lambda i, rows: ...)在回传的值里带上i这样在 driver 端就能看到分区维度的分布。这个技巧在排查数据倾斜时特别有用——你能一眼看出是不是某个分区的数据量是其他分区的几十倍。8. 把累加器用得不那么野几条工程约定最后分享几条我在项目里沉淀下来的约定都是踩过坑之后加的。给每个 job 封一个 Metrics 类。不要让累加器对象散落在代码各处那样 driver 端的读取点会失控。把所有累加器收在一个类里提供唯一的report()方法job 结束时调一次。这样executor 侧能读 value这种错误在结构上就不可能发生。给自定义 Param 带上 label并且把它打进日志。前面提过PySpark 没有具名累加器Spark UI 里的编号又对不上号。在 Param 上挂一个可读字符串在report()里输出能省掉大量这是哪个计数器的困惑。所有 collection 类累加器都必须设上限。不管是样本列表还是错误消息列表都要有limit。这个上限不只是防止内存问题更重要的是让这个累加器在语义上明确地表达我只关心前 N 条而不是暗示我要全量数据。单元测试里用本地模式和固定数据集断言累加器的值。累加器的逻辑错误尤其是自定义 Param 的addInPlace返回值在本地单线程环境下能完整复现。写个小测试def test_accumulator_merge(spark): acc spark.sparkContext.accumulator({count: 0, samples: []}, RejectParam(limit2)) data spark.sparkContext.parallelize([1, 2, 3, 4], 2) data.foreach(lambda x: acc.add({count: 1, samples: [x]})) assert acc.value[count] 4 assert len(acc.value[samples]) 2注意最后一条断言——样本只留 2 条这正是 limit 生效的证据。这类测试跑一次只要几秒但能挡住大部分低级错误。生产环境在 job 结束时把累加器的值落一份结构化日志。累加器的值只活在这次 SparkContext 的生命周期里context 一停就没了。我习惯在report()里同时输出一行 JSON方便后续用日志系统做告警和趋势图。这个习惯在定位昨天开始脏数据突然变多这类问题时特别管用。最后一个是关于版本差异的。PySpark 的累加器实现从 2.x 到 3.x 有过重写Accumulator的序列化方式和回传机制都有调整。如果你在排查一个2.4 上跑得好好的、升到 3.x 就不对了的问题先把怀疑对象放在版本差异上写一个最小复现脚本在两个版本上各跑一遍对比sc.accumulator的初始值语义和自定义 Param 的zero()调用时机。自定义 Param 的zero()在序列化时会被调用如果你的zero()里带了副作用或者返回了共享对象两个版本下的表现可能不一致。我在实际使用中最大的体会是累加器的难点从来不在 API而在你对它的语义边界有多清楚。把它当成监控指标用它非常好用、几乎没有替代品把它当成精确的计数器用它会在你最不希望的时候给你一个差一点点的数字而且这种错误在测试环境里根本不会出现。想清楚这一点剩下的事情就都简单了。