最近有好几个团队过来问我同一个问题:手里的 Spark 集群跑着一堆离线数据,现在想把大模型接进来做文本分类、实体抽取、情感打分这类活儿,到底该怎么做?有的同学上来就打算在 UDF 里把模型 load 进去,结果集群直接 OOM;有的走 HTTP 逐条调大模型接口,跑了一晚上一看进度条才走了几百万条。今天就把我在多个项目里摸爬滚打总结出来的方案和细节整理出来,覆盖方案选型、代码写法、资源规划和最容易被坑坏的那些细节。
先说明白这篇文章适合谁看:你已经会写 PySpark,跑过简单的 join 或 groupBy,现在需要把大模型推理任务和大数据 pipeline 结合在一起;或者是刚接手一个所谓“LLM + Spark”项目,想知道别人是怎么做的。全文以实操为主,尽量不堆概念。
1. 为什么偏偏要在 Spark 里调用大模型
1.1 把“存量数据拥抱大模型”这件事搬上生产
先说场景。现在的数据团队手里普遍压着大量历史数据:几亿条用户评论、合同扫描件转出来的文本、客服对话记录、商品描述,甚至是一堆日志字段。过去这些数据要么洗完之后做统计报表,要么用 TF-IDF 和词向量凑合出一个“语义相似度”,但真正想要的理解型任务,比如判断一段评论到底是“质量抱怨”还是“物流抱怨”,传统方法做不动。
大模型出来之后,大家第一个想到的是让 LLM 干这个活。但几亿条数据不可能靠脚本一条条调接口,也不可能全部塞进 Excel。这时候 Spark 作为离线批处理的事实标准,自然而然就成了很多人眼里“调度这些数据”的首选。你只需要在 Spark 里写个 UDF,把 DataFrame 的一部分字段传进去,调用大模型,把结果作为新列写出来,逻辑上非常顺。
但逻辑顺不等于跑得顺。我见过不少项目就在这一步开始炸:模型参数几百 G,每个 Executor 都 load 一遍,内存直接翻车;或者从 Hive 读一张表,每条 row 都发起一次 HTTP 请求,网络 IO 成了瓶颈,跑了好几天都跑不完。所以先搞清楚 Spark 里调大模型到底有哪些靠谱套路,比直接写代码更重要。
1.2 先理解 UDF 在集群里的真实运行位置
很多人对 Spark UDF 有一个误解,以为 UDF 是在 Driver 上统一执行的,其实完全不是。Spark 跑批任务时,数据被切成多个 partition,分布在不同的 Executor 上;每一个 partition 会被一个 Task 处理。你写的 UDF 会被序列化之后分发到每个 Executor,然后由 Executor 里的 Python worker 进程逐行执行。
这带来一个非常关键的影响:如果 UDF 里要访问某个对象,比如一个大模型、一个 HTTP 客户端、一个数据库连接池,那么这个对象要么能跟着 UDF 一起被序列化分发,要么就必须在 Executor 的 worker 进程里独立初始化。
理解了这一点,后面很多坑就都能解释了:为什么模型不能直接在 UDF 里 load?因为每个 Executor 都要 load 一份,你有 100 个 Executor,就得占 100 份显存或内存;为什么很多人的 HTTP 请求特别慢?因为每个 Task 可能在不同时间启动,连接没法复用,甚至每次请求都要重新建立 TCP 连接。没有这种全局视角,后面调优基本靠瞎试。
2. 四种主流接入方案:选型对比与推荐场景
2.1 方案一:在 UDF 里直接加载本地或内网模型
这是新手最容易想出来的方案:既然要调大模型,那直接把模型放到集群的每台机器上,在 UDF 里用transformers库加载,然后逐条推理。听起来很美好,实际只适合小模型。
具体的做法是一个 map 类型的 UDF:
from pyspark.sql.functions import udf from pyspark.sql.types import StringType @udf(StringType()) def local_model_predict(text): from transformers import pipeline # 注意:如果在这里加载模型,每条 task 第一次执行时都会触发加载 classifier = pipeline("sentiment-analysis", model="/path/to/model") result = classifier(text, truncation=True)[0] return result["label"]这种写法最大的问题在于模型加载时机。如果 pipeline 在 UDF 内部初始化,意味着每个 Python worker 的每个 task 第一次调用时都可能触发一次模型加载。即使模型只有几百 MB,100 个 Executor 同时加载也会把磁盘 IO 和内存打满。更别提真正的 LLM,动辄 7B、13B 参数,光权重就要十几 GB 甚至几十 GB。
所以这个方案我只建议用在两类场景:一是模型很小,比如几百万参数的 MiniLM、BERT-base 这类;二是你只需要在单个 Executor 上做一次性处理,而不是全集群高并发推理。如果非要在大模型场景里用,也可以把模型放在共享文件系统(比如 HDFS 或对象存储),让 Executor 从远端加载,但启动时间会非常感人,而且多个 worker 同时拉模型文件可能把带宽占满。
2.2 方案二:逐条或小批量调用在线 LLM API
这是目前生产环境里最常见的做法。大模型统一以一个在线服务的形式部署在内网,Spark 这边不去管模型本身,只管发 HTTP 请求、拿结果。这个在线服务可以是外部厂商的 API,也可以是团队自己用 vLLM、TGI、SGLang 之类框架部署的模型服务。
最简单的实现长这样:
import requests from pyspark.sql.functions import udf from pyspark.sql.types import StringType def call_llm(text): resp = requests.post( "http://llm-service:8000/v1/completions", json={"prompt": text, "max_tokens": 128}, timeout=30, ) resp.raise_for_status() return resp.json()["choices"][0]["text"] llm_udf = udf(call_llm, StringType()) df.withColumn("llm_result", llm_udf(df["content"]))逻辑很简单,但性能一言难尽。每条数据一次请求,如果数据量是一亿条,哪怕每次请求只要 200 毫秒,单并发也要跑两百多天。再加上网络开销、限流、超时,这个方案不经过任何优化就跑全量数据,基本等于自杀。
所以这个方案必须搭配批量优化。我们要做的核心是两件事:一是尽量多利用模型服务的并发能力,二是减少 HTTP 请求的次数或提高每次请求的处理量。下面第三章我会详细讲怎么用 pandas UDF 和 Iterator 风格实现真正的批量调用。
2.3 方案三:Spark 调度 + 外部离线推理服务
如果说方案二是 Spark 直接跟在线服务同步交互,那方案三是把一个完整的异步离线推理链路搭起来。Spark 的角色变成“数据组织和结果回收方”,真正的大模型推理发生在一个独立的服务集群里。
大致流程是这样的:Spark 批量读数据,筛选出需要推理的字段,把数据写入消息队列(Kafka、Pulsar)或中间表;一个常驻的推理worker从队列里拿数据,调用部署好的大模型服务做推理,再把结果写回结果表;Spark 那边可以用流式任务或者定时批任务把结果读回来,跟原始表做 join,得到最终带推理结果的数据。
这个方案的优点是 Spark 不再直接依赖大模型服务的实时性能和稳定性,推过去的数据可以慢慢消化,模型服务即使重启、扩容,也不影响 Spark 主任务。缺点是架构复杂度明显上升,你需要额外维护队列和推理 worker。如果你们团队已经有现成的推理平台,或者大模型服务经常抖动,我建议优先考虑这个方案;如果只是临时跑一次数据,就太重了。
2.4 方案四:引入 Ray 或其他分布式推理引擎做协同
比方案三更激进一点的做法是直接用 Ray 这类通用分布式计算框架管理大模型推理,Spark 和 Ray 各自负责自己擅长的部分。比如 Spark 负责从 Hive 抽数、做复杂的 SQL 清洗和 join,然后你可以把要推理的数据转换成 Ray 的 object ref,Ray 侧用多个 GPU actor 并行加载模型、做推理,最后结果再转回 Spark DataFrame 做下游统计。
这个方案适合的场景有两个特征:一是模型推理量非常大,二是你们团队的 GPU 资源本身已经用 Ray 或者类似框架在管理。如果只是为了一个批处理任务,专门搭一套 Ray 集群,代价会比较高。而且 Spark 和 Ray 之间的数据打通没那么自然,通常需要走 Parquet 文件或者 Redis 之类的中间存储,增加了延迟。
2.5 各方案选型对比
下表是我个人在项目里做选型时的参考框架:
| 方案 | 优点 | 缺点 | 适合场景 |
|---|---|---|---|
| UDF 加载本地小模型 | 部署简单,无外部依赖 | 大模型内存爆炸,加载慢 | 百 MB 级小模型,低并发 |
| 逐条/小批量调用在线 API | 逻辑简单,模型可集中管理 | 性能差,易被限流 | 数据量小、临时验证 |
| Spark + 外部离线推理服务 | 解耦好,模型服务可独立扩展 | 组件多,链路长 | 大规模常态化推理任务 |
| Spark + Ray 分布式推理 | 推理性能上限高 | 架构重,运维复杂 | GPU 密集型、高吞吐需求 |
从我自己的实践来看,大多数团队的选型会落在方案二和方案三之间。临时跑一版用方案二,常态化跑批就要升级成方案三。方案四适合已经有该基础设施的团队,不建议现搭。
3. 实战演练:给上亿条评论做大模型情感分类
3.1 设定一个具体任务
我拿一个真实案例来演示。假设我们有一张 Hive 表app_comments,里面存了电商平台的用户评论,核心字段是comment_id和comment_text。现在要做两件事:给每条评论打一个情感标签(positive / negative / neutral),再提取一个话题标签,比如“物流”、“质量”、“价格”、“服务”。全量数据大约 1 亿条,我们部署了一个内网模型服务,兼容 OpenAI 的 Chat 接口风格,QPS 上限约 200。
这个任务如果按方案二直接逐条调,显然不行。就算按 200 QPS 跑满,1 亿条也要 5 万秒,差不多 14 个小时。但实际网络、限流、重试损耗一叠加,跑 24 小时都算运气好。所以这里我用 pandas UDF 加 Iterator 模式来写,尽可能把批量优势发挥出来。
3.2 先封装一个独立的 Python 调用函数
不管 Spark 那边怎么包,最关键的是先把“调用大模型”这件事做成一个纯净的 Python 函数。这样方便本地调试,也方便后面放到 Spark 的不同封装模式里。
这里用一个 OpenAI 兼容接口的写法,假设请求参数是model、messages、temperature那些。如果你们有自己的协议,替换成requests.post的 body 就行。
import requests import time from typing import List, Dict API_URL = "http://llm-service:8000/v1/chat/completions" API_KEY = "internal-token" def llm_chat(prompt: str, max_tokens: int = 256) -> str: payload = { "model": "qwen2.5-7b-instruct", "messages": [ {"role": "system", "content": "你是一个文本分析助手,只输出 JSON。"}, {"role": "user", "content": prompt}, ], "temperature": 0.0, "max_tokens": max_tokens, } resp = requests.post( API_URL, json=payload, headers={"Authorization": f"Bearer {API_KEY}"}, timeout=60, ) resp.raise_for_status() return resp.json()["choices"][0]["message"]["content"]这个函数我没加任何重试逻辑,因为重试要放到批量层去做,不然单个函数里写死退避逻辑会拖慢整体吞吐。这里高亮一个关键点:超时必须设,不要用默认的无限等待。大模型服务有时会因为排队过长导致单个请求十几秒不回,如果不设超时,整个 Spark Task 就会一直挂着,最后变成一副“集群活着但任务不动”的诡异画面。
然后再加一个解析函数:大模型输出并不总是稳定的 JSON,有时会带解释文字,有时直接给你一段 Markdown。所以我一般会写一个容错解析,提取 JSON 片段再解析,解析失败就返回一个默认结果。
import json import re def parse_llm_json(text: str) -> Dict: text = text.strip() # 去掉可能的 ```json 包装 text = re.sub(r"^```(?:json)?|```$", "", text, flags=re.MULTILINE).strip() try: return json.loads(text) except json.JSONDecodeError: # 尝试找到第一个 { 到最后一个 } start = text.find("{") end = text.rfind("}") if start != -1 and end != -1 and end > start: try: return json.loads(text[start:end+1]) except json.JSONDecodeError: pass return {"sentiment": "unknown", "topic": "unknown"}为什么这么写?因为实际跑过就知道,大模型偶尔会在 JSON 前后加解释,或者中途截断。你要是直接把json.loads的结果当成最终结果,出错率不高,但一亿条数据哪怕 1% 的脏数据也是 100 万条要重跑,很痛。
3.3 用 pandas UDF 包装批量推理
接下来是 Spark 部分。我不建议用普通 UDF,因为逐行调用llm_chat的话,每行都是一次完整的 HTTP 往返,延迟被完全暴露。更合理的是用 pandas UDF 的 Iterator 模式,每次处理一批数据,这样可以在批内复用逻辑、控制并发,甚至可以把一批文本拼接成一次请求。
基础的 Iterator 写法是这样的:
from pyspark.sql.functions import pandas_udf import pandas as pd @pandas_udf("string") def sentiment_topic_batch(iterator): for batch in iterator: # batch 是一个 pandas Series,代表一批 comment_text texts = batch.tolist() results = [] for text in texts: prompt = ( "请对以下用户评论进行情感和话题分析," "返回 JSON,格式为 {\"sentiment\": \"positive|negative|neutral\", " "\"topic\": \"物流|质量|价格|服务|其他\"}。\n" f"评论:{text}" ) raw = llm_chat(prompt) parsed = parse_llm_json(raw) # 把 dict 转成字符串,后面统一解析 results.append(json.dumps(parsed, ensure_ascii=False)) yield pd.Series(results)然后在主代码里这样调用:
df = spark.sql("SELECT comment_id, comment_text FROM app_comments WHERE comment_text IS NOT NULL") result_df = df.withColumn("llm_result", sentiment_topic_batch(df["comment_text"]))跑起来之后你会发现,它比逐条 UDF 快,但还没达到理想状态。因为这里的llm_chat仍然是每条文本一个请求,只是把 Spark Task 的序列化开销和行处理开销省掉了一部分。真正的吞吐瓶颈还是在大模型服务的 QPS 上。
所以如果要进一步提速,有两种常见做法:一种是把文本改成批量提示词,让模型一次处理多个文本,减少请求数;另一种是在 Executor 里做并发请求,一次批数据拆成几个线程同时调接口,充分利用模型服务的 QPS。我下面把两种都说一下。
3.4 批量提示词与并发请求的取舍
批量提示词的思路是:把这一批 50 条评论拼成一个 JSON 数组,塞进一个 prompt,让模型一次性返回 50 条结果。
def llm_chat_batch(texts: List[str]) -> List[str]: body = json.dumps([{"id": i, "text": t} for i, t in enumerate(texts)], ensure_ascii=False) prompt = ( "请逐条分析以下评论数组中的每条评论,返回一个 JSON 数组," "每个元素是 {\\\"sentiment\\\": ..., \\\"topic\\\": ...},下标保持对应。\n" f"数组:{body}" ) raw = llm_chat(prompt, max_tokens=1024) try: parsed = json.loads(raw) # 按长度对齐,不完整时用 unknown 补 return [json.dumps(x, ensure_ascii=False) for x in parsed[:len(texts)]] except json.JSONDecodeError: return [json.dumps({"sentiment": "unknown", "topic": "unknown"}) for _ in texts]这种做法能大幅降低请求次数,比如原来 50 条要 50 个请求,现在只需要 1 个。但风险在于:模型可能不严格对齐下标,或者输出被截断,导致整个 batch 解析失败。一旦失败,这一批全部要重来,反而更慢。所以我把 batch 大小控制在 20~50 之间,并且解析容错做得比较强。
另一种做法是并发请求。在 pandas UDF 里,对一批 100 个文本,用 ThreadPoolExecutor 同时发 5~10 个请求,等待全部返回。这样可以比较充分地利用模型服务的 QPS,而且因为是在 Executor 端并发,不会受 Spark 单 Task 串行限制。
from concurrent.futures import ThreadPoolExecutor, as_completed def call_llm_with_limited_concurrency(texts: List[str], max_workers: int = 8) -> List[str]: results = [None] * len(texts) with ThreadPoolExecutor(max_workers=max_workers) as executor: future_map = { executor.submit(llm_chat, build_prompt(texts[i])): i for i in range(len(texts)) } for future in as_completed(future_map): idx = future_map[future] try: raw = future.result() results[idx] = json.dumps(parse_llm_json(raw), ensure_ascii=False) except Exception: results[idx] = json.dumps({"sentiment": "unknown", "topic": "unknown"}) return results这里有个非常重要的事:并发数不要拍脑袋定。如果模型服务 QPS 上限是 200,平均每个请求耗时 300ms,那么单个 worker 的并发上限差不多是 60(因为 1 秒 / 0.3 秒 × 20 QPS 其实这个数字要根据模型吞吐算)。在实际项目里更稳妥的办法是拿到模型服务的压测数据再定。第一版可以先保守一点,并发设 4~8,跑五分钟看看服务端监控的 QPS 和延迟,再逐步上调。
完整的 pandas UDF 代码就变成了:
@pandas_udf("string") def sentiment_topic_batch(iterator): for batch in iterator: texts = batch.tolist() # 20 条一组拼接,或直接并发逐条 yield pd.Series(call_llm_with_limited_concurrency(texts, max_workers=8))3.5 全量任务怎么跑:从抽样到扩容
我不建议拿 1 亿条直接开跑。第一次调试时,先读 1 万条数据,确认结果格式正确、服务端没有报错,再逐步放大。比如先跑 100 万条,看单 Task 的平均耗时,推算全量所需时间,再决定要不要加资源。
假设 100 万条数据,跑了 3 分钟,那 1 亿条就差不多要 300 分钟,5 个小时。如果你觉得太慢,可以加大spark.sql.shuffle.partitions没用,关键是增加 Executor 数量和并发度。比如从 50 个 Executor 加到 200 个,单 Task 处理的数据量不变,但并行度上去了,总时间理论上是原来的四分之一。
但也要注意,加太多 Executor 不一定是好事。大模型服务端如果 QPS 扛不住,加再多 Spark Executor 也只会看到大量超时和重试。正确的做法是先确认服务端瓶颈在哪。我常用的思路是把 Spark UI 里的 Task 耗时和服务端监控对照看:如果 Task 平均耗时接近请求平均耗时,说明瓶颈在网络或服务端,加 Spark 资源没用;如果 Task 里有大量时间花在处理和排队上,那确实需要提高并行度。
4. 我踩过的那些坑:序列化、超时、限流与资源配置
4.1 为什么请求函数经常传不进去
很多人写完 UDF 提交任务,直接在 Driver 端报错,类似Could not serialize object: TypeError: cannot pickle 'xxx' object。原因是你把某个无法被 pickle 的对象定义在了 UDF 所在的作用域里,比如在类里初始化了一个 HTTP 连接池,或者把一个大模型的pipeline对象定义成全局变量。
PySpark 在分发 UDF 时会把相关对象通过 pickle 序列化后发给 Executor。像requests.Session虽然勉强能 pickle,但序列化之后连接池状态会丢;一些模型对象、GPU 句柄根本无法序列化。解决办法就是不要在 UDF 对象里直接持有这些资源,而是把它们的初始化逻辑放到 Executor 内部惰性执行。比如在 UDF 里判断全局变量_client是不是 None,是就现场建一个:
_client = None def get_client(): global _client if _client is None: _client = LLMClient() return _client这样 UDF 被序列化时只发送了一个函数引用,真正重的客户端对象在每个 Executor 的 worker 进程里各自实例化一次。这也是为什么很多人说“在 mapPartitions 里初始化一次连接池”才是正确姿势。
4.2 Executor 核数、内存与并行度怎么配
Spark 调大模型推理和普通的 ETL 任务不一样,CPU 和内存都不是绝对瓶颈,网络 IO 和模型服务端 QPS 才是。但这不代表资源随便配。我见过不少同事开 100 个 Executor、每个 Executor 8 核,结果大模型服务被打挂,任务疯狂重试。反过来,如果你只开 5 个 Executor,每个核数也不高,那又白白浪费了模型服务的吞吐能力。
一个比较稳的先手配置是:先限制spark.executor.cores=2,每个 Executor 内存 4~8G,这样单个 Executor 的并发请求数基本在 2~4(取决于你在 UDF 里开的线程数)。跑 5 分钟看一下模型服务监控的 QPS 是否接近上限,如果离上限还很远,再增加 Executor 数量,或者把executor.cores调高。注意executor.cores调高不等于 Spark 里每个 Task 就会自动并发发起 HTTP 请求,除非你在 pandas UDF 内部自己开了线程。
另外有个容易被忽略的点:spark.sql.execution.arrow.enabled最好保持默认开启,因为 pandas UDF 依赖 Arrow。如果你的 Spark 版本比较老,Arrow 性能有问题,可以升级 Spark 或 pyarrow 版本,而不是关掉 Arrow 硬跑。
4.3 超时、限流、熔断,一个都不能少
在线模型服务不像你本地写的函数,它受队列长度、GPU 显存、batch size 等因素影响,响应时间起伏很大。所以请求必须设置超时,最好还能区分连接超时和读超时。
我常用的请求配置是timeout=(10, 120),意思是连接超时 10 秒,读超时 120 秒。大模型长文本生成可能确实需要较长时间,所以读超时不能太短。
对于限流,模型服务一般会返回 429 或者 503。429 是请求太频繁,503 是服务暂时不可用。处理方式不同:429 可以稍微退避后重试,503 要等更长时间,甚至放弃这一批。我会在批量函数里加一个简单的重试逻辑,最多重试 3 次,退避时间按 1 秒、2 秒、4 秒递增:
import time def call_with_retry(prompt, retries=3): for attempt in range(retries): try: return llm_chat(prompt) except requests.exceptions.HTTPError as e: if e.response.status_code in (429, 503) and attempt < retries - 1: time.sleep(2 ** attempt) continue raise except requests.exceptions.ReadTimeout: if attempt < retries - 1: time.sleep(1) continue raise这里有一个容易犯错的地方:不要对每条请求都在内部重试,因为重试会叠加在批量并发之上,可能瞬间造成更大的请求洪峰。正确做法是把重试次数压到很低(1~2 次),并且配合并发数上限来控制总体 QPS。如果错误率持续超过 5%,应该停下来查模型服务状态,而不是让重试逻辑硬扛。
再补一个熔断思路:在做超大任务之前,先往模型服务发 50 个测试请求,统计成功率。如果成功率低于 90%,就不要启动 Spark 全量任务,先解决模型服务的问题。这个前置检查写成一个单独的 Python 脚本,不进 Spark,用来给平台做发布前检查非常有用。
4.4 token 长度、输出截断与结果校验
大模型服务的输入输出都有 token 上限。文本如果太长,接口会直接报错,也可能自动截断。我一般在调用前做一次文本长度截断,比如保留前 1000 个字符,因为情感分类和话题识别对文本开头的信息最敏感。这样做有两个好处:一是减少 token 消耗,变相降低成本;二是避免因为超长文本导致请求失败。
输出截断是另一个问题。max_tokens设太小,模型输出到一半被切断,JSON 不完整,解析就会失败。所以解析函数要非常健壮,解析不了就返回 unknown。千万不要让一个坏 JSON 导致整个 Spark Task 失败,那样重跑的成本高得多。
我还会做一道结果完整性校验:跑完 Spark 任务后,统计llm_result里包含"unknown"的比例。如果 unknown 比例超过 1%,说明模型输出质量有问题,需要检查是提示词还是服务端配置出了问题。这个比例可以作为质量监控指标写进平台报表。
4.5 失败重试的粒度:批内重试优于批级重试
很多人写 pandas UDF 时,把整个 batch 包在 try except 里,一旦 batch 里有一条请求超时,整个 batch 都标记失败,然后 Spark 重跑这个分区。这个做法非常浪费。因为一个 batch 里其他 99 条可能都已经成功返回了,只是某一条超时,你重跑整个 batch,等于把 99 条成功的请求又打了一遍,白白浪费资源,还可能压垮模型服务。
我建议在 batch 内部做单条重试。也就是上面call_llm_with_limited_concurrency里给每条请求单独 try except,异常时单独重试。即使这条最终失败,也只影响这一条,输出 unknown 或默认值,不影响任务整体跑完。后续如果对 unknown 结果有要求,可以再单独捞出来重跑一次。
这里也解释一下为什么建议把结果解析函数写得尽量宽松:你的目标是让全量任务稳定跑完,而不是追求单条结果完美。稳定跑完意味着任何异常都要被兜住,绝不能因为某条脏数据让整个 Executor 崩溃。
4.6 常见问题速查表
| 现象 | 可能原因 | 排查与处理 |
|---|---|---|
| 任务一启动就报序列化错误 | UDF 持有无法 pickle 的对象 | 把客户端初始化改为惰性全局变量 |
| 内存持续上涨直到 Executor 挂掉 | 每条 Task 内 load 了大模型或连接池过多 | 改为 mapPartitions 初始化,控制连接池大小 |
| 数据量很大但跑得极慢 | 逐条调用 HTTP,网络往返成为瓶颈 | 改为批量提示词或并发请求 |
| 任务卡住不动,Spark UI 显示 Task 仍在运行 | 请求没设超时,模型服务长时间不返回 | 设置 read timeout,并加重试 |
| 模型服务返回大量 429/503 | 并发过高,超过服务 QPS | 减少 Executor 或线程并发数,加退避重试 |
| 输出结果大量是 unknown | 提示词不合适或 max_tokens 太小 | 调整提示词,增大 max_tokens,增加校验 |
4.7 成本与性能估算:先算账再开跑
最后再分享一个经验:大模型调用是有成本的,不管是外部 API 按 token 计费,还是内网 GPU 资源占用,跑之前最好算一笔账。估算方式很简单:
假设单条评论平均 200 个 token,输入输出加起来大约 400 token。如果批量提示词方案,会稍微多一点,但可以忽略。1 亿条就是 400 亿 token。外部 API 按百万 token 计费,假设每百万 token 0.5 块,总成本就是 2000 块。如果觉得贵,就要考虑是否抽样子集、降低任务频率,或者用小模型替代大模型。
时间估算也有公式:总请求数 / (并发数 × 单请求 QPS 贡献)。假设并发 8,单请求 1 秒,8 并发每秒 8 个请求,1 亿条需要 12500 秒,约 3.5 小时。这个估算结果可以帮你说服业务方接受调度时间,也能帮你自己决定要不要加资源。
5. 架构层面的一些个人体会
踩过几次大坑之后,我现在的习惯是把“在线”和“离线”分开。真正跑全量批处理的任务,尽量不要让 Spark 直接对高频在线推理服务发请求,因为两者生命周期不一样:Spark 任务跑完就结束了,在线服务却要常年稳定。如果 Spark 侧并发控制不好,一次全量任务就可能把在线服务的资源打满,影响线上业务。更稳的做法是把推理请求打到独立的离线模型服务池上,或者在中间加一个削峰队列。
另外,prompt 的管理不要散落在代码里。我见过很多团队把 prompt 写死在 UDF 里,后面想改一点措辞就要改代码重跑全量。更合理的做法是把 prompt 模板存到配置中心或 HDFS 上的一个配置表里,Spark 任务启动时读取,这样调 prompt 不用重新发布代码,只需要重启任务。对于任务频率高、团队人员流动大的项目,这个优化非常值。
最后一点是监控。Spark 任务跑多久、失败多少条、unknown 比例多少、平均耗时多少,这些指标最好都能汇总到团队已有的监控平台上。大模型推理跟普通 ETL 不一样,它的结果质量波动很大,今天 prompt 还能用,明天换了个模型版本可能结果就变了。没有监控和对比,后面排查问题会非常痛苦。
这篇文章里提到的代码和方案,都是我实际跑过之后沉淀下来的。如果你现在正要动手,建议不要贪多,第一版先用 pandas UDF 加并发请求,把端到端链路跑通,再逐步往异步队列或 Ray 方向演进。大模型在 Spark 里调用这件事,最大的难点从来不是写代码,而是把资源、并发、稳定性、成本这些细节都考虑周全。希望这篇分享能帮你少走几步弯路。