Model extraction is straightforward in principle: query the API, collect input-output pairs, train a student model on them. For a classifier with a well-defined decision boundary, a few tens of thousands of well-chosen queries can produce a student that agrees with the victim on the large majority of inputs.
What makes it efficient is active learning. The attacker does not sample randomly — they concentrate queries near the decision boundary, where each answer is maximally informative. Uncertainty sampling can cut the query budget by an order of magnitude versus random sampling.
Two things follow that shape every defence:
What leaks scales with output granularity. A full probability vector leaks far more per query than a top-1 label. Confidence scores are the attacker’s gradient signal.
Prevention is not achievable. If the API is useful, it is extractable — the information the legitimate user needs is the information the attacker needs. The goal is to make extraction cost more than licensing, and to detect it while it is in progress.
Defence 1: reduce what each answer reveals
from future import annotations
import hashlib
import math
from dataclasses import dataclass
@dataclass
class OutputPolicy:
top_k: int = 3 # never return the full vector
round_to: int = 2 # decimal places on probabilities
temperature: float = 1.0 # >1 flattens, leaking less boundary detail
deterministic_noise: bool = True
def defend_output(
probs: dict[str, float], request_key: str, policy: OutputPolicy
) -> dict[str, float]:
"""Truncate, round, and add small DETERMINISTIC noise.
Determinism matters: random noise per call is averaged away by an
attacker who queries the same point repeatedly, and it costs you
reproducibility. Derive the perturbation from the input so the same
input always gets the same answer."""
if policy.temperature != 1.0:
t = policy.temperature
logits = {k: math.log(max(v, 1e-9)) / t for k, v in probs.items()}
m = max(logits.values())
exp = {k: math.exp(v - m) for k, v in logits.items()}
total = sum(exp.values())
probs = {k: v / total for k, v in exp.items()}
top = dict(sorted(probs.items(), key=lambda kv: -kv[1])[: policy.top_k])
if policy.deterministic_noise:
seed = int.from_bytes(
hashlib.blake2b(request_key.encode(), digest_size=4).digest(), "big")
for i, k in enumerate(top):
# +/- half a rounding unit, stable per input.
jitter = (((seed >> (i * 5)) % 1000) / 1000 - 0.5) * (10 ** -policy.round_to)
top[k] = max(0.0, min(1.0, top[k] + jitter))
return {k: round(v, policy.round_to) for k, v in top.items()}
Rounding to 2 decimal places sounds trivial. It is not — it removes most of the fine-grained boundary information that uncertainty sampling depends on, and it is invisible to nearly every legitimate consumer. This is the highest value-per-effort defence on the list, and the one most often skipped because it feels like degrading the product.
Defence 2: detect the query distribution, not the volume
An extractor’s query pattern is statistically distinct from a user’s. Legitimate traffic clusters around a few common cases. Extraction traffic is unusually uniform over the input space, unusually close to the decision boundary, and has unusually low repetition.
@dataclass
class ExtractionSignals:
"""Per-principal, over a rolling window. None of these is conclusive
alone; the combination is."""
queries: int = 0
unique_inputs: int = 0
near_boundary: int = 0 # confidence in [0.4, 0.6]
mean_pairwise_distance: float = 0.0
distinct_classes_returned: int = 0
def score(self) -> tuple[float, list[str]]:
reasons: list[str] = []
s = 0.0
if self.queries < 200:
return 0.0, ["insufficient volume to assess"]
uniqueness = self.unique_inputs / self.queries
if uniqueness > 0.97:
s += 0.3
reasons.append(f"near-zero repetition ({uniqueness:.2f})")
boundary_frac = self.near_boundary / self.queries
if boundary_frac > 0.35:
s += 0.4
reasons.append(f"{boundary_frac:.0%} of queries near the decision boundary")
# Real users do not sweep the whole input space evenly.
if self.mean_pairwise_distance > 0.8:
s += 0.2
reasons.append("queries uniformly spread across input space")
if self.distinct_classes_returned >= 0.9 * 10:
s += 0.1
reasons.append("coverage of nearly all output classes")
return min(s, 1.0), reasons
The boundary-concentration signal is the strongest. A user classifying real documents gets confident predictions most of the time; an attacker doing uncertainty sampling is deliberately querying where the model is unsure, and that shows up immediately as an anomalous confidence histogram.
Defence 3: watermark the decision boundary
You cannot stop extraction, so make a stolen model provable. Train the victim to produce specific, arbitrary outputs on a secret set of trigger inputs — points far from the natural data distribution, where the behaviour costs nothing on real traffic. A student distilled from the API inherits them.
def verify_ownership(suspect_predict, triggers: list[tuple], alpha: float = 1e-6) -> dict:
"""Under the null hypothesis (independent model), matching the trigger
labels is chance. With enough triggers the binomial p-value is tiny."""
n = len(triggers)
matches = sum(1 for x, y in triggers if suspect_predict(x) == y)
p_chance = 1 / 10 # number of classes
p_value = sum(
math.comb(n, k) * p_chance*k * (1 - p_chance) * (n - k)
for k in range(matches, n + 1)
)
return {"triggers": n, "matches": matches, "p_value": p_value,
"claim_supported": p_value < alpha}
With 100 triggers over 10 classes, chance gives about 10 matches; 40 matches produces an astronomically small p-value. That is evidence you can put in front of a lawyer, which is the actual remedy — technical defences delay extraction, legal ones address it.
What to actually deploy
In order of value per unit of effort:
Round and truncate outputs. Free, instant, removes most of the leak.
Authenticate everything and rate limit per principal. An unauthenticated endpoint cannot be defended at all, because the attacker rotates IPs and there is nothing to bind the budget to.
Log the confidence distribution per principal and alert on anomalies. Cheap, and it catches extraction while it is running rather than after.
Watermark. Moderate effort, and the only one that gives you a remedy after the fact.
Differential privacy on outputs. A real guarantee, a real accuracy cost. Reach for it when the model itself is the product.
The honest limitation
None of this stops a determined, well-resourced attacker who is patient. They can spread queries across accounts, sample slowly, mimic legitimate distributions and accept a lower-fidelity student.
What the defences do is change the economics. Extraction that needs 500,000 queries across 50 accounts over three months, with a detection risk at every step, is a different proposition from one that needs 20,000 queries in an afternoon. For most models that difference is the whole game — and for the handful where it is not, the answer is not to serve the model over a public API at all.
We build and secure ML systems at SoluLab — more on our machine learning development work.
Top comments (0)