外观
批处理推理管道:Airflow 调度下的离线批量推理
一句话定义:批处理推理(batch inference)是对一批积累的数据离线执行模型推理、结果落库,由调度器(如 Airflow)按时触发,跑完即出结果——与在线推理(请求来一个算一个)在架构与成本上是两条路线。
为什么值得动手:同样一次模型推理,批处理成本可以比在线低一个数量级——在线服务为扛突发流量常备 3 倍冗余 GPU,批处理用少量机器分片慢慢跑、利用率还能打满。日活画像更新、日报评分、风控回捞、推荐候选离线预生成,都是批推理的典型场景。本文用 Airflow + Python(含 Spark 选项)搭一套可重试、可断点续跑的管道,重点是幂等、分片、checkpoint 三个工程动作。
一、批推理 vs 在线推理:先算经济账
| 维度 | 在线推理 | 批处理推理 |
|---|---|---|
| 触发方式 | 请求驱动,QPS 波动 | 定时调度,任务量已知 |
| 资源 | 常驻,按峰值 3 倍冗余 | 按批起停,跑完释放 |
| 成本 | 高(GPU 常年开着) | 低一个数量级(见下表) |
| 延迟要求 | 毫秒~秒级 | 分钟~小时级 |
| 数据 | 单条请求 | 全量/增量数据集 |
| 出错影响 | 单个请求失败 | 全批重跑,需幂等与断点 |
一组成本数字(同一张 A10 GPU,同一 7B 模型):
| 部署方式 | 利用率 | 单条推理成本(含摊销) |
|---|---|---|
| 在线常驻(扛 3× 峰值冗余) | 平均 20%~30% | 基准 1.0 |
| 在线 + 自动扩缩 | 40%~60% | 约 0.4~0.6 |
| 批处理(分片跑满 GPU) | 85%~95% | 约 0.1~0.2 |
结论:能离线算的不要在线算。批推理的适用判断——结果可接受分钟级延迟、数据按天积累、计算量可预估。这也是部署架构模式里"批处理模式"与"在线模式"的分界。
二、架构设计
text
┌────────────┐ ┌──────────────────────────┐ ┌──────────────┐
│ Airflow │ │ 推理集群(K8s/裸机) │ │ 数据层 │
│ (调度器) │───▶│ ┌─────────┐ ┌────────┐ │ │ │
│ 每天 02:00 │ │ │ Spark/ │ │ 推理 │ │───▶│ 结果库/对象存储 │
└────────────┘ │ │ 分片任务 │─▶│ worker │ │ └──────────────┘
│ └─────────┘ └────────┘ │
└──────────────────────────┘三个职责分离:调度器只负责"到点触发 + 状态跟踪",分片任务负责"把数据切成可并行的小块",推理 worker 只干"加载模型 + 批量前向"。这样任意一块坏掉都能单独重试。上游衔接见MLOps 部署流水线(训练产物如何进仓库、如何触发本管道)。
三、实战一:Airflow DAG(每日画像打分)
python
# dags/user_profile_scoring.py
from datetime import datetime, timedelta
from airflow import DAG
from airflow.operators.bash import BashOperator
from airflow.operators.python import PythonOperator
from airflow.operators.empty import EmptyOperator
default_args = {
"owner": "ml-platform",
"depends_on_past": False,
"retries": 2, # 失败自动重试
"retry_delay": timedelta(minutes=10),
"execution_timeout": timedelta(hours=4),
}
with DAG(
dag_id="user_profile_scoring",
schedule="0 2 * * *", # 每天 02:00(避免与在线高峰抢资源)
start_date=datetime(2025, 1, 1),
catchup=False, # 不补跑历史(除非想要)
default_args=default_args,
max_active_runs=1, # 同 DAG 不并发跑两轮,防写冲突
) as dag:
start = EmptyOperator(task_id="start")
# 分片:按日期把用户表切成 16 片,输出分片清单
split = BashOperator(
task_id="split_shards",
bash_command="python /opt/pipeline/split_shards.py "
"--dt {{ ds }} --n-shards 16 --out /data/shards/{{ ds }}",
)
# 并行推理:每个分片一个 task,Spark 跑分片任务(见第四节)
# 用动态 task mapping(Airflow 2.4+)按分片清单展开
from airflow.operators.python import PythonOperator
def score_shard(shard: str):
# 实际调推理 worker,见第四节
return run_inference_shard(shard)
score_all = PythonOperator.partial(
task_id="score_shard",
python_callable=score_shard,
).expand(op_kwargs=[{"shard": f"/data/shards/{__import__('datetime').date.today()}/{i:02d}"}
for i in range(16)])
checkpoint = PythonOperator(
task_id="mark_checkpoint",
python_callable=write_checkpoint, # 记下已成功的数据范围
)
publish = PythonOperator(
task_id="publish_results",
python_callable=publish_to_feature_store, # 结果写回,供在线读取
)
start >> split >> score_all >> checkpoint >> publish要点:max_active_runs=1 防止调度重叠;catchup=False 避免误触发历史补跑;重试与超时在 default_args 里统一声明。调度体系与告警的关系见监控与可观测性。
四、实战二:分片与并行
分片的核心是把"一个巨型任务"拆成"多个可独立重跑的小任务"。按用户 ID 哈希分片最常用:
python
# split_shards.py 片段
import hashlib
def shard_of(user_id: str, n_shards: int) -> int:
"""按用户 ID 哈希分片:同一用户永远落同一片,天然支持增量"""
return int(hashlib.md5(user_id.encode()).hexdigest(), 16) % n_shards
# 结果:users_20250101/ 下生成 16 个分片文件
# shard_00.parquet ~ shard_15.parquet选 Spark 还是 Pandas,先给结论:数据量 < 千万行、单机内存放得下 → Pandas/Polars 更简单;数据量更大或需要分布式 shuffle → Spark。Pandas 的分片推理 worker:
python
# worker.py 片段
import pandas as pd
import torch
from model import CTRModel
def run_inference_shard(shard_path: str, batch_size: int = 512):
df = pd.read_parquet(shard_path)
model = CTRModel.load("/models/ctr_v2/")
model.eval()
results = []
# 分批前向:控制单次显存,同时保证 batch 足够大
for i in range(0, len(df), batch_size):
batch = df.iloc[i:i + batch_size]
features = model.build_tensor(batch) # 与训练一致的预处理
with torch.inference_mode():
logits = model(features)
results.append(model.postprocess(batch, logits)) # 后处理+回填 ID
out = pd.concat(results)
out.to_parquet(shard_path.replace("shards", "results"))Spark 版本(适合大规模、直接用 Spark ML 或外部模型):
python
from pyspark.sql import SparkSession
from pyspark.sql.functions import pandas_udf, col
spark = SparkSession.builder.appName("batch-infer").getOrCreate()
@pandas_udf("double")
def predict_udf(embeddings: pd.Series) -> pd.Series:
"""Pandas UDF:Spark 自动分批传入,内部可调 batch 推理"""
import torch
model = _get_model() # 惰性加载,UDF 进程内只加载一次
tensors = torch.from_numpy(embeddings.to_numpy())
with torch.inference_mode():
return pd.Series(model(tensors).numpy().tolist())
df = spark.read.parquet("s3://data/users/20250101/")
df.select("user_id", predict_udf(col("user_embedding")).alias("score")) \
.write.mode("overwrite").parquet("s3://results/users/20250101/")五、模型加载与 batch 大小
批推理的 GPU 利用率 = f(batch size, 分片并行度)。规则:
- batch size 往大调,直到延迟拐点或显存上限。7B 模型单卡通常 32~128 条/批;embedding 小模型可以 1024+;
- 并行度:一个 worker 跑一个分片、吃满一张卡。多分片并行 = 多卡并行;
- 估算公式:
总耗时 ≈ 数据量 ÷ (batch size × 吞吐 × 卡数)。一张 A10 跑 7B 量化模型,吞吐约 1000~3000 token/s,一千万条短文本打分约 3~8 小时。
batch 与吞吐的理论关系见性能优化与容量规划。如果分片太多、单分片太小,调度和落盘开销会超过推理本身——分片数 ≈ 卡数 × 3~5 是经验甜点。
六、checkpoint 与失败重试:幂等 + 断点续跑
批推理最容易翻车的是跑了一半挂了,全量重跑。三件套解决:
- 幂等写:结果写
mode="overwrite"到以日期为 key 的路径(s3://results/users/20250101/),重跑不产生脏数据; - 分片级 checkpoint:每个分片成功即标记(如写
.done文件或记录到任务表),下次调度只跑没成功的分片; - 重试上限:Airflow
retries自动重试;超时与数据问题要区别对待——超时重试即可,坏数据要告警人工介入。
python
def score_shard(shard: str):
done_flag = shard + ".done"
if Path(done_flag).exists():
return # 已成功,跳过(断点续跑的关键)
run_inference_shard(shard)
Path(done_flag).touch() # 成功后打标全量 vs 增量
能增量就别全量。每日画像通常只重算"有变化的用户"(当天活跃、资料变更),增量命中率取决于业务;全量重算成本高但逻辑简单。折中方案:增量为主 + 每周一次全量对账,校验增量结果与全量一致。数据一致性问题详见常见陷阱与反模式。
七、结果写回与监控告警
写回两种形态:
- 结果库:特征/画像写 Feature Store 或 Redis(供在线读取),注意写在线库会引入在线侧依赖,务必控制并发与幂等;
- 对象存储:评分报表写 Parquet,供 BI/下游取数。
监控三条线(见监控与可观测性):
- 任务成功率:Airflow 自带状态;失败即告警(企业微信/钉钉/Slack webhook);
- 数据新鲜度:告警"结果表最后更新时间 > 26 小时",防"任务成功但数据没写全"的假成功;
- 质量校验:每个分片输出后校验行数、空值率,与上期对比,波动超阈值告警。
Spark 资源配置一例
用 Spark 跑批推理时的资源给法,直接决定跑多快、花多少钱:
bash
# 提交 Spark 任务:8 个 executor、每 executor 2 核 8G
# driver 只做调度,推理负载全部在 executor
spark-submit \
--master k8s://https://k8s.example.com \
--deploy-mode cluster \
--conf spark.executor.instances=8 \
--conf spark.executor.cores=2 \
--conf spark.executor.memory=8g \
--conf spark.dynamicAllocation.enabled=false \
pipeline/batch_infer.py先给结论:批推理的 executor 数量按"数据量 ÷ 单 executor 吞吐"定,别无脑调大。8 个 executor 跑不动的任务,16 个通常能提速约 1.5~1.8 倍(有 shuffle/IO 开销),再往上收益递减;Spark 里真正的瓶颈常是 shuffle 和资源等待,而不是算子本身。
告警接入一例(Airflow 失败回调)
python
# airflow 插件:任务失败即发 webhook(企业微信/钉钉/Slack 通用格式)
def on_failure(context):
dag_id = context["dag"].dag_id
task_id = context["task_instance"].task_id
ts = context["ts"]
requests.post(WEBHOOK_URL, json={
"msgtype": "text",
"text": f"[批推理失败] {dag_id}/{task_id} @ {ts}",
})
# DAG 参数里挂上
default_args = {**default_args, "on_failure_callback": on_failure}常见坑与排查
| 坑 | 现象 | 排查 |
|---|---|---|
| 数据倾斜 | 某分片巨慢,其余早跑完 | 检查哈希分片是否均匀;热点用户拆细粒度分片 |
| 单 worker OOM | 分片太大或 batch 太大 | 减小 batch、增加分片数 |
| 假成功 | 任务绿但数据缺失 | 加数据新鲜度/行数校验,见第七节 |
| 重跑重复写 | 结果翻倍或污染 | mode="overwrite" + 日期 key + 幂等 |
| 调度重叠 | 两轮任务同时写同一批数据 | max_active_runs=1 |
| 与在线抢资源 | 在线 P99 被批推理拖垮 | 批推理错峰调度(夜间)、限 GPU 份额 |
| 断点续跑失效 | 挂了重跑还是全量 | checkpoint 以分片为单位,成功即打标 |
延伸阅读
- 部署架构模式 —— 批处理模式在部署架构谱系中的位置
- 监控与可观测性 —— 任务成功率、数据新鲜度、质量校验的实现细节
- 性能优化与容量规划 —— batch size 与吞吐曲线的理论支撑
- MLOps 部署流水线 —— 批管道上游的训练产物管理与触发链路
- 推理:从前向传播到推理引擎 —— 批前向的底层机理
- NVIDIA Triton 多模型服务 —— 批推理同样可以用 Triton 的 dynamic batching 做加速
参考资料
- Apache Airflow 官方文档:https://airflow.apache.org/docs/
- Apache Spark 官方文档:https://spark.apache.org/docs/latest/
- Pandas 官方文档:https://pandas.pydata.org/docs/
- Airflow Core Concepts:Dags(含 Dynamic Dags / Task Mapping):https://airflow.apache.org/docs/apache-airflow/stable/core-concepts/dags.html