PySpark调用大模型全指南:从UDF到mapPartitions并发实践
先说句实在话我刚接触这个需求的时候也被“分布式框架”和“大模型”这两个词搞得有点心虚。Spark 是个纯计算引擎擅长跑 SQL、跑 ETL、做特征工程大模型是另一套生态吃 GPU、吃显存、讲究服务化部署。这两者凑到一起第一反应往往是“能行吗”但搞完几个项目之后我可以直接给结论能行而且只要把调用方式、并发模型和容错策略想清楚它能比想象中稳得多。这篇内容主要面向两类人一类是已经用 Spark 在做数仓和数据分析现在想把语义理解、文本分类、信息抽取这类能力也塞进既有流程的工程师另一类是已经在大模型应用层写了不少 Python 脚本但一跑到“每天几千万条”这个量级就被卡死、想迁移到集群上跑的开发。两种视角我都经历过文章里的代码和结论均来自我自己踩过的坑不是教科书式的理论推演。1. 为什么非要把 Spark 和大模型凑到一块1.1 真正的业务场景不是炫技是被数据逼的过去一年我接手的大部分文本处理需求都长这样每天几千万条客服工单、用户评论、日志片段、商品描述要按不同维度做分类、打标签、抽取关键词、判断情感倾向。以前这类活主要靠关键词规则加正则再做点统计特征准确率卡在 70% 到 80%每次换一个业务场景就得重新洗数据、重新调规则非常磨人。大模型普及之后很多团队顺理成章想用模型来做这批活。单条文本丢给大模型返回分类结果或者摘要效果确实比规则好一个档次。但问题随之而来你拿一个 Python 脚本循环调用模型接口跑 100 条、1000 条都能忍耐跑到 100 万条无论是耗时还是费用都开始失控。再往上走到千万级单机脚本基本就是自杀式写法一台机器跑个两三天中途还随时可能被限流打断。这时候 Spark 的价值就出来了。Spark 的底层模型就是把大任务切分成小任务分发到多台机器的多个 Executor 上并行处理。大模型推理则是“把一段文本送出去拿回一段结果”的 IO 等待过程。把两者叠加本质上是用 Spark 的分布式调度能力去统一管理大模型调用的并发、重试、限流和成本。这也是“在 Spark 里调用大模型”这个需求的真正核心而不是简单地在 UDF 里塞一个 HTTP 请求。我见过太多团队翻车的路径几乎一样先单机脚本跑通然后直接塞进 Spark结果被序列化错误、限流、Executor 失联、重复计费轮番教育。这篇文章把方案和坑都梳理清楚至少能帮你少走两三个星期的弯路。1.2 Spark 与大模型天然“不合拍”三个绕不过的坎第一个坎计算模型不匹配。Spark 的任务设计是“短平快”的一个 Task 通常几秒到几十秒结束。但一次大模型 HTTP 调用从请求发出到拿到结果可能就要两三秒甚至更久。如果你在 UDF 里同步调用每个 Task 都卡在网络等待上Executor 虽然有几个核实际吞吐能力却低得可怜。第二个坎结果不确定性和 Spark 重试机制叠加。大模型输出天然不稳定同一个 prompt 可能这次返回 A下次返回 B还可能返回超长文本、JSON 片段、甚至是空内容。而 Spark 默认的容错方式是“Task 失败就重算”它不会管你上一次调用是否已经付费。如果不做好幂等和缓存失败重跑几次账单会很酸爽。第三个坎资源与成本换算。GPU 是稀缺资源大多数 Spark 集群压根没有 GPU。就算有也不是每个 Executor 进程都能直接访问 GPU 设备。于是你必须做抉择把推理请求通过网络发给一个独立的模型服务还是自己在 GPU 机器上部署推理服务然后让 Spark 通过客户端调用。无论是哪种Spark 本身都不承担推理计算它只负责“调度 数据搬运 结果汇总”。把这三点想透了再看市面上的各种方案思路就清晰了。2. 四种主流程调用方案一次讲清楚2.1 方案AUDF 里直接调 API这是最直白、也最容易踩坑的方案。大概写法是from pyspark.sql.functions import udf from pyspark.sql.types import StringType udf(StringType()) def classify_with_llm(text: str) - str: return call_llm_api(text)好处是代码量极少适合快速验证、跑小批量样本验证模型效果能否满足业务要求。坏处也很明确每个 Task 一次只处理一条数据完成一次网络调用。如果你有 20 个 Executor每个 4 核理论上最多 80 个并发请求而每个请求的耗时都在秒级数据处理速度完全被网络延迟拖死。而且这种逐行调用模式还特别容易撞上模型服务的 QPS 限制。一旦触发限流UDF 抛异常Spark 开始重试重试又换来更多限流最终整个 Job 卡成一个死循环。所以我的结论是方案A只适合几千条样本验证场景不适合任何有规模预期的生产任务。2.2 方案BExecutor 内批量并发请求方案B解决的就是方案A的并发问题。核心思路是不要让每个 Task 只处理一条数据而是让一个 Task 处理一个分区分区内把多条文本攒起来再用ThreadPoolExecutor同时发起多个模型请求全部返回后合并结果。同样的集群规模方案B的吞吐量能比方案A提升 5 到 10 倍。因为它把单条串行等待变成了线程级并发等待网络 IO 的空闲时间被充分填满。方案B是最适合绝大多数 Spark 批处理任务的答案。它不需要引入额外集群组件代码复杂度可控而且能精准控制并发数不容易打爆模型服务。2.3 方案C集群本地部署推理服务如果数据不能出内网或者调用外部 API 的成本实在压不住那就得走方案C在 Spark 集群能访问的同一内网区域部署推理服务比如用 vLLM、Ollama 把开源大模型跑在 GPU 节点上再通过内部 endpoint 提供服务。Spark 在这里退化为纯粹的客户端。方案C的好处很明显延迟低、数据不出域、模型版本和并发策略都掌握在自己手里。但它也有很高的运维成本。你得自己管理 GPU 节点、模型加载、并发队列、显存占用还要处理多套模型之间的资源隔离。对很多团队来说这不是一件轻松的事。另外需要注意本地部署开源大模型的效果未必比得上外部 API因为开源模型和商用模型的指令跟随能力、生成质量还是存在差距。它更适合的数据处理场景是“抽取、分类、标准格式化”这些任务对模型能力要求不高但对成本和延迟敏感。2.4 方案D调用模型厂商的批量接口现在很多模型厂商都提供“离线批量处理”能力。你只需要把数据组织成 jsonl 文件上传后发起一个批处理任务等它跑完再下载结果文件。在这个方案里Spark 只做前置的数据准备、清洗、去重以及后置的结果校验和回填。方案D的优点非常明显不占用在线服务并发成本通常比在线调用便宜一个量级而且由厂商内部调度稳定性高。缺点是时效性差T1 甚至数小时级别的等待是常态不适合实时性要求高的场景。如果你的业务允许延迟比如“每天凌晨处理前一天的数据”方案D几乎是性价比之王。2.5 方案对比总表方案实现复杂度吞吐量成本控制数据安全延迟推荐场景UDF 直连 API低低差依赖外部服务秒级并发低POC、小样本验证Executor 批量并发中中高中依赖连接方式分钟级吞吐高百万级批处理本地推理服务高高规模后可摊薄高低延迟生产长期稳定、敏感数据厂商批量接口低高好看厂商协议小时级离线 T1 场景3. 实战PySpark 调用大模型的完整流程3.1 先把运行环境打好开始写代码之前有几个环境细节我建议先确认。Spark 版本至少 3.xPython 3.8 以上。依赖库上需要requests用于 HTTP 调用并发部分用 Python 自带的concurrent.futures即可不需要额外安装。如果你是走spark-submit提交任务还要注意 Python 依赖的同步。最简单的做法是先在每台 worker 机器的 Python 环境里装好依赖或者用--py-files把依赖包一起提交上去。我遇到过一种情况Driver 端调通了Executor 端一跑就ModuleNotFoundError就是因为依赖只装在了提交任务的那台机器上。假设你的模型服务是标准的 OpenAI 兼容接口地址是http://model-service:8000/v1/chat/completions下面都按这个接口来写。如果你们模型服务用的是别的协议替换 endpoint 和请求体即可整体思路完全一致。3.2 写一个能扛故障的调用函数很多人写大模型调用的第一个版本往往只有一个requests.post()没设置超时没做重试失败全靠 Spark 兜底。我早期也这么干结果任务跑了一半一批 Executor 因为网络超时被 kill整个分区重算模型接口被重复调用了几万次费用直接超预算。现在我的习惯是所有外部调用先封装成“可重试 可观测”的函数。下面这个函数只做一件事调一次接口失败抛异常不做重试。重试逻辑放外层这样职责清晰便于测试。import os import requests MODEL_ENDPOINT os.getenv(LLM_ENDPOINT, http://model-service:8000/v1/chat/completions) TIMEOUT 60 def call_llm_once(prompt: str) - str: payload { model: qwen2.5-7b, messages: [{role: user, content: prompt}], temperature: 0.1, max_tokens: 512, } resp requests.post( MODEL_ENDPOINT, jsonpayload, headers{Authorization: Bearer os.getenv(LLM_API_KEY, )}, timeoutTIMEOUT, ) resp.raise_for_status() data resp.json() return data[choices][0][message][content]这里有个小经验temperature在批处理场景下调得非常低一般 0.1 甚至 0。你要的是稳定可复现的分类和抽取结果不是让模型发挥创意。另外max_tokens要设一个上限防止某条文本触发模型无限生成把耗时和费用都拉爆。3.3 从逐行 UDF 升级到 mapPartitions 并发有了单次调用函数接下来要解决“批量并发”。我强烈建议不要用逐行 UDF而是用mapPartitions在分区内部做批量处理。mapPartitions的含义是一次处理一个分区里的全部数据。分区内的数据是本地集合你可以自由地做缓存、批量请求、并发调度。from concurrent.futures import ThreadPoolExecutor, as_completed def call_llm_batch(prompts, max_workers8): results [None] * len(prompts) def single_item(idx, p): try: return idx, call_llm_once(p) except Exception as e: return idx, f__ERROR__: {e} with ThreadPoolExecutor(max_workersmax_workers) as pool: futures [pool.submit(single_item, i, p) for i, p in enumerate(prompts)] for fut in as_completed(futures): idx, res fut.result() results[idx] res return results def process_partition(rows): buffered list(rows) prompts [row[0] for row in buffered] # 假设第一列是文本 llm_results call_llm_batch(prompts, max_workers8) for row, res in zip(buffered, llm_results): yield row (res,)使用方式df spark.read.parquet(hdfs:///data/comments.parquet) result_df df.rdd.mapPartitions(process_partition).toDF([text, label, llm_result])这里有一个足够重要的经验值得单独强调分区内批量处理时不要一次性把分区所有数据无脑加载进内存。list(iterator)会把整个分区全部载入如果分区过大内存会爆。在实际生产里我会控制每个分区在一万条以内超额就先repartition()。Spark 默认分区数往往不是为“单分区内做批量 IO”设计的需要自己调。3.4 设置合理的分区与并行度Spark 的并行度由分区数决定分区数是任务调度的最小单位。分区太少Executor 大部分在空转分区太多又会产生大量调度开销。我的经验公式分区数 ≈ Executor 数量 × 单个 Executor 的内核数可以再乘 1.5 到 2 留出缓冲。假设你申请了 30 个 Executor每个 4 核理想并行度 120分区数设置在 180 到 240 之间比较合适。具体提交命令参考spark-submit \ --master yarn \ --deploy-mode cluster \ --num-executors 30 \ --executor-cores 4 \ --executor-memory 8g \ --driver-memory 4g \ --conf spark.sql.shuffle.partitions200 \ --conf spark.task.cpus1 \ llm_spark_job.py--executor-cores 4表示每个 Executor 分配 4 个内核spark.task.cpus1表示每个 Task 占 1 个内核这样单个 Executor 能同时运行 4 个 Task。很多环境里如果不显式配置spark.executor.cores默认只有 1 个可用核心这也是后来我在排障时频繁发现“Spark on YARN 只有 1 个 CPU”的关键原因。3.5 结果回写写入 Hive 或 Parquet 的设计要点推理结果拿到手之后别直接写回业务表。大模型返回的文本很不规矩可能带换行符、控制字符、前后空格甚至是不完整的 JSON。直接写入下游表跑数时解析出错排查起来非常痛苦。我的习惯是分两层落库。一层是原始返回raw_result完整保留模型输出的原始文本另一层是下游用的parsed_result在写入前统一做清洗和格式化。如果真的解析失败宁可让parsed_result为空也不要往里面塞半截 JSON。output_df result_df.select( id, text, F.col(llm_result).alias(raw_result), ) output_df.write.mode(overwrite).partitionBy(dt).parquet(hdfs:///data/llm_output)按业务日期dt做分区是个好习惯。一方面下游按天取数方便另一方面任务重跑只需要处理对应分区不用全表覆盖成本控制更加精细。4. 性能优化与稳定性设计4.1 用线程池把单 Executor 吞吐拉满前面代码里的max_workers8是我常用的默认值但不要照抄。这个数字的确定依赖两个约束模型服务的 QPS 上限和单次请求的平均耗时。举个例子。如果模型服务 QPS 上限是 100单次请求平均耗时 2 秒那么在途请求上限就是 200QPS 乘以耗时。假设 Spark 有 20 个 Executor平均每个 Executor 需要的并发就是 10。我通常会给 20% 缓冲开 8 个线程总并发 160既能逼近服务端能力上限又不会一窝蜂把服务压垮。线程池数量不是越大越好。如果开 50 个线程服务端排队堆积单个请求的 P99 延迟会快速飙升进而触发大量超时重试重试又进一步加剧排队最终雪崩。这里的关键是并发控制是在保护服务端也是在保护任务本身。4.2 缓存与幂等防止重复计费的内服药大模型按调用量计费或者至少是配额受限的。Spark 的容错机制是“失败重跑”而重跑必然伴随重复调用。为了避免重复付费必须加一层“请求指纹缓存”。最简单的做法在调用函数里维护一个缓存key 是hash(prompt model 参数)value 是上一次返回结果。如果一个文本在缓存中命中直接返回不再发起请求。但分布式环境里的内存缓存只对当前 Executor 有效跨 Executor 不共享。更有效的做法是在 Spark 任务前置阶段做“输入去重”。把原始 DataFrame 按文本列去重只对去重后的样本做推理再把结果join回原始表。这个操作在 Spark 里就是一次groupBy成本很低但能直接降低推理次数。对于评论分类、工单打标这类场景重复文本的比例往往不低去重省下的费用常常让人惊喜。4.3 限流、重试与熔断三板斧外部模型服务一定有限流。最常见的表现是 HTTP 429也可能是连接被重置。如果请求端不做主动限流一股脑打过去服务端限流之后任务大批量失败Spark 开始重试重试又引发新一轮限流这就是典型的“重试风暴”。限流的核心是在客户端控制并发加一个信号量或令牌桶。遇到明确限流和 5xx 错误时做指数退避重试import time import random def call_llm_with_retry(prompt, max_retry4): for attempt in range(max_retry): try: return call_llm_once(prompt) except requests.exceptions.HTTPError as e: status e.response.status_code if status in (429, 500, 502, 503): sleep_time 2 ** attempt random.uniform(0, 1) time.sleep(sleep_time) continue raise raise RuntimeError(f请求失败重试 {max_retry} 次后仍然失败)熔断是另一层保护。如果连续失败次数超过阈值比如 1 分钟内连续 20 次 5xx直接暂停模型调用 60 秒返回预置兜底结果同时把异常行落表等待人工检查和补跑。没有熔断的调用链在高并发故障场景下就像刹车失灵的车很吓人。4.4 监控别等任务挂了才知道出了问题推理任务是典型的 IO 密集型任务最容易出现“Job 看起来在跑实际一个 Task 都推不动”的假死状态。不监控很难定位问题。我在生产里会埋这五类指标打到 Prometheus 或 PushGatewayQPS每秒实际发出的请求数成功率成功请求占总请求的比例平均延迟与 P95 延迟重试次数分布熔断触发次数。有了这些指标你才能回答“要不要扩容”和“为什么任务又卡住了”这类问题。单靠stages页面看到 task 一直 RUNNING根本定位不到根因。5. 常见问题与排障实录5.1 Spark on YARN 上 CPU 只有 1 个我收到过很多类似的提问spark-submit里明明写了--executor-cores 4为什么 YARN 控制台上每个 Executor 只有一个 vCore这个问题的背后一般有三个原因。第一个原因就是没设置spark.executor.cores。Spark 在 YARN 模式下不显式配置这个参数时默认值是 1。如果在提交命令里显式加上--executor-cores 4问题大概率解决。第二个原因是 YARN 队列容量限制。管理员给某个用户或队列设置了 vCore 上限你申请的并行度超过队列上限之后队列不会报错而是默默把分配值降下来。这种情况只能看 YARN 调度日志和队列配置改 Spark 参数没有用。第三个原因是分区数太少。哪怕 Executor 有 4 核每个 Task 占 1 核如果整个 Job 只有 30 个分区那么只有 30 个 Task 同时运行大量核心依然空闲。注意看 Spark UI 里的 Active Tasks 数量通常能立刻发现问题。5.2 PySpark 序列化报错Could not serialize objectPySpark 在提交任务时会把 UDF 函数连同闭包一起序列化到 Executor。如果你的函数内部引用了无法 pickle 的对象比如requests.Session、threading.Lock、数据库连接就会报Could not serialize object或Cannot pickle。解决思路是懒加载不要在函数外定义连接对象而是在mapPartitions内部第一次使用时才创建。每个 Executor 进程保持自己的连接池既避免了序列化问题又能复用 TCP 连接效果更好。简单说就是连接对象在线程函数内部创建而不是在 Driver 端创建。5.3 上游 API 限流与网络抖动遇到 429 限流我会先看重试是否生效。如果重试次数正常再判断是不是初始并发设置得过高。如果是网络抖动导致连接重置需要确认 TCP 连接是否复用、是否有CLOSE_WAIT堆积。还有一个容易被忽略的现象模型服务在长时间运行后内存碎片和显存碎片增加响应延迟慢慢变高。这种渐进式恶化光靠监控平均值是看不出来的要看 P95/P99 趋势。一旦发现趋势持续走高就该考虑服务重启或扩容。5.4 Spark 任务超时和 Executor 被 kill推理请求如果慢到超过 Executor 心跳上报周期Spark 会判定 Executor 失联直接将其 kill。分区内的任务还没跑完就没了整个分区的数据都得重来代价极高。应对办法有三个层面。第一把单次请求的超时时间调短宁可重试也不要让 Executor 无限等待第二单个 Task 处理的数据量要控制避免一个慢请求拖垮整个 Executor第三在模型服务端设置排队上限如果排队时间超过阈值直接拒绝新请求让客户端快速失败重试而不是傻等。5.5 结果顺序错乱与缺失使用ThreadPoolExecutor时as_completed的返回顺序不是提交顺序。我在前面的代码里特意用results[idx] res按下标回填就是为了避免结果对应错位。如果不用下标记录靠list顺序去 zip一旦有请求失败结果和输入就对不上下游数据会全面错乱。另外要注意如果某个请求失败返回了__ERROR__占位符后续解析必须显式判断不要直接把它当正常结果写进数仓。5.6 常见问题速查表问题表现排查重点常用解法CPU 只有 1 个Executor 资源申请异常spark.executor.cores / 队列限制显式配置提交参数序列化失败Job 提交阶段直接报错闭包内是否有不可 pickle 对象连接对象在 Executor 内懒加载限流 429大量任务失败重试服务端返回头 / 埋点指标退避重试 信号量控制Executor 被 kill任务偶发全挂心跳超时 / 内存溢出缩短请求超时、减小分区结果错乱下游数据对不上并发写回顺序按下标回填、任务内排序6. 选型建议和我的实操心得6.1 到底选哪种方案我判断方案只看三条数据量、数据敏感度、时效性。数据量在十万条以下用方案A快速验证完全没问题到了百万级以上方案B 是底线方案C 是更稳的选择。数据不能出内网就别考虑外部 API老老实实在内网部署模型服务为了省事把数据往外送一旦出问题后续沟通成本远高于那点算力费用。时效性要求 T1 可接受优先看厂商批量接口需要小时级甚至分钟级结果本地推理服务加并发调用才是正解。开源大模型和商业 API 之间怎么选我的经验是对模型输出质量要求极高的场景商业 API 更稳对成本极敏感、文本内容相对规整的场景本地部署开源模型更划算。你有没有 GPU 资源决定了这个选择的自由度。没有 GPU 就老老实实走 API硬着头皮上 CPU 推理速度会让你怀疑人生。6.2 绕开这五个坑至少节省 30% 开发时间如果你今天开始动手下面五个细节建议先想清楚所有外部调用参数包括 endpoint、模型名、超时时间、最大重试次数全部走配置中心或者环境变量不要写死在 Python 代码里。模型版本一升级改配置比改代码安全得多。任何推理结果都要保留原始返回字段。只存清洗后的值一旦后续需求变了或者发现清洗逻辑有 bug还得重新跑一遍大模型那个费时费钱。跑大规模任务之前先抽样 1% 数据做全链路验证确认输出格式没问题再放开全量。抽样这一步花费的时间不超过十分钟但能避免整个大任务毁于一个字段解析。任务开始前先做输入去重。如果你的数据里有大量重复文本这一步能直接减少推理成本。业务上如果允许丢重复优先做如果不允许也建议对重复项只推理一次再广播结果。建一张“推理调用日志表”每次调用记录输入 hash、时间戳、结果、状态码、耗时。出问题时这张表就是你的账本和诊断依据。我做过的项目里成本失控和任务不稳定往往不是哪一个巨大的架构错误而是这些不起眼的细节一个一个累积出来的。日志、缓存、重试三板斧打牢Spark 调用大模型没有想象中那么玄。最后分享一个小技巧即使你用了最稳的方案也一定在代码里保留一个“失败样本落表”的逻辑把每一条调用失败的数据单独保存下来这样出了问题不用全量重跑精准补数就够了。踩过几次坑之后你会发现真正可靠的批处理推理系统功夫都在 Spark 之外在调用设计、监控和成本审计这些地方。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →