Skip to content

Batch Inference Pipelines

At a glance Offline batch inference (profile updates, daily report scoring) costs an order of magnitude less than online inference. Build a retryable, checkpoint-resumable batch inference pipeline with Airflow + Python/Spark.

Batch Inference Pipelines: Offline Batch Inference on an Airflow Schedule ​

One-line definition: batch inference runs model inference offline over an accumulated batch of data and writes the results to storage; a scheduler such as Airflow triggers it on a schedule, and the results are ready when the run finishes — architecturally and economically a different track from online inference, which scores one request at a time.

Why it's worth doing: the same model inference can cost an order of magnitude less in batch than online — online services keep roughly 3× redundant GPUs standing by to absorb traffic spikes, while batch processing runs at its own pace on a small fleet of machines and can still saturate utilization. Daily active-user profile updates, daily report scoring, risk-control lookback scoring, and offline pre-generation of recommendation candidates are all classic batch-inference workloads. This guide builds a retryable, checkpoint-resumable pipeline with Airflow + Python (with a Spark option), focusing on three engineering moves: idempotency, sharding, and checkpointing.

1. Batch vs. Online Inference: Run the Numbers First ​

DimensionOnline inferenceBatch inference
TriggerRequest-driven, QPS fluctuatesScheduled, workload known in advance
ResourcesAlways-on, ~3× peak redundancySpun up per batch, released when done
CostHigh (GPUs running around the clock)An order of magnitude lower (table below)
LatencyMilliseconds to secondsMinutes to hours
DataOne request at a timeFull or incremental datasets
Failure impactA single request failsThe whole batch reruns; needs idempotency and checkpoints

A set of cost figures (same A10 GPU, same 7B model):

DeploymentUtilizationCost per inference (amortized)
Online, always-on (3× peak redundancy)20–30% averageBaseline 1.0
Online + autoscaling40–60%~0.4–0.6
Batch (shards keep the GPU saturated)85–95%~0.1–0.2

Conclusion: anything you can compute offline, don't compute online. Batch inference fits when results tolerate minute-level latency, data accumulates by the day, and the compute volume is predictable. This is also where deployment architecture patterns draws the line between "batch mode" and "online mode".

2. Architecture ​

text
┌────────────┐    ┌──────────────────────────────┐    ┌────────────────┐
│  Airflow   │    │  Inference cluster (K8s or   │    │   Data layer   │
│(scheduler) │───▶│  bare metal)                 │    │                │
│ daily 02:00│    │  ┌───────────┐  ┌─────────┐  │───▶│ Results store/ │
└────────────┘    │  │ Spark/    │─▶│Inference│  │    │ object storage │
                  │  │ shard jobs│  │ workers │  │    └────────────────┘
                  │  └───────────┘  └─────────┘  │
                  └──────────────────────────────┘

Three responsibilities, cleanly separated: the scheduler only handles "fire on schedule + track state"; shard jobs handle "split the data into independently parallelizable chunks"; inference workers only do "load the model + run batched forward passes." Any one piece can break and be retried on its own. For how the upstream connects — how training artifacts enter the registry and trigger this pipeline — see the MLOps deployment pipeline.

3. Example 1: The Airflow DAG (Daily Profile Scoring) ​

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,                       # automatic retries on failure
    "retry_delay": timedelta(minutes=10),
    "execution_timeout": timedelta(hours=4),
}

with DAG(
    dag_id="user_profile_scoring",
    schedule="0 2 * * *",               # daily at 02:00 (avoids competing with online peak traffic)
    start_date=datetime(2025, 1, 1),
    catchup=False,                       # no historical backfills (unless you want them)
    default_args=default_args,
    max_active_runs=1,                   # never run two rounds of the same DAG at once; prevents write conflicts
) as dag:
    start = EmptyOperator(task_id="start")

    # Sharding: split the user table into 16 shards by date and emit the shard list
    split = BashOperator(
        task_id="split_shards",
        bash_command="python /opt/pipeline/split_shards.py "
                     "--dt {{ ds }} --n-shards 16 --out /data/shards/{{ ds }}",
    )

    # Parallel inference: one task per shard; Spark runs the shard jobs (see Section 4)
    # Expanded from the shard list via dynamic task mapping (Airflow 2.4+)
    from airflow.operators.python import PythonOperator

    def score_shard(shard: str):
        # Actually invokes the inference worker; see Section 4
        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,   # record which data ranges succeeded
    )

    publish = PythonOperator(
        task_id="publish_results",
        python_callable=publish_to_feature_store,  # write results back for online reads
    )

    start >> split >> score_all >> checkpoint >> publish

Key points: max_active_runs=1 prevents overlapping runs; catchup=False avoids accidentally triggering backfills of history; retries and timeouts are declared centrally in default_args. For how scheduling ties into alerting, see monitoring and observability.

4. Example 2: Sharding and Parallelism ​

The core of sharding is turning "one giant job" into "many small jobs you can rerun independently." Hash-sharding by user ID is the most common choice:

python
# snippet from split_shards.py
import hashlib

def shard_of(user_id: str, n_shards: int) -> int:
    """Hash-shard by user ID: the same user always lands in the same shard, which naturally supports incremental runs"""
    return int(hashlib.md5(user_id.encode()).hexdigest(), 16) % n_shards

# Result: 16 shard files generated under users_20250101/
# shard_00.parquet ~ shard_15.parquet

