Appearance
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
| Dimension | Online inference | Batch inference |
|---|---|---|
| Trigger | Request-driven, QPS fluctuates | Scheduled, workload known in advance |
| Resources | Always-on, ~3× peak redundancy | Spun up per batch, released when done |
| Cost | High (GPUs running around the clock) | An order of magnitude lower (table below) |
| Latency | Milliseconds to seconds | Minutes to hours |
| Data | One request at a time | Full or incremental datasets |
| Failure impact | A single request fails | The whole batch reruns; needs idempotency and checkpoints |
A set of cost figures (same A10 GPU, same 7B model):
| Deployment | Utilization | Cost per inference (amortized) |
|---|---|---|
| Online, always-on (3× peak redundancy) | 20–30% average | Baseline 1.0 |
| Online + autoscaling | 40–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 >> publishKey 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.parquetVerdict 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:
- Idempotent writes: write results with
mode="overwrite"to a date-keyed path (s3://results/users/20250101/), so reruns never produce dirty data; - Shard-level checkpoints: mark each shard as it succeeds (e.g. a
.donefile or a row in a job table), so the next scheduled run only processes the shards that didn't finish; - Retry limits: Airflow's
retrieshandles 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 successFull 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):
- Task success rate: Airflow reports this out of the box; alert on failure (WeCom/DingTalk/Slack webhook);
- 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";
- 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.pyVerdict 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
| Pitfall | Symptom | Fix |
|---|---|---|
| Data skew | One shard crawls while the rest finish early | Check whether hash shards are balanced; split hot users into finer-grained shards |
| Worker OOM | Shard too large or batch too big | Reduce batch size, increase shard count |
| False success | Task is green but data is missing | Add freshness/row-count checks; see Section 7 |
| Duplicate writes on rerun | Doubled or corrupted results | mode="overwrite" + date key + idempotency |
| Overlapping schedules | Two rounds writing the same data at once | max_active_runs=1 |
| Competing with online traffic | Online P99 dragged down by batch inference | Schedule batch off-peak (overnight), cap its GPU share |
| Resume-from-checkpoint broken | A rerun after failure still reprocesses everything | Checkpoint at shard granularity; mark each shard on success |
Further Reading
- Deployment architecture patterns — where batch mode sits in the deployment architecture spectrum
- Monitoring and observability — implementation details for task success rate, data freshness, and quality checks
- Performance optimization and capacity planning — the theory behind batch size and throughput curves
- The MLOps deployment pipeline — artifact management and triggering upstream of the batch pipeline
- Inference: from forward pass to inference engine — the mechanics of batched forward passes
- NVIDIA Triton multi-model serving — batch inference can also be accelerated with Triton's dynamic batching
References
- Apache Airflow documentation: https://airflow.apache.org/docs/
- Apache Spark documentation: https://spark.apache.org/docs/latest/
- Pandas documentation: https://pandas.pydata.org/docs/
- Airflow Core Concepts: Dags (incl. Dynamic Dags / Task Mapping): https://airflow.apache.org/docs/apache-airflow/stable/core-concepts/dags.html