凌晨两点半,手机在床头柜上疯狂震动。生产集群的告警说,我们新上线的PySpark机器学习评分任务连续失败了6次。打开YARN界面,错误信息翻来覆去就那几行:Py4JError、PickleException,偶尔还夹杂着AttributeError: 'NoneType' object has no attribute 'predict'。第一反应是模型文件在HDFS上损坏了——毕竟那个LightGBM的pkl足有200MB,上传半路断掉也不奇怪。但一连串排查下来,发现模型文件完好,本地跑同样的代码一点问题没有,一旦提交到集群就随机抽风。那天的经历,让我第一次正视了一个此前一直被忽略的词:跨节点依赖。
如果你也在用PySpark跑机器学习推理或训练任务,大概率迟早会遇到类似问题。代码在本地环境跑得行云流水,一到集群就各种奇怪的序列化错误、内存爆炸、结果不一致。这篇文章不打算讲PySpark的基础API,而是聚焦生产化过程中最磨人的一个环节:如何正确处理driver节点和executor节点之间的模型对象、状态、数据传递。我会结合自己实际踩过的坑,把7种解决方案掰开揉碎讲清楚,并给你一套能直接抄作业的选型思路。
1. 深夜告警之后:一次跨节点依赖事故的完整复盘
1.1 事故现场:日志里那些看不懂的异常
那天的任务链路其实很简单:Spark读取HDFS上的用户行为parquet表,经过特征拼装后,调用一个训练好的LightGBM排序模型,输出每个候选物品的点击概率,再写回HDFS。模型是用joblib.dump保存的pkl文件,大约200MB,放在HDFS的/models/目录下。本地测试时,我用同样的数据切片跑了三遍,输出正常、速度正常、内存正常。
上集群之后,任务在最后一个stage反复失败。奇怪的是失败并不是每次都在同一批数据上,而是随机飘的。Spark UI里能看到某些task成功,某些task反序列化失败,整个stage被反复重试。当时日志里最扎眼的是这几类:
PickleException: Could not serialize object,后面跟着一长串Py4J错误EOFError: Ran out of input,发生在读取模型文件的pickle时AttributeError: 'NoneType' object has no attribute 'predict',说明有人拿到了一个None模型对象
这三种错误交替出现,很难让人第一时间联想到同一个根因。我甚至一度怀疑是某个坏节点上的磁盘有问题,申请换了两台机器,现象依旧。
1.2 排查链路:从怀疑模型文件到盯上闭包序列化
排查过程大概走了一个多小时,步骤是这样的:
先确认模型文件完整性。我在driver端用joblib.load加载了一次,能正常加载,predict也正常。怀疑是HDFS读文件时网络抖动导致executor读到半截,于是在executor端加了一个校验逻辑:加载后检查模型的n_features_in_字段是否等于预期值。加了校验之后,偶尔能捕捉到加载的模型确实是个损坏的或者不完整的对象——但文件本身明明没问题。
接着盯上了并发。在UDF里,我每来一行数据就调用一次模型加载函数。你可以算一下:5000万条数据,默认200个partition,每个partition对应一个task,而并行度开到200的话,同一时刻可能有200个task在各自执行joblib.load。200个进程同时从HDFS拉一个200MB的pkl文件,HDFS的NameNode和DataNode压力瞬间拉满,部分连接超时、半包读取,就会产生上面那些随机且诡异的异常。
再往下挖,发现更大的问题在UDF本身。当时的代码长这样:
# driver端加载模型 model = joblib.load("hdfs:///models/lgb.pkl") @F.udf("double") def predict_prob(features): # 这里引用了driver端的model变量! return float(model.predict_proba(features.reshape(1, -1))[0][1]) result = df.withColumn("score", predict_prob(F.col("feature_vector")))这段代码在本地跑没问题,因为本地模式下driver和executor在同一进程里,model就是内存里的同一个对象。但提交到YARN集群后,driver和executor是不同进程甚至不同机器,PySpark必须把predict_prob这个UDF引用的所有外部变量——也就是model——一起序列化后塞进task里分发出去。当闭包里携带的是一个200MB的模型对象时,会发生两件事:每个task都得先反序列化一遍模型,耗时爆炸;同时这个模型对象里如果引用了不可序列化的底层资源(比如libgomp的线程句柄),就会直接报PickleException或EOFError。
1.3 根因定性:分布式环境下的"闭包陷阱"
Spark的分布式执行模型决定了,凡是UDF或算子函数中引用的外部对象,都必须能通过Java的pickle机制传给executor。这个机制叫闭包序列化。很多人写PySpark时没有意识到:你在driver端定义的一个普通局部变量,一旦被算子里引用,就会成为整个task二进制的一部分,被复制到每一个执行单元里。
问题在于,机器学习模型对象往往很重。一个200MB的pkl文件,序列化进闭包后,每个task都要带着这份拷贝;200个task就是40GB的传输和反序列化开销。更麻烦的是模型内部可能持有线程池、OpenMP运行时状态、文件句柄等无法pickle的东西,导致序列化直接失败。这就是跨节点依赖的第一层含义:driver端对象无法安全、高效地跨越节点到达executor。
还有第二层含义,是executor端数据或状态很难可靠地回到driver。比如你想统计每个partition处理了多少条数据,在executor里更新一个driver端变量,以为最后能读到总数,结果发现driver端变量根本没变。这也是跨节点依赖——反向的数据流动被分布式内存模型阻断了。后面几章里,这两种方向的问题都会有对应解法。
2. 跨节点依赖的四种形态:本地能跑、上集群就爆的秘密
2.1 形态一:UDF闭包捕获了driver端的模型对象
这是最常见也最容易踩的一类。凡是把模型加载放在driver端、然后在UDF或map里直接引用模型变量的写法,都属于闭包捕获。本地模式下因为进程没隔离,感知不到问题;一到集群,要么模型被序列化进每个task导致开销巨大,要么序列化失败导致任务崩溃。
更隐蔽的变体是:模型加载不在UDF里,而是在一个模块的顶层。比如你写了一个utils/model_loader.py,在模块加载时执行MODEL = joblib.load(...)。当executor import这个模块时,它会执行这段加载代码。如果HDFS路径在所有节点上可见,这个方案在"每个executor只加载一次"的前提下勉强可行——但executor上的Python worker进程可能被反复创建和销毁,实际加载次数仍然不可控,而且加载失败的异常会直接炸掉整个task。这类写法最大的问题是:你没法精确控制加载时机和加载次数。
2.2 形态二:driver端的可变状态,executor看不见
我曾经在代码里用了一个全局计数器:
processed = 0 def process(row): global processed processed += 1 return row df.rdd.map(process).count() print(processed) # 永远是0!原理很简单:executor上跑的是process函数的一个反序列化副本,它修改的processed是executor进程里的全局变量,和driver端的processed毫无关系。驱动端的processed只有在task执行完后,通过累加器或shuffle结果才能回传。很多人不理解"为什么我的计数器没有累加",本质上就是没搞清楚分布式环境下没有共享内存这一铁律。
2.3 形态三:非序列化对象被夹带进闭包
类似threading.Lock、数据库连接、文件句柄、socket等对象,一旦被闭包捕获,序列化阶段就会报错。还有一种情况是对象本身能序列化,但里面嵌套了不可序列化的属性。比如LightGBM的Booster对象在某些版本里可以pickle,但如果你在模型上挂了自定义的logger或回调函数,就可能触发序列化异常。这类问题比形态一更棘手,因为报错信息往往不直接指向你的模型对象,而是指向某个深层依赖。
2.4 形态四:collect把整个结果集倒灌回driver
还有一类跨节点依赖不涉及模型对象,而是数据回传。常见的操作是:
all_rows = df.collect() for row in all_rows: do_something(row)当DataFrame特别大时,collect会把所有executor的数据通过网络传输到driver端,driver内存直接被打爆,或者GC时间暴长。这种做法本质上把"分布式计算"退化成了"单机计算",和跨节点依赖的语义是反的。就算你用了toLocalIterator,如果后续处理逻辑里还是需要全局状态,依然会踩形态二的坑。
2.5 统一视角:driver、executor之间只有三种数据流动
我把上面的形态归纳一下:跨节点依赖本质是driver和executor之间的数据流动出了问题,而流动方式只有三种:
- driver → executor:依赖闭包序列化,问题表现为模型加载慢、序列化失败、重复加载
- executor → driver:依赖collect或累加器,问题表现为内存爆炸、状态不同步、计数不准
- executor → executor:依赖shuffle,问题表现为数据倾斜、shuffle开销大,但这块有专门的文章讲,本文不展开
搞清楚这三种流动,后面所有解决方案都是围绕"如何让数据更安全、更高效地在这些路径上流动"展开的。
3. 方案一至三:让模型和参数安全抵达每一个executor
3.1 方案一:broadcast广播,一次打包全节点共享
当模型对象在20MB到200MB之间、推理逻辑简单、且你希望避免重复序列化时,广播变量是第一选择。
# driver端加载模型 model = joblib.load("hdfs:///models/lgb.pkl") model_bc = spark.sparkContext.broadcast(model) @F.udf("double") def predict_prob(features): # executor端通过 .value 拿到模型 m = model_bc.value return float(m.predict_proba(features.reshape(1, -1))[0][1]) result = df.withColumn("score", predict_prob(F.col("feature_vector")))广播变量解决了两个问题:一是模型只在driver端序列化一次,然后通过Spark内置的TorrentBroadcast协议在executor之间P2P分发,不再随每个task重复传输;二是executor拿到后会在本地缓存,同一个executor上的多个task共享这份副本。
需要注意的坑:
- 广播变量是只读的,你不能在executor端修改
model_bc.value的内容再期望driver拿到。 - Spark默认会在任务结束后自动清理广播变量,如果你在多个action里反复使用同一个UDF,可能遇到"broadcast已被unpersist"的报错,需要重新广播。
- 广播变量过大会导致executor内存被模型占满。我一般用2GB作为心理红线,超过1GB就开始考虑架构层面的方案(后面会讲方案六)。
- 模型对象必须可pickle。如果你的模型里有lambda表达式或自定义回调,先想办法移除它们,否则广播在创建阶段就失败。
3.2 方案二:mapPartitions,让加载次数从task数降到partition数
广播适合模型不大、推理逻辑简单的场景。但有些模型真的没法放进广播,或者广播后每个task还是要做大量初始化工作(比如初始化一个特征处理器、加载一个字典表),这时用mapPartitions更合适。
mapPartitions的核心思路是:给每个partition执行一次函数,函数内部可以只做一次加载和初始化,然后批量处理这个partition里的所有数据。
def predict_partition(rows): # 每个partition只加载一次模型 model = joblib.load("hdfs:///models/lgb.pkl") tokenizer = load_tokenizer() # 其他一次性初始化也放这里 for row in rows: features = tokenizer.transform(row["raw_text"]) pred = model.predict_proba(features.reshape(1, -1))[0][1] yield (row["id"], float(pred)) result = df.repartition(200).rdd.mapPartitions(predict_partition).toDF(["id", "score"])这段代码里,模型加载次数 = partition数,而不是task数。200个partition就只加载200次,相比逐行加载、逐task加载已经是数量级的优化。
但这里有个容易忽略的问题:每股partition的函数只执行一次,但如果这个partition的数据被物化到磁盘后又重新读取(比如shuffle写失败),函数可能被再次调用,模型就会再加载一次。所以mapPartitions里的初始化还是要尽量轻量,不要把启动时间拖到分钟级。
写的时候还有个细节:toDF需要显式指定列名和类型,如果你的推理结果要保留原始列,最好在yield时把需要保留的字段一起带出来,否则还要再做一次join,白白增加一次shuffle。
3.3 方案三:pandas UDF + 惰性单例,向量化推理的工程红利
mapPartitions虽然好,但它要求你把推理逻辑写成"基于行的迭代器"风格,处理批量特征时不够方便,向量化程度也不够。生产环境里我更喜欢用pandas UDF(也叫Vectorized UDF)来做推理,它有两个优势:一是通过Arrow批量传输数据,省掉了逐行pickle的开销;二是在同一个Python worker进程内,可以用模块级单例缓存模型,避免反复加载。
import pandas as pd from pyspark.sql.functions import pandas_udf # 惰性加载:只在第一次调用时真正加载模型 _model_cache = None def _get_model(): global _model_cache if _model_cache is None: _model_cache = joblib.load("hdfs:///models/lgb.pkl") return _model_cache @pandas_udf("double") def predict_prob_pd(features: pd.Series) -> pd.Series: model = _get_model() # features是pd.Series,每个元素是array X = np.vstack(features.to_numpy()) return pd.Series(model.predict_proba(X)[:, 1]) result = df.withColumn("score", predict_prob_pd(F.col("feature_vector")))关键点在于_get_model的惰性单例模式。pandas UDF运行在executor里的Python worker进程中,进程启动时会加载这个模块,第一次推理时触发模型加载,之后同一进程内所有批次都复用这一个模型对象,不再重复加载。这比mapPartitions更省心,因为你不需要关心partition边界,只需关心worker进程数量。
使用时有几个坑要留意:
- pandas UDF里函数的输入输出类型声明要准确,
"double"表示返回浮点数;特征列如果是array<double>类型,拿到的pd.Series每个元素是np.ndarray,要用np.vstack转成二维矩阵,否则predict_proba会报维度错误。 - Arrow的开启配置:
spark.sql.execution.arrow.pyspark.enabled要设为true,spark.sql.execution.arrow.maxRecordsPerBatch控制每个批次的记录数,直接影响单批推理的内存峰值。默认值可能偏大,建议根据你的特征维度调成500~2000。 - pandas UDF对Python worker进程的内存占用要求较高,因为每个批次都要把数据转换到pandas结构。如果executor内存紧张,可以把
maxRecordsPerBatch调小,以增加批次数为代价换取更低的峰值内存。
3.4 模型形态与序列化边界:能广播不等于适合广播
很多人在方案一和方案二之间纠结:模型到底该用广播还是mapPartitions?我的判断标准是:如果模型在50MB以内且推理频率高,优先广播;如果模型在50MB到500MB之间,用mapPartitions或pandas UDF惰性加载;超过500MB,直接考虑方案六的服务化部署。
有几个与模型形态相关的细节值得说道。第一,joblib.dump保存的pkl文件在加载时会重建整个对象图,如果模型内部有大量NumPy数组,加载时间很长,占用内存也大。这时可以考虑用model.save_model()这类原生格式(比如LightGBM的txt格式、XGBoost的json格式),加载时用原生API,内存占用更小。第二,如果你在一个Spark任务里同时用到多个模型(比如一个排序模型加一个过滤模型),尽量让它们走同一种加载方式,否则每个executor的内存里堆着多个模型,OOM风险成倍增加。第三,如果模型内部有sklearn的Pipeline,里面包含StandardScaler之类的转换器,务必确认所有转换器都能pickle;不能pickle的组件(比如自定义Transformer),想办法换成官方实现。
4. 方案四和方案五:把executor的结果"还"回driver的正确姿势
4.1 方案四:累加器,只做"写了多少条"这类监控
如果说广播是driver向executor传递数据的正向通道,那么累加器就是反向通道里最轻量的一种。它允许executor向driver累积更新一个值,但限制是:executor只能做"增加"操作,driver才能读取最终值。
processed_cnt = spark.sparkContext.accumulator(0) skipped_cnt = spark.sparkContext.accumulator(0) def process_row(row): if row["label"] is None: skipped_cnt += 1 # executor端只能 += return None processed_cnt += 1 return row df.rdd.map(process_row).count() print("processed:", processed_cnt.value) # driver端才能 .value print("skipped:", skipped_cnt.value)累加器非常适合做运行监控:统计处理条数、异常条数、缺失特征条数、耗时总和等。但有一个在生产里很要命的细节:task失败重试时,累加器会重复计数。比如某个partition的task执行到一半挂了,Spark会重新调度这个task,之前那段代码对累加器产生的更新不会被回滚。所以累加器只能当作"近似指标"来用,不能作为精确的业务计数。真要精确统计,用groupBy或agg来算。
累加器还有一个用途是给driver端发信号。比如某个executor检测到数据分布异常,想要中断整个任务,可以在executor端给累加器加一个大数,然后定期检查这个值。但这种"轮询中断"模式不够优雅,生产里我更推荐用spark.sparkContext.cancelJobGroup()配合自定义异常来处理。
4.2 方案五:foreachPartition + 外部存储,结果不要回driver
当数据量大到不能collect回driver,或者你需要把结果写进数据库、消息队列、对象存储时,首选foreachPartition。它的语义是:每个partition在executor端执行一次函数,函数内部拿到这个partition的所有行,可以批量写入外部系统。
def write_partition(rows): conn = psycopg2.connect(CONN_STR) batch = [] for row in rows: batch.append((row["id"], row["score"], row["ts"])) if len(batch) >= 500: insert_batch(conn, batch) batch.clear() if batch: insert_batch(conn, batch) conn.close() df.foreachPartition(write_partition)这个方案的精髓在于:外部系统取代了driver成为端到端数据汇集的终点。executor不再依赖driver来收集结果,driver只负责调度和监控。写入数据库时注意几个工程细节:
- 每个partition建立一个连接即可,不要每行都建连接,否则数据库连接池瞬间被打满。
- 批次提交能显著提升吞吐,但批次大小不是越大越好。我在PostgreSQL上实测,500~1000条一批比较稳妥,超过这个范围锁竞争和内存占用都会上升。
- 要考虑写入幂等性。如果某个task失败重试,这个partition的数据可能会被写两遍。要么在目标表上建唯一键,要么在insert语句里用
ON CONFLICT DO UPDATE,要么先把结果写到临时路径,全部成功后再原子rename到最终路径。
如果目标是HDFS或S3,更推荐直接用DataFrame的write方法,因为Spark原生写入方式能规避很多重复写问题。foreachPartition更多用于数据库、Redis、消息队列这类不能直接用DataFrame写入的系统。
4.3 如果一定要collect:先count再拉,拉完立刻释放
有些场景确实绕不开collect,比如要把模型评估指标汇总到driver端做可视化。这时我建议遵守三条纪律:
- 先
df.count()确认行数在你的driver内存承受范围内,再执行collect()。 - 只collect需要的列,不要
select("*")然后拉一堆大字段(比如原始文本、向量列)。 collect()之后立刻对driver端的引用赋值成局部变量,用完之后置None,避免长期占用堆内存。
顺带一提,toLocalIterator看起来是"懒加载"的,好像内存压力小,但它本质上是逐partition拉回driver,如果你的下游逻辑一定要全量数据,它不会比collect省多少内存。
5. 方案六和方案七:架构层面釜底抽薪,绕开跨节点依赖
5.1 方案六:把模型做成服务,Spark只做调度和ETL
如果模型大到广播放不下、加载时间动辄几分钟、或者你需要频繁更新模型而任务又不想停,那就不要在executor里加载模型了。把模型部署成一个独立的推理服务,PySpark通过HTTP或gRPC调用它,是最干净的解耦方式。这也是我在上一家公司最终采用的方案,原因很简单:把跨节点依赖变成了跨服务调用,属于网络问题,有一整套成熟的监控、重试、限流方案可以套用。
import requests INFER_URL = "http://ml-serving.internal:8080/v1/models/score:predict" def infer_batch(rows): payload = {"instances": [r["features"] for r in rows]} resp = requests.post(INFER_URL, json=payload, timeout=30) resp.raise_for_status() preds = resp.json()["predictions"] for row, pred in zip(rows, preds): yield (row["id"], float(pred)) df.repartition(200).rdd.mapPartitions(infer_batch).toDF(["id", "score"])注意这里我用的是mapPartitions而不是普通的map或UDF,目的是在partition级别聚合成一个batch请求,避免逐条调HTTP接口,否则网络往返会把吞吐拖垮。一次请求几十条甚至几百条样本,服务端返回对应的预测列表,效率会高很多。
服务化方案的运维要点:
- 推理服务要支持横向扩容。Spark的并行度一开就是几百个task,服务端必须能扛住几百并发。否则Spark任务会大量超时,看起来像"模型崩了"。
- 请求超时和重试策略要设计好。推荐连接超时3秒、读超时30秒,重试2次并加指数退避。重试时要注意幂等:服务端最好不依赖请求顺序和次数,只根据请求内容返回结果。
- gRPC比REST更适合大批量推理,因为protobuf的编码和解码效率远高于JSON。但gRPC的调试成本高一些,如果团队没接触过,从REST起步也可以。
这样改造后,Spark的任务不会因为模型更新而重启,模型发布和回滚完全由服务团队独立控制,PySpark的UDF里也不再有任何模型加载逻辑,自然就没有跨节点依赖的问题了。
5.2 方案七:用DataFrame和外部存储代替点对点回传
最后一类方案有点"意识流",但用得好非常省事。核心思路是:与其想着怎么把对象安全地在节点间传来传去,不如直接把状态物化成分布式数据集或外部存储,用Spark的算子去处理。
举个例子。假设你要在每个partition上训练一个小模型,然后把所有小模型的评估指标汇总到driver端。最初我是想用collect()把所有指标拉回来,后来改成这样:
def train_submodel(rows): X = np.vstack([r["features"] for r in rows]) y = np.array([r["label"] for r in rows]) model = train_small_model(X, y) yield (model.score(X, y), float(model.feature_importances_.mean())) metrics_df = df.repartition(16).rdd.mapPartitions(train_submodel).toDF(["acc", "importance"]) metrics_df.show() # 用DataFrame聚合,而不是collect到driver再手算你看,这里我们让每个partition产出一行指标,最终这些指标天然就是分布式DataFrame里的数据,可以直接用Spark SQL做统计。我们没有手动"把结果回传",而是让结果以数据的形式存在,由Spark引擎保证一致性——跨节点依赖被框架消化掉了。
另一个更常见的形式是:中间结果不要留在Spark内存里等下一个action用,而是直接写HDFS或S3,下一个job再读。这个模式看起来多了一次磁盘IO,但好处是任务之间彻底解耦,模型文件、特征表、预测结果都变成了稳定的中间产物,既方便排查问题,也能让不同团队各自运行自己的任务。
6. 七种方案的选型地图与最终排坑清单
6.1 一张表看懂7种方案的适用范围和代价
我把这7种方案放在一张表里,方便你根据实际情况快速选型。
| 方案 | 核心思路 | 适用场景 | 主要代价 |
|---|---|---|---|
| 1. 广播变量 | 模型在driver端序列化一次,executor本地缓存 | 模型<200MB,推理频繁,模型更新不频繁 | 广播数据量大时driver分发有压力;内存占用高 |
| 2. mapPartitions | 每个partition初始化一次模型/连接 | 模型较大,不能在executor间共享,适合pyfunc或复杂特征工程 | 加载次数=partition数,partition多时仍有开销 |
| 3. pandas UDF | 向量化推理 + worker进程内单例模型 | 特征列多、推理可用批量矩阵运算、内存可控 | 需要处理Arrow转换;worker内存占用高 |
| 4. 累加器 | 轻量级反向通道,只做增加操作 | 监控任务进度、统计异常条数、近似计数 | 不保证精确一次,task重试会重复计数 |
| 5. foreachPartition + 外部存储 | executor直接写数据库/消息队列/HDFS | 结果集大不能回driver,需要落库的ETL任务 | 外部系统的可用性、写入幂等性要额外设计 |
| 6. 模型服务化 | 推理服务独立部署,Spark远程调用 | 模型超大、更新频繁、需要服务级SLA | 引入网络开销,需要配套限流、熔断、监控 |
| 7. 分布式中间结果 | 状态物化为DataFrame或中间存储,用算子处理 | 需要聚合跨节点结果、需要分层解耦的任务链路 | 多一次IO,SQL表达有一定的迁移成本 |
选型时我一般按这个顺序问自己:模型能不能广播?能,用方案一。不能,模型能不能在partition级别加载?能,用方案二或三。都不能,模型能不能服务化?能,用方案六。如果问题不是模型,而是结果回传,那就看数据量——小数据用方案四或直接collect,大数据用方案五或方案七。
6.2 生产环境里我踩过最深的5个坑
这些坑不是每个项目都会遇到,但遇到了不处理好,会让你在集群上熬好几个通宵。
第一个坑:广播变量被Spark自动清理。我用广播模型跑了一个多stage任务,第一个stage推理没问题,第二个stage再用同一个broadcast时直接报Broadcast variable ... was unpersisted。解决方案是确认broadcast变量要跨多个action使用时,不要在stage结束后手动unpersist,也不要依赖Spark的自动清理。真遇到清理,就重新广播一次,或者把两个stage合并成一个。
第二个坑:pandas UDF的返回值长度对不上。有一次我写了个pandas UDF,内部对某些行做了过滤,返回的pd.Series比输入的短了几行,Spark直接报Result vector from pandas_udf was not the required length。后来养成了习惯:pandas UDF的输入输出必须一一对应,任何过滤、去重、采样操作都要在Spark DataFrame层面完成,不要塞进UDF里。
第三个坑:foreachPartition写入数据库导致连接风暴。最初我按默认并行度200跑写入任务,每个partition建一个连接,数据库瞬间被打挂。后来我强制先repartition(50),把写入任务的并行度降下来,同时在写入端加了连接池,问题才解决。记住:写入外部系统的并行度不等于Spark任务的并行度。
第四个坑:模型加载路径在executor上不可见。如果用mapPartitions加载HDFS上的模型,要注意executor端有没有配HDFS的core-site.xml和hdfs-site.xml,以及hdfs://协议是否被解析。最稳妥的做法是把模型路径做成参数传进SparkSubmit的--files选项,让模型文件随任务分发到每个executor的本地目录,然后用相对路径加载。
第五个坑:用累加器做精确计数导致数据对不上。这是我最肉疼的经历:用累加器统计推理失败的行数,由于某个stage重试了一次,计数翻了一倍,下游监控误判成生产事故。从那以后,所有精确统计我都改用DataFrame的agg,累加器只保留在"够用就行"的场景。
七种方案说完了,翻来覆去其实就一句话:PySpark机器学习生产化的核心难点,从来不是算法本身,而是怎么让模型、数据、状态在分布式环境下安全高效地流动。希望这篇文章能帮你少熬几个深夜。