Verdict first on Spark vs. Pandas: under roughly ten million rows, when the data fits in one machine's memory, Pandas/Polars is simpler; for larger volumes or when you need a distributed shuffle, use Spark. A Pandas sharded-inference worker:

python
# snippet from 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 = []
    # Forward in batches: keeps peak GPU memory in check while batches stay large enough
    for i in range(0, len(df), batch_size):
        batch = df.iloc[i:i + batch_size]
        features = model.build_tensor(batch)       # preprocessing identical to training
        with torch.inference_mode():
            logits = model(features)
        results.append(model.postprocess(batch, logits))  # post-processing + reattach IDs
    out = pd.concat(results)
    out.to_parquet(shard_path.replace("shards", "results"))

The Spark version (for large scale, or when you use Spark ML directly or an external model):

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 feeds data in batches automatically; run batched inference inside"""
    import torch
    model = _get_model()   # lazy load: loaded once per UDF process
    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/")

5. Model Loading and Batch Size ​

GPU utilization in batch inference = f(batch size, shard parallelism). The rules:

  • Push batch size up until you hit the latency knee or the memory limit. A 7B model typically takes 32–128 samples per batch on a single card; small embedding models can go 1024+;
  • Parallelism: one worker per shard, one shard per card. Many shards in parallel = many cards in parallel;
  • Back-of-envelope: total runtime ≈ data volume ÷ (batch size × throughput × GPU count). One A10 running a quantized 7B model delivers roughly 1,000–3,000 tokens/s, so scoring ten million short texts takes about 3–8 hours.

The theory linking batch size and throughput is covered in performance optimization and capacity planning. If you have too many shards and each one is too small, scheduling and write overhead will exceed the inference itself — shard count ≈ 3–5× GPU count is the empirical sweet spot.

6. Checkpoints and Failure Retry: Idempotency + Resumable Runs ​

The easiest way batch inference falls over is dying halfway through and rerunning everything. Three practices fix it:

  1. Idempotent writes: write results with mode="overwrite" to a date-keyed path (s3://results/users/20250101/), so reruns never produce dirty data;
  2. Shard-level checkpoints: mark each shard as it succeeds (e.g. a .done file or a row in a job table), so the next scheduled run only processes the shards that didn't finish;
  3. Retry limits: Airflow's retries handles automatic retries; treat timeouts and data problems differently — retry a timeout, but page a human for bad data.
python
def score_shard(shard: str):
    done_flag = shard + ".done"
    if Path(done_flag).exists():
        return  # already succeeded; skip (the key to resumable runs)
    run_inference_shard(shard)
    Path(done_flag).touch()   # mark on success

Full Recompute vs. Incremental

Go incremental whenever you can. Daily profile jobs usually only need to recompute "users who changed" (active today, profile updated); the incremental hit rate depends on the business, and full recomputation is expensive but simple. A middle ground: incremental by default plus one full run per week for reconciliation, checking that incremental results match the full computation. Data-consistency issues are covered in common pitfalls and antipatterns.

7. Writing Results Back, Monitoring, and Alerting ​

Two write-back forms:

  • Result store: features/profiles go into a Feature Store or Redis (for online reads). Watch out: writing to an online store introduces an online-side dependency — keep concurrency and idempotency under tight control;
  • Object storage: scoring reports go to Parquet for BI and downstream consumers.

Three monitoring lines (see monitoring and observability):

  1. Task success rate: Airflow reports this out of the box; alert on failure (WeCom/DingTalk/Slack webhook);
  2. Data freshness: alert when "the results table hasn't updated in over 26 hours" — this catches the false success of "the task went green but the data never fully landed";
  3. Quality checks: after each shard's output, validate row counts and null rates, compare against the previous period, and alert when drift exceeds a threshold.

A Spark Resource Configuration Example ​

How you allocate resources for Spark batch inference directly determines how fast it runs and what it costs:

bash
# Submit a Spark job: 8 executors, 2 cores / 8 GB each
# The driver only coordinates; all inference load runs on the executors
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

Verdict first: size the executor count as "data volume ÷ per-executor throughput" — don't just crank the number up. A job that can't finish on 8 executors usually gets about 1.5–1.8× faster on 16 (shuffle/IO overhead eats into the gain); beyond that, returns diminish. In Spark the real bottleneck is often shuffle and resource waiting, not the operators themselves.

An Alerting Hook Example (Airflow Failure Callback) ​

python
# Airflow hook: fire a webhook the moment a task fails (generic format for WeCom/DingTalk/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"[Batch inference failed] {dag_id}/{task_id} @ {ts}",
    })

# Attach it in the DAG args
default_args = {**default_args, "on_failure_callback": on_failure}

Common Pitfalls and Troubleshooting ​

PitfallSymptomFix
Data skewOne shard crawls while the rest finish earlyCheck whether hash shards are balanced; split hot users into finer-grained shards
Worker OOMShard too large or batch too bigReduce batch size, increase shard count
False successTask is green but data is missingAdd freshness/row-count checks; see Section 7
Duplicate writes on rerunDoubled or corrupted resultsmode="overwrite" + date key + idempotency
Overlapping schedulesTwo rounds writing the same data at oncemax_active_runs=1
Competing with online trafficOnline P99 dragged down by batch inferenceSchedule batch off-peak (overnight), cap its GPU share
Resume-from-checkpoint brokenA rerun after failure still reprocesses everythingCheckpoint at shard granularity; mark each shard on success

Further Reading ​

References ​