Training vs Inference in Machine Learning: Cost, Latency, Drift
8 min readBytePatterns
Training vs inference: one slow loop that changes the parameters, then a fast pass that only reads them. Costs, batching, train/serve skew, with runnable code.
Every machine learning model lives two lives. In training, it is shown labelled examples over and over and its parameters are nudged until its predictions match. In inference, those parameters are frozen and the model answers one request after another. Same model, yet the two phases have different costs, hardware and failure modes, and interviewers ask about the split because so many production bugs live on it.
The problem it solves
A model that classifies sensor readings as hot or cold might take an hour to fit, once a week. Answering takes microseconds, but it happens on every request, millions of times a day, with a user waiting. Design both phases the same way and one of them is wrong.
So each is optimised for what it is. Training is optimised for quality per unit of compute. Inference is optimised for latency and cost per request. What connects them is one artifact: the trained model, plus everything it needs to read its inputs exactly as it did in training.
The intuition
Five contrasts carry the topic:
- What runs. Training runs a forward pass to predict, compares with the label, then a backward pass to compute how each parameter should move. Inference runs the forward pass only.
- What happens to the parameters. Training changes them on every step. Inference reads them and changes nothing; a serving model does not learn from the requests it answers.
- How often. Training runs on a release cadence: daily, weekly, or once for a big model. Inference runs per request, forever.
- What is measured. Training is judged by the quality reached per unit of compute. Inference is judged by latency percentiles, throughput and cost per answer.
- How it fails. Training fails loudly: the loss will not go down. Inference fails quietly: inputs drift, or serving prepares inputs differently from training, and accuracy slides with no error anywhere.
A rule of thumb for large neural networks (from memory): training costs roughly six floating-point operations per parameter per training token, a forward pass roughly two. But inference never stops, so over a popular model's lifetime the serving bill can exceed the training bill, which is why techniques exist purely to make inference cheaper, such as quantization and the KV cache.
Watch it run
The animation uses the lesson's four examples on one axis. Training fits parameters from labelled examples: two are labelled cold, at 1.0 and 2.0, and their average is 1.5; two are labelled hot, at 8.0 and 9.0, and their average is 8.5. Halfway between them is 5.0, and that one number is the entire model. Training is expensive, repeated and full of corrections, and then it is over: the parameter is frozen. Inference is one pass over those frozen parameters. A new input arrives, 6.5, and one comparison settles it: 6.5 is above 5.0, so the answer is hot. No parameter moved and nothing was learned. That single pass runs on every request, while training ran once: the same model with two cost profiles. Months later the readings drift upward and the old threshold starts calling cold things hot. The learned rule is stale, so it is refitted on newer data, which moves the threshold to 7.3: the fix is retraining, not a bigger prompt.
Training vs Inference
Step 1 of 11
Training fits parameters from labelled examples. Here are four of them.
The same interactive animation as the lesson — step through it with the controls.
The code
A toy model: logistic regression on one feature. Training standardises the feature with the training set's mean and spread, then makes 200 passes, a forward and a backward step per example. Inference is the forward step alone:
import math
import random
def make_data(rng, n, shift=0.0):
"""Temperature readings; the label is what a person wrote down (5% wrong)."""
data = []
for _ in range(n):
x = rng.uniform(0, 10) + shift
y = int(x - shift > 5)
if rng.random() < 0.05:
y = 1 - y
data.append((x, y))
return data
def train(data, epochs=200, lr=0.5):
"""TRAINING: many passes over every example, forward and backward."""
xs = [x for x, _ in data]
mean = sum(xs) / len(xs)
std = (sum((x - mean) ** 2 for x in xs) / len(xs)) ** 0.5
w = b = 0.0
passes = {"forward": 0, "backward": 0}
for _ in range(epochs):
for x, y in data:
z = (x - mean) / std
p = 1 / (1 + math.exp(-(w * z + b))) # forward
passes["forward"] += 1
w -= lr * (p - y) * z # backward: nudge the parameters
b -= lr * (p - y)
passes["backward"] += 1
return {"w": w, "b": b, "mean": mean, "std": std}, passes
def predict(model, x):
"""INFERENCE: one forward pass over frozen parameters."""
z = (x - model["mean"]) / model["std"]
return int(model["w"] * z + model["b"] > 0)
rng = random.Random(36)
train_set, test_set = make_data(rng, 800), make_data(rng, 200)
model, passes = train(train_set)
print(passes) # {'forward': 160000, 'backward': 160000}
print(sum(predict(model, x) == y for x, y in test_set) / len(test_set)) # 0.925
frozen = dict(model)
for x, _ in make_data(random.Random(1), 10_000):
predict(model, x) # 10,000 requests later...
print(model == frozen) # True
Training cost 160,000 forward and 160,000 backward steps; each answer costs one forward step, and ten thousand requests later the model is bit for bit unchanged. Note that mean and std are part of the model. Here is what happens when serving code normalises each batch with its own statistics, the classic train/serve skew, then batching and drift:
def predict_batch(model, xs):
"""Batched inference: the same forward pass, many inputs per call."""
return [predict(model, x) for x in xs]
def predict_skewed(model, xs):
"""A serving bug: normalise with the batch's own statistics, not training's."""
mean = sum(xs) / len(xs)
std = (sum((x - mean) ** 2 for x in xs) / len(xs)) ** 0.5
return [int(model["w"] * (x - mean) / std + model["b"] > 0) for x in xs]
warm = [6 + 4 * random.Random(s).random() for s in range(100)] # a hot afternoon
print(sum(predict_batch(model, warm)), sum(predict_skewed(model, warm))) # 100 54
def total_ms(calls, inputs): # illustrative: 4 ms per call, 0.05 ms per input
return calls * 4 + inputs * 0.05
print(total_ms(32, 32), total_ms(1, 32)) # 129.6 5.6
drifted = make_data(random.Random(2), 1000, shift=2.0) # the sensor now reads 2 high
score = lambda m: sum(predict(m, x) == y for x, y in drifted[800:]) / 200
new_model, _ = train(drifted[:800]) # retrain on fresh labels
print(score(model), score(new_model)) # 0.705 0.885
All 100 warm readings are hot and the correct pipeline says so; the skewed one calls 46 of them cold, because inside that batch they look average. Nothing crashed. Batching 32 requests cuts the illustrative total from 129.6 ms to 5.6 ms, but each request waits for the batch to fill. And when the sensor drifts, the frozen model falls to 0.705 until retraining restores it.
Finally, cross-checks on 300 seeded cases: batched answers equal one-at-a-time answers, training makes exactly epochs × n steps each way, and the lesson's midpoint threshold agrees with a brute-force "nearest class mean" on every point:
ok = True
for seed in range(300):
r = random.Random(seed)
m = {"w": r.uniform(-3, 3), "b": r.uniform(-3, 3),
"mean": r.uniform(0, 10), "std": r.uniform(0.5, 4)}
xs = [r.uniform(-5, 15) for _ in range(r.randint(1, 64))]
ok &= predict_batch(m, xs) == [predict(m, x) for x in xs]
d, epochs = make_data(r, r.randint(2, 40)), r.randint(1, 5)
ok &= train(d, epochs=epochs)[1] == {"forward": epochs * len(d), "backward": epochs * len(d)}
cold = [r.uniform(0, 5) for _ in range(r.randint(1, 6))]
hot = [r.uniform(5, 10) for _ in range(r.randint(1, 6))]
cm, hm = sum(cold) / len(cold), sum(hot) / len(hot)
threshold = (cm + hm) / 2 # the lesson's model
for x in (r.uniform(0, 10) for _ in range(50)):
if abs(x - threshold) > 1e-9: # skip exact ties
ok &= (x > threshold) == (abs(x - hm) < abs(x - cm))
print(ok) # True
The complexity
- Training (this toy):
O(epochs × n)forward and backward steps, plusO(n)for the statistics. - Inference:
O(1)per request here; in a neural network, a fixed amount of arithmetic proportional to the parameter count, per input or per generated token. - Memory: training holds parameters, gradients and optimiser state; inference holds the parameters and the current request's working data.
Where it goes wrong
- Train/serve skew. Preprocessing written twice, in a training notebook and in the service, drifts apart. Ship it with the model and test both paths on the same inputs.
- Statistics from the wrong data. Normalising with the test set or the live batch leaks information or, as above, changes answers.
- Expecting a serving model to learn. Corrections reach it only through a new training run or a deliberately designed online-learning system.
- Silent drift. Monitor live inputs and accuracy proxies; retraining is the fix, as the machine learning basics article shows.
- Batching without a deadline. A batch that waits to fill can blow the latency budget; cap the wait.
When it shows up in interviews
In ML system design ("serve this model at 10,000 requests a second"), the interviewer wants an offline training pipeline and an online serving path joined by a versioned model artifact. Follow-ups: batching, rollback to the previous model, detecting drift. In AI engineering rounds the split reappears as fine-tuning versus serving, which is why cheaper fine-tuning is its own topic. As of October 2026, large language models follow the same split: their parameters are frozen while they answer.
How to say it in an interview
"Training is the expensive loop of forward and backward passes that updates the parameters until held-out quality is good; it runs on a schedule. Inference is a forward pass over frozen parameters, once per request, so I optimise it for latency and cost per answer, batching under a deadline. A versioned artifact joins them and includes the preprocessing, because train/serve skew is the classic silent bug. When inputs drift, the fix is retraining on fresh data."