Course Content
Live Coding Interview Prep
7 sections · 50 lessons
Implement zero-shot classification using embeddings.
What you need to know
Zero-shot classification means classifying into labels the system was never trained on. With embeddings the idea is simple: if a text and a label description mean similar things, their vectors point in similar directions.
Three ways to describe a label, from weakest to strongest:
| label representation | example for "billing" | quality |
|---|---|---|
| The bare name | "billing" | weak: one word embeds vaguely |
| A description | "invoices, charges, refunds and failed payments" | good |
| A prototype | the average vector of 5 real billing messages | usually best |
A prototype (or centroid) is the mean of several example vectors. It captures how customers actually write, not how the product team describes the category.
The "unknown" threshold matters as much as the labels. Without it, "hello" and "your app is great" get forced into whichever label is least far away.
Alternatives worth naming: an NLI-based zero-shot model (for example facebook/bart-large-mnli) is usually more accurate but slower; an LLM prompt with the label list is stronger still and costs a model call per text. With 20–50 labelled examples per class, a trained classifier on the same embeddings beats all of them.
1from collections.abc import Callable2import numpy as np34class ZeroShotClassifier:5 """Nearest label by cosine similarity between text and label embeddings."""67 def __init__(self, labels: dict[str, str], embed_fn: Callable, threshold: float = 0.3):8 self.names, self.embed_fn, self.threshold = list(labels), embed_fn, threshold9 self.matrix = self._unit(embed_fn([labels[n] for n in self.names])) # (k, d)1011 @staticmethod12 def _unit(vectors) -> np.ndarray:13 v = np.asarray(vectors, dtype=np.float32)14 return v / (np.linalg.norm(v, axis=1, keepdims=True) + 1e-10)1516 def predict(self, texts: list[str]) -> list[tuple[str, float]]:17 if not texts:18 return []19 scores = self._unit(self.embed_fn(texts)) @ self.matrix.T # (n, k)20 best = scores.argmax(axis=1)21 return [(self.names[j] if s[j] >= self.threshold else "unknown", round(float(s[j]), 3))22 for s, j in zip(scores, best)]The tricky parts:
scores = texts @ labels.Tcomputes every text-label cosine in one product: an (n, d) matrix times a (d, k) matrix gives (n, k).argmax(axis=1)picks the best label for each row (each text).- Label vectors are computed once in
__init__; only the texts are embedded per call. - A
staticmethodbecause normalising does not depend on the instance.
Complexity: setup is one embedding call for k labels. Prediction is one embedding call for n texts, O(n·k·d) for the product and O(n·k) for the argmax. Memory is O(k·d) for the labels.
A real-life example
Hand-made 3-dimensional vectors (money, delivery, account) make the scores checkable:
1VECS = {"invoices, charges, refunds and failed payments": [1.0, 0.1, 0.1],2 "late, missing or damaged deliveries": [0.1, 1.0, 0.0],3 "login, OTP and profile problems": [0.1, 0.0, 1.0],4 "I was charged twice for one order": [0.9, 0.3, 0.0],5 "OTP never arrives": [0.0, 0.1, 0.9],6 "Nice app!": [0.2, 0.2, 0.2]}7labels = {"billing": "invoices, charges, refunds and failed payments",8 "delivery": "late, missing or damaged deliveries",9 "account": "login, OTP and profile problems"}10clf = ZeroShotClassifier(labels, lambda ts: [VECS[t] for t in ts], threshold=0.7)11print(clf.predict(["I was charged twice for one order", "OTP never arrives", "Nice app!"]))12# [('billing', 0.971), ('account', 0.989), ('unknown', 0.686)]| text | billing | delivery | account | result |
|---|---|---|---|---|
| charged twice | 0.971 | 0.409 | 0.094 | billing |
| OTP never arrives | 0.109 | 0.110 | 0.989 | account |
| Nice app! | 0.686 | 0.632 | 0.632 | best 0.686 is below 0.7 → unknown |
"Nice app!" is roughly equally close to everything — the signature of a text that belongs to no class. The threshold catches it; without it, a compliment would open a billing ticket.
A support team setting up ticket routing on day one, before anyone has labelled data, can ship exactly this, and replace it with a trained model once a few hundred tickets have been labelled.
Follow-up questions to expect
- "How do you pick the threshold?" — Label 100 real messages, plot precision and recall of the "unknown" bucket against the threshold, and choose the trade-off the business wants.
- "A text belongs to two labels?" — Multi-label: return every label above the threshold instead of only the argmax.
- "Two labels keep getting confused?" — Rewrite their descriptions to contrast them ("refunds for cancelled orders" vs "charges on active orders"), or merge them if the business does not need the split.