Skip to content

批处理推理管道:Airflow 调度下的离线批量推理

本页速览 离线批量推理(画像更新、日报评分)与在线推理的成本差一个数量级。本文实战用 Airflow + Python/Spark 构建可重试、可断点续跑的批推理管道。

批处理推理管道: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 与失败重试:幂等 + 断点续跑

批推理最容易翻车的是跑了一半挂了,全量重跑。三件套解决:

  1. 幂等写:结果写 mode="overwrite" 到以日期为 key 的路径(s3://results/users/20250101/),重跑不产生脏数据;
  2. 分片级 checkpoint:每个分片成功即标记(如写 .done 文件或记录到任务表),下次调度只跑没成功的分片
  3. 重试上限: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/下游取数。

监控三条线(见监控与可观测性):

  1. 任务成功率:Airflow 自带状态;失败即告警(企业微信/钉钉/Slack webhook);
  2. 数据新鲜度:告警"结果表最后更新时间 > 26 小时",防"任务成功但数据没写全"的假成功;
  3. 质量校验:每个分片输出后校验行数、空值率,与上期对比,波动超阈值告警。

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 以分片为单位,成功即打标

延伸阅读

参考资料