1. 场景先行:什么时候你才真的需要“Spark 调大模型”
先讲个我这边真实发生过的事。去年底我们接了一个需求:把电商平台上近一年的用户评价全部过一遍大模型,做情感分析、标签抽取、投诉原因归类。数据量不大不小,大概两亿多条短文本。一开始团队里有同学说“这还不简单,写个 Python 脚本,requests 调 API 就行”。结果一算:单机脚本一条请求平均 800 毫秒,两亿条要跑一年多。哪怕把并发拉到 32,也要上百天。最后换成了 Spark 集群 + 大模型推理服务的架构,压缩到几个小时跑完。那会儿我才真正意识到,所谓“Spark 调大模型”不是一个炫技题目,而是数据规模到一定程度之后被逼出来的必答题。
那什么场景下你会碰到这种需求?我归纳为三类。
1.1 批量文本清洗与字段抽取
这是最常见的一类。业务数据通常是脏的、非结构化的,例如客服工单里套着订单号、用户ID、商品型号、退回原因,甚至还有情绪词。用正则写规则能覆盖 60% 的情况,剩下 40% 就得上大模型。再比如用户评价里抽“物流速度”“商品质量”“客服态度”三个维度的观点句,用传统 NLP 管线做一遍分词、依存句法分析,调参能调到怀疑人生。而大模型只需要一个 prompt,加一个 JSON schema,输出结构化结果,准确率还比规则高不少。
这类需求的基本特征:数据量大、单条推理计算量不大、结果结构固定。非常适合 Spark 批量处理。
1.2 非结构化数据转结构化数据
日志解析、文档解析、OCR 结果清洗,都属于这一类。我见过一个项目,要把几千份 PDF 合同转成结构化表格——甲方、乙方、金额、日期、违约条款。你能用 pdfplumber 把文字抽出来,但抽出来之后怎么映射到字段,那就是另外一回事了。传统方案是写一堆 if-else 加正则,合同版本一变就崩。后来用大模型做“文档摘要 + 字段抽取”,每条记录调一次模型,输出 JSON,Spark 端收到 JSON 直接拍平,一张干净的 DataFrame 就出来了。
还有一类是从嵌入式日志里提取异常事件。日志不是给人读的,但大模型能读。你要做的就是把一堆堆原始日志通过 Spark 读进来,切成固定窗口,交给大模型去判断“这段日志说明服务发生了什么”。
1.3 大规模数据增强与伪标签生成
机器学习团队大概率会用到。做分类任务,标注数据不够,拿大模型生成一批增强样本;做检索任务,拿大模型为 query 生成多个改写版本;做半监督,拿大模型给无标签数据打伪标签。这些场景的本质是:大模型扮演一个“不需要睡觉的标注员”,Spark 负责把海量数据喂给它,再把结果收回来。
1.4 千万别说“这种情况”你也想用 Spark
反过来也提醒一下:如果你的数据只有几千条,或者对延迟要求很高(每个请求必须 200ms 内返回),那别用 Spark。Spark 的调度开销、序列化开销、任务启动开销在大批量小文件面前会吃掉很多收益。这种实时轻量场景,老老实实写个异步服务,或者用 Flink / Kafka Streams 做流式处理。Spark 适合的是离线、批量、稳定吞吐的任务,不是在线推理网关。
2. 方案对比:直连 API、本地推理服务、GPU 原生推理,谁更适合生产
“在 Spark 里调用大模型”这句话其实掩盖了至少四种完全不同的技术路线。我在不同项目里把这几种都试过,各有各的坑。
2.1 方案 A:在 UDF 里直接调云端大模型 API
最直接的想法,就是写一个 PySpark UDF,里面用 requests 去请求 GPT 或 Claude 这类云端 API:
from pyspark.sql.functions import udf from pyspark.sql.types import StringType import requests def call_llm(text: str) -> str: resp = requests.post( "https://api.xxx.com/v1/chat/completions", headers={"Authorization": "Bearer sk-xxx"}, json={ "model": "gpt-4o-mini", "messages": [{"role": "user", "content": text}], "temperature": 0 }, timeout=30 ) data = resp.json() return data["choices"][0]["message"]["content"] llm_udf = udf(call_llm, StringType()) df = df.withColumn("llm_result", llm_udf(df["text"]))这段代码跑通容易,但上生产等于裸奔。首先是限流问题:云端 API 都有 QPS 限制,Spark 默认一启动就是几十上百个并发任务,每个任务都在调 API,配额瞬间打满,429 满天飞。其次是超时问题:UDF 里一个请求 30 秒超时,如果一个 partition 里有一条卡住,整个 task 就要等 30 秒,集群里有几万个这样的 task,任务总时长直接被放大。第三是成本失控:UDF 是逐行调用的,一旦遇到重复数据(比如多张表 JOIN 之后同一文本出现多次),同一个 prompt 会被重复计费。
方案 A 比较适合的场景是:数据量小、一次性探索、不追求稳定。真要上生产,最少也得配套做结果缓存、失败重试、并发控制,这就把代码复杂度拉上去了。
2.2 方案 B:Spark 接本地推理服务(Ollama / vLLM / TGI)
这是我现在最推荐的生产方案,思路很简单:不要让每个 Spark task 都去直连远端,而是把大模型推理服务收敛为一个独立的中间服务,Spark 作为客户端去调用它。
本地推理服务可以是:
- Ollama:部署门槛最低,一条命令拉起模型,适合中小数据量和验证阶段;
- vLLM:吞吐量优势明显,支持 PagedAttention、continuous batching,适合高并发批量推理;
- TGI(Hugging Face Text Generation Inference):功能完整,内置量化、流式输出,适合和已有推理体系集成。
架构上大概长这样:
Spark Cluster <--HTTP/GRPC--> 推理服务(vLLM/Ollama, 部署在 GPU 节点)这样做的好处非常直接。第一,内网调用延迟比公网低一个数量级,而且没有公网 API 的 QPS 限制,你能根据自己的 GPU 资源精准控并发。第二,推理服务独立于 Spark 集群,Spark job 挂掉不影响模型服务,模型升级也不需要重启 Spark。第三,成本可控,模型跑在自己的卡上,按 token 计价这件事就彻底消失了。
说个具体案例。我们有个任务要给一亿张商品图片的标题生成短描述,用的就是 Spark + vLLM。vLLM 部署在一个 4 卡 A100 节点上,Spark 集群 20 个 executor,每个 executor 开 8 个线程并发请求,实测吞吐量稳定在每秒 150 个请求左右,整个 job 跑完不到一天。要是走云端 API,这个量级的成本会非常难看。
2.3 方案 C:在 Spark Executor 内部直接跑 GPU 推理
这个方案从理论上说最“Monolithic”——把模型放进 Spark 进程里,一个 pipeline 搞定。我试过在 PySpark 的 mapPartitions 里加载 transformers 模型,然后用 GPU 批量推理。原以为这样能避免网络开销,结果踩了一堆坑:
- executor 数量和 GPU 数量很难匹配。Spark 的调度单位是 executor,它不知道你的 GPU 长什么样。你让 10 个 executor 都在跑,每个都想 loading 同一个 7B 模型,显存直接爆掉。
- 模型加载时间被反复浪费。每个 executor 启动时要重新加载权重,7B 模型在 HDD 上加载一次就要几分钟,整个 job 光加载模型就占了三分之一时间。
- 后续维护困难。模型升级要重新构建整个 Spark 环境,依赖冲突能把人逼疯。
结论很明确:Spark 是分布式数据处理框架,不是模型推理引擎。让它直接和 GPU 驱动打交道,就好比让你家的路由器去算天气预报,能算,但没必要。如果非要在这个方向走,更合理的是搭配 NVIDIA 的 Spark-RAPIDS 系列工具,让 GPU 承担 Spark 内部的算子加速,而不是把模型塞进去推理。
2.4 方案 D:通过向量化或嵌入 API 做中间层
有些场景你其实不需要大模型的完整“智能”,只需要一个向量表示,比如做相似度搜索、聚类、去重。这种时候可以不用调用 chat/completions 接口,而是用 Embedding API 或者本地 embedding 模型(例如 SentenceTransformer、BGE 系列)把文本转成向量,再落回向量数据库或特征存储。
这个方案的好处是:embedding 模型参数量小(通常 300M 到 1B),推理速度快,资源占用低;坏处是:功能单一,只能做特征抽取,不能做生成。如果业务目标是“把文本转成向量存起来供后续检索”,这个方案比硬调大模型省钱省力得多。
2.5 方案对比表与选型结论
我把四种方案的核心差异整理成一张表:
| 方案 | 延迟 | 吞吐量 | 成本 | 运维复杂度 | 适用场景 |
|---|---|---|---|---|---|
| A 云端 API 直调 | 高 | 低 | 高(按token计费) | 低 | 小数据量、一次性探索 |
| B 本地推理服务 | 低 | 高 | 中(GPU机器成本) | 中 | 生产环境、批量推理 |
| C Executor 内推理 | 低 | 低 | 低 | 极高 | 不推荐 |
| D 嵌入模型/向量化 | 低 | 高 | 低 | 低 | 向量检索、聚类、去重 |
我的选型结论很明确:凡是生产环境,一律优先方案 B。如果预算紧张或者业务刚起步,可以用 Ollama 代替 vLLM;如果数据量很小且只想快速验证,方案 A 也不是不能用,但要记得做缓存和限流。方案 C 我劝你直接放弃,方案 D 只在你明确知道“我只需要向量”时使用。
3. 从逐行调用到批量推理:一个重构实录
选定方案 B 之后,真正的工程挑战才刚开始。很多人会写一个 UDF 往推理服务发请求,逻辑没毛病,但性能和稳定性差到离谱。下面我把我自己从逐行调用一步步演进到批量推理的完整过程写出来。
3.1 第一步:常规 UDF,能跑,但别这么上
第一版代码通常长这样:
from pyspark.sql.functions import udf from pyspark.sql.types import StringType, StructType, StructField import requests def ask_model(content: str) -> str: resp = requests.post( "http://llm-service:8000/v1/completions", json={"prompt": content, "max_tokens": 512}, timeout=30 ) return resp.json()["choices"][0]["text"] df_result = df.withColumn("answer", udf(ask_model, StringType())(df["content"]))注意,这个 UDF 是逐行调用的。Spark 拿到一个 partition,就会在该 partition 内逐条执行 UDF,等于一个 task 内部串行调用 N 次 HTTP。你就算把集群扩到一百个 executor,每个 executor 内部依然是串行,整体吞吐量被 HTTP 往返时间死死卡住。
更糟的是,requests 每次都会新建一个 TCP 连接。HTTP 握手 + TLS 握手的开销占了大头,模型本身根本没在忙。实测下来,这种写法单 executor 每秒只能完成 2-3 个请求,和直接用单机脚本没什么区别。
3.2 第二步:用 mapPartitions 替代 map,把串行改成小批量
正确的做法是用mapPartitions。它的意思是对整个 partition 的数据一次性处理,这样你可以:
- 复用 HTTP 连接(Session 复用 TCP);
- 将多条 prompt 合并成一个请求发给推理服务(如果你的推理服务支持 batch);
- 在 partition 内部做并发。
先看简化版:
import json import requests from concurrent.futures import ThreadPoolExecutor, as_completed def process_partition(rows): rows = list(rows) # 将 partition 内的行分组,每 20 条一组 for i in range(0, len(rows), 20): batch = rows[i:i+20] texts = [row["content"] for row in batch] with requests.Session() as session: resp = session.post( "http://llm-service:8000/v1/batch_completions", json={"prompts": texts, "max_tokens": 512}, timeout=120 ) if resp.status_code != 200: raise RuntimeError(f"batch request failed: {resp.status_code}") batch_results = resp.json()["choices"] for row, result in zip(batch, batch_results): yield (row["id"], row["content"], result["text"]) rdd_result = df.rdd.mapPartitions(process_partition)这样写之后,一个 partition 内的 1000 条数据只需要发起 50 次 HTTP 请求(假设每批 20 条),比逐行调用少了整整一个数量级的网络往返。如果你的推理服务不支持 batch 接口(Ollama 默认不强制支持,vLLM 支持 OpenAI 兼容的 batch),那就退而求其次,在 partition 内部用线程池并发发单个请求,效果也不错。
3.3 第三步:partition 内用 ThreadPoolExecutor 做并发
很多推理服务支持并发请求,但 Spark 一个 task 是单线程的。为了把 partition 内的并发能力也利用起来,我一般会在process_partition内部维护一个线程池:
from concurrent.futures import ThreadPoolExecutor, as_completed import requests def call_single(session, text): resp = session.post( "http://llm-service:8000/v1/completions", json={"prompt": text, "max_tokens": 256}, timeout=60 ) resp.raise_for_status() return resp.json()["choices"][0]["text"] def process_partition(rows): rows = list(rows) texts = [row["content"] for row in rows] ids = [row["id"] for row in rows] results = [None] * len(rows) # 保序 with requests.Session() as session: with ThreadPoolExecutor(max_workers=8) as executor: future_to_idx = { executor.submit(call_single, session, text): idx for idx, text in enumerate(texts) } for future in as_completed(future_to_idx): idx = future_to_idx[future] results[idx] = future.result() for idx in range(len(rows)): yield (ids[idx], texts[idx], results[idx])有几个细节值得说。
为什么用 ThreadPoolExecutor 而不是纯并发 requests?因为我们是 I/O 密集型任务,等待 HTTP 响应时 CPU 是空闲的。线程切换成本低,8 个线程就能把一个 GPU 推理服务的吞吐打满。为什么线程数取 8?这个不是拍脑袋定的。先看推理服务能接受多少并发,再看网络往返时间和 GPU 计算时间的比例。一般 4-8 线程是一个不会出错的起步值,后续再根据实际瓶颈调。
为什么 yield 的时候按原顺序返回?因为 Spark 后续的处理逻辑可能依赖原始顺序(比如 join 回业务表)。用results = [None] * len(rows)存结果,最后按索引输出,保证顺序稳定。
3.4 第四步:把重试、超时、熔断封装成公共方法
一旦进入生产,你会发现网络请求永远不会一次成功。推理服务过载、网络闪断、prompt 里有特殊字符导致服务端解析异常,各种问题都有。所以我会把调用逻辑封装成一个带重试和退避的公共函数:
import time import random import requests class LLMServiceClient: def __init__(self, endpoint, max_retries=3, base_delay=1.0): self.endpoint = endpoint self.max_retries = max_retries self.base_delay = base_delay self.session = requests.Session() def complete(self, prompt, **kwargs): for attempt in range(self.max_retries): try: resp = self.session.post(self.endpoint, json={ "prompt": prompt, **kwargs }, timeout=60) if resp.status_code == 429: # 限流,等待更长时间 time.sleep(self.base_delay * (2 ** attempt) + random.uniform(0, 1)) continue resp.raise_for_status() return resp.json()["choices"][0]["text"] except requests.exceptions.Timeout: # 超时重试,但要防止雪崩 time.sleep(self.base_delay * (2 ** attempt)) except requests.exceptions.ConnectionError: time.sleep(self.base_delay * (2 ** attempt)) raise RuntimeError(f"LLM service call failed after {self.max_retries} retries")这里的“指数退避 + 抖动”是关键。指数退避保证不会把服务端打得更死;加随机抖动是为了防止多个任务同时重试导致波形叠加。这个细节很细,但生产环境有没有抖动,区别非常大。
4. 性能调优与成本控制:把并发、超时、缓存讲透
很多人在 Spark 里调大模型,第一反应是“我加 executor 数量”。实际上瓶颈往往不在 Spark 端,而在推理服务端和网络端。下面把几个真正决定性能的变量拆开讲。
4.1 并发数怎么算:别让推理服务被打死
假设你的推理服务(比如 vLLM)最高能承受 32 并发。那么整个 Spark 侧的有效并发数就必须控制在 32 以内。否则多出来的请求全部排队超时,白白浪费时间。
有效并发数的估算公式是:
有效并发数 ≈ 活跃的 executor 数量 × 每个 executor 内的线程数假如集群有 20 个 executor,每个 executor 内 partition 并行度为 1,每个 partition 内我们开了 8 个线程,那有效并发就是 20 × 8 = 160,远超推理服务的 32。这时候你会看到服务端延迟急剧增加,客户端超时率飙升。
解决办法有三种:
- 把线程池大小降到 1-2;
- 减少 executor 数量(不划算,浪费资源);
- 加服务端限流逻辑,在推理服务前面加一个并发信号量,超过即拒绝。
我实际用得最多的是第三种。因为线程池大小这东西在 Spark 里不是全局可调的,改起来麻烦。更好的是在推理服务前面加一个简单的限流模块,比如用 Redis 计数器或者服务端信号量,保证服务端永远不会被打爆。
4.2 超时设置:太短会误杀,太长会拖死任务
调用大模型接口的超时时间设置是个技术活。太短,模型生成长文本时容易被误判为超时;太长,网络故障时整个 task 卡在等待上。
我的经验值:
| 参数 | 建议值 | 理由 |
|---|---|---|
| 单请求超时 | 60-120 秒 | 大模型生成 token 需要时间,7B 模型生成 512 token 大约需要几十秒 |
| 重试次数 | 2-3 次 | 再多会放大雪崩风险 |
| 指数退避基础值 | 1-2 秒 | 太短等于没退避,太长浪费等待时间 |
| 最大退避时间 | 30 秒 | 超过就放弃,等下一轮重跑 |
另外注意,requests的timeout要区分connect和read。我一般这样设置:
session.post(url, json=payload, timeout=(5, 120))5 秒是建立连接的超时,120 秒是等待响应的超时。这样网络不通时能快速失败,模型生成慢时又不会被误杀。
4.3 结果缓存:一次性算完,重跑别再付账
生产环境里,Spark 任务不可能每次跑成功。代码有 bug、数据有新版本、模型参数要调,任务重新跑是常态。如果不做缓存,每次重跑都要重新调用大模型,时间和钱都白花了。
我常用的做法是:把推理结果落地成一个明细表,用输入文本的哈希值作为主键。每次推理之前先查表,命中了直接用历史结果,没命中才调模型。
from pyspark.sql.functions import sha2 df_with_hash = df.withColumn("text_hash", sha2("content", 256)) # 先和已有结果做 left anti join,只处理新增或变化的数据 new_df = df_with_hash.join( cached_results.select("text_hash"), on="text_hash", how="left_anti" ) # 跑推理,写入结果 new_results = new_df.rdd.mapPartitions(process_partition).toDF(...) new_results.write.mode("append").saveAsTable("llm_cache")这样重跑时只有新增数据会被模型处理,已经算过的直接读表。实测下来,迭代 5 轮优化 prompt,实际模型调用量只有第一轮的 1.2 倍。
4.4 成本估算:别等到月底账单出来才后悔
最后说一下成本。用云端 API 时,成本可以这样粗算:
单条成本 = 输入 token 数 / 1000 × 输入单价 + 输出 token 数 / 1000 × 输出单价 总成本 = 单条成本 × 总条数假设一亿条数据,平均每条 300 token 输入、100 token 输出,用某云端模型,输入 0.15 美元/百万 token,输出 0.6 美元/百万 token,总成本大约:
输入成本 = 1e8 × 300 / 1e6 × 0.15 = 4500 美元 输出成本 = 1e8 × 100 / 1e6 × 0.6 = 6000 美元 合计 = 10500 美元这个数字已经不小了。这也是为什么我强烈推荐本地推理服务——自建 GPU 节点(哪怕是租的)跑大模型,一次离线批量推理的成本可能只是云端 API 的十分之一。当然,本地有维护成本,需要权衡。但如果你的业务量让我上面这个量级,本地推理几乎一定是更优解。
5. YARN 资源调度与大模型任务的资源协调
接下来聊一个偏离“推理逻辑”但直接影响任务成败的话题:资源调度。很多人在 Spark 上跑大模型调用时莫名其妙发现任务卡慢、OOM、或者明明集群有几十个核但 CPU 利用率上不去。我盘点几个高频坑。
5.1 “为什么每个 container 只分配一个 vcore?”的排查过程
网上经常看到这样的问题:Spark on YARN 起来之后,每个 executor 只占了一个 vcore,集群里明明还有很多核闲置。很多人第一反应是“YARN 调度器有问题”,其实大概率是你没有配置 Spark 的 executor 核数。
YARN 默认的行为是:一个 container 代表一个执行单元,如果 Spark 没告诉你每个 executor 要多少核,它就给一个默认值。在 Spark 2.x/3.x 中,spark.executor.cores的默认值在 YARN 模式下是 1。也就是说,每个 executor 只会向 YARN 申请 1 个 vcore,哪怕你的机器有 32 核。
解决的方法很简单,提交任务时加上:
spark-submit \ --master yarn \ --deploy-mode cluster \ --executor-cores 4 \ --num-executors 20 \ --executor-memory 16g \ --conf spark.yarn.executor.memoryOverhead=2g \ job.py--executor-cores 4表示每个 executor 申请 4 个 vcore,结合 Worker 节点的物理核数合理设置。如果你发现集群总核数是 64,executor 核数设为 4,最多能起 16 个 executor(理论上限,还得结合内存看)。这样 CPU 资源才不会被白白浪费。
但这里也有个反向坑:executor 核数设太大,并发度反而下降。设成 4 意味着每个 node manager 最多能容纳几个 executor 受核数限制,同时 executor 内部的任务并行度也会过高,导致大量任务同时请求推理服务,把服务端打爆。所以--executor-cores别贪多,4 是稳妥值。
5.2 Spark 内存模型与推理结果处理
再讲内存。Spark 的内存分配分为执行内存(execution)和存储内存(storage),默认由spark.memory.fraction控制。当你在大模型场景下处理推理结果,有几种情况会让内存爆掉:
- 推理结果很长(比如 2048 token),DataFrame 里一个字段几 KB,一亿行就是几百 GB;
mapPartitions里你把整个 partition 都list(rows)出来,partition 太大时单 task 内存直接飙高;- 推理结果里包含大 JSON 字符串,后续做解析和展开时产生大量中间对象。
我的建议是:
- 尽量推理完直接输出精简字段,把模型的“长篇大论”在 service 端就截断或结构化,不要全量落回 Spark;
- 合理设置
spark.executor.memory和spark.yarn.executor.memoryOverhead。用 PySpark 时,Python 进程的内存不归 JVM 管,所以要额外为 PySpark 进程预留内存,memoryOverhead 通常要设到 2-4G,否则 Python 端容易因为内存不足被 YARN 杀死; - 如果实在需要处理大量长文本,考虑在
mapPartitions里直接压缩结果(gzip)再落盘,减少 shuffle 和存储压力。
5.3 GPU 节点与 Spark 集群的拓扑关系
最后强调一个架构层面的认知:推理服务节点不一定要和 Spark 集群在同一个 YARN 集群里。
如果推理服务只部署在某个 GPU 节点上,而 Spark 集群分布在不同机器,那么所有请求都汇聚到那个 GPU 节点的网络入口,网络带宽会成为新的瓶颈。我之前遇到过 20 个 executor 同时向一个 vLLM 服务发请求,服务端 GPU 利用率只有 60%,但网卡先打满了,延迟从 200ms 涨到 2s。
解决方案要么是把推理服务部署到多个 GPU 节点,前面加负载均衡;要么把 Spark 集群和推理服务放在同一个机架内,走内网高速网络。如果你用的是 NVIDIA DGX Spark 这类个人 AI 工作站/桌面级超算设备做推理,可以把 Spark 的 Driver/Executor 调度到同一台机器,减少网络跳数,但这种模式更适合小规模实验,大规模生产还是得靠横向扩展 GPU 节点。
6. 生产环境踩过的坑和排查思路
理论讲完,最后聊几个实战中遇到的具体问题。每个问题我都会从表象、排查链路、根因、解决手段四个维度讲,希望能给你一个完整的排查思路,而不是直接丢结论。
6.1 本地 Ollama 突然“假死”,请求全部超时
表象:Spark job 跑了一半,所有 mapPartitions 任务开始超时,Ollama 服务进程还活着,但响应时间从 500ms 涨到 60s 以上。
排查链路:先看了 Ollama 服务端日志,发现有大量context length exceeded的错误,再看了 GPU 显存利用率,发现显存几乎被打满。进一步定位到是某个批次里的 prompt 特别长,超过了模型上下文窗口,Ollama 在尝试做 prefix caching 时把显存刷爆了。
根因:批量推理时,个别长文本的 token 数接近甚至超过模型上限,导致显存峰值飙升,服务端 OOM 后不再响应新请求。
解决方案:在 Spark 端对输入文本长度做预处理,超过阈值就截断或者丢弃,同时给推理服务加一个显存保护,比如 vLLM 的--gpu-memory-utilization限制到 0.9,避免显存吃满。还有一劳永逸的做法是把输入文本按长度分桶,长文本单独一批、短文本单独一批,避免长短文本混合导致推理批次的显存波动剧烈。
6.2 幂等性:BigJSON 结果被截断,下游解析全挂
表象:部分行产出的 JSON 是残缺的,下游from_json解析失败,任务报错。
排查链路:先看报错堆栈,发现解析失败的行大多是长文本生成的长结果。再检查模型输出,发现结果是半截——模型生成了 512 token 就被截断了。
根因:我在调用推理服务时max_tokens设置得不够,模型输出超过了这个限制被硬切。LLM 的文本补全没有必然的终止符,你设了 512,它就只生成 512,后面的半个 JSON 自然没法解析。
解决方案:一是把max_tokens调大(比如 2048),满足绝大多数场景;二是在 Spark 端增加对输出 JSON 的合法性校验,解析失败的重试一次,重试时把max_tokens翻倍;三是如果你的下游只需要 JSON 的某个字段,让模型输出纯 JSON 字段值而不是整个 JSON 对象,可以显著降低被截断的概率。
6.3 API 限流导致的“整个任务被拖死”
表象:任务开始时速度正常,跑了一小时之后越来越慢,最终大部分 task 都在等待重试。
排查链路:看 Spark UI 的 task 状态,大量 task 卡在RUNNING状态很久。再看推理服务日志,发现一堆 429 限流响应。进一步查了限流策略,发现服务端限制的是“每分钟 1000 次请求”,我们没有加客户端限流,所以瞬时请求量远超配额。
根因:Spark 天然有“并行放大”效应——一个 stage 里几百个 task 同时启动,每个 task 内的线程池再放大 8 倍,瞬时请求量直接打穿配额。
解决方案:在客户端实现全局信号量,控制整个 Spark job 的总并发数不超过服务端配额。我用的是 Python 的multiprocessing.Manager或者 Redis 计数器,在 driver 上初始化一个限流器,广播给各 executor:
import threading import time from pyspark import SparkContext # 在 driver 端初始化 _rate_limiter = None def init_rate_limiter(rate_per_minute): global _rate_limiter _rate_limiter = RateLimiter(rate_per_minute) def process_partition_with_ratelimit(rows): for row in rows: _rate_limiter.wait_if_needed() ...当然,更直接的做法是把请求量控制在服务端容量的 70% 以内,留出余量。如果服务端是自己部署的 vLLM,也可以在客户端做自适应流控:实时监测最近 10 秒的平均延迟,超过阈值就自动降低并发,低于阈值再恢复。这个思路能从根上解决限流问题。
6.4 异构数据源接入:达梦数据库这类场景的适配技巧
最后说一个比较冷门但实用的点。很多企业数据不光存在 Hive 和 Spark SQL 里,还存在达梦(DM)、Oracle 这类国产或传统数据库里。你要让 Spark 调用大模型,首先得把这些数据拉进 DataFrame。
用 JDBC 是最常规的方式:
df = spark.read \ .format("jdbc") \ .option("url", "jdbc:dm://x.x.x.x:5236/SCHEMA") \ .option("dbtable", "table_name") \ .option("user", "xxx") \ .option("password", "xxx") \ .option("driver", "dm.jdbc.driver.DmDriver") \ .load()这里有个坑:如果表很大,默认情况下 JDBC 是单分区读的,速度极慢。需要指定分区键和分区数:
df = spark.read \ .format("jdbc") \ .option("numPartitions", 20) \ .option("partitionColumn", "id") \ .option("lowerBound", 0) \ .option("upperBound", 10000000) \ .load()这样 Spark 会把查询拆成 20 个子查询并行去读,吞吐量提升非常明显。读完数据后统一走 DataFrame 管线,后续的数据清洗、过滤、推理调用就和别的数据源完全一致了。
写在最后的一点个人体会
干完这一圈“Spark 调大模型”的活,我最大的感受是:这个场景真正的难点不在“深度学习”,而在“分布式系统与模型服务的协作”。Spark 擅长高速并行计算,但它的假设是任务之间相互独立;大模型推理服务恰好是一个有状态、有并发上限、有失败概率的外部依赖。把这两者拼接起来,本质上是个流控与容错的工程问题。
我的建议是:先用小数据量验证 prompt 和输出格式,确认结果质量达标;再用中等数据量(几万条)压测推理服务的吞吐上限和延迟分布;最后才上全量任务。期间务必做好结果缓存和失败重试,否则每次调参都等于多烧一遍钱。管线稳定之后,再考虑把推理服务接入带优先级的队列、或者换成吞吐更强的引擎,这些才是后续真正的优化方向。