Multiple instance learning
When only a whole slide is labelled and no single tile is, a model has to learn which few tiles actually earned that label on their own.
- 10 min read
- 3 reading levels
- Published
Read these first
On this page 6
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Multiple instance learning means a model gets one label for a whole group of items.
Think about a bag of grapes. If even one grape inside is rotten, you call the whole bag "bad." Every other grape in it might be perfectly fine.
A pathology slide works the same way. A slide gets one label: "cancer" or "no cancer." Nobody marks which tiny patch of tissue, out of thousands of tiles, made it cancer.
Why it exists
The previous lesson showed that a single slide can produce well over a hundred thousand tiles. Labelling each tile by hand, at that scale, is not realistic for any pathology department.
What pathologists actually record is a diagnosis for the whole slide — sometimes for the whole patient. Multiple instance learning (MIL) was built for exactly this gap. A group of items — a "bag" — shares one label. A model has to find which items inside actually explain it, without ever being told directly.
How it works
Bag (one slide) = many tiles
tile 1: normal tile 2: normal tile 3: normal
tile 4: CANCEROUS <-- this one tile made the whole bag positive
tile 5: normal tile 6: normal
Whole bag label: "cancer" (because of tile 4 alone)A trained model learns to give more "attention" to the tiles that matter. It gives less to the ones that do not. Nobody tells it in advance which is which.
Where you have already seen it
- A spam folder rule can flag an entire email thread as spam, because of one suspicious message inside it. That is the same "one bad item spoils the bag" logic.
- A batch quality check on a shipment works the same way. One faulty item fails the whole crate.
An honest warning
A model trained this way can be right about the whole slide, yet wrong about which tile actually mattered. Checking that a model's attention lands on genuinely relevant tissue is a real, necessary validation step. It should never land on scanning artifacts or irrelevant structures instead.
No pathology MIL model belongs in a real diagnosis without a pathologist's review. Regulatory clearance for that specific use matters too.
Remember this
- Multiple instance learning trains on one label per group of items, not one label per item.
- A trained model learns to weight the group's individual items differently, favoring the ones that explain the label.
- Getting the group-level answer right does not guarantee the model is looking at the right item for the right reason.
What to learn next
- Gigapixel pathology slides — the scale problem that makes this learning setting necessary.
- What is a neural network? — the general building block this lesson's model is made from.
- Why your model fails at the next hospital — a different failure mode that can affect the same pathology pipeline.
Developer — Code and libraries.
Setup
pip install torchMinimal runnable code
import torch
import torch.nn as nn
torch.manual_seed(0)
# A "bag" is one slide: a set of instances (patches), with only ONE label
# for the whole bag. A bag is positive if AT LEAST ONE patch inside it is
# cancerous -- exactly like a bag of grapes being "bad" if one grape is rotten.
def make_bags(n_bags, bag_size=8, feat_dim=4):
bags, labels = [], []
for _ in range(n_bags):
bag = torch.randn(bag_size, feat_dim)
is_positive = torch.rand(1).item() < 0.5
if is_positive:
# plant exactly one "cancerous" patch, shifted in feature space
idx = torch.randint(0, bag_size, (1,)).item()
bag[idx] += 3.0
bags.append(bag)
labels.append(float(is_positive))
return bags, torch.tensor(labels)
class AttentionMIL(nn.Module):
"""Scores every instance, turns scores into attention weights, and
pools instances into one bag-level prediction."""
def __init__(self, feat_dim=4, hidden=8):
super().__init__()
self.attention = nn.Sequential(nn.Linear(feat_dim, hidden), nn.Tanh(), nn.Linear(hidden, 1))
self.classifier = nn.Linear(feat_dim, 1)
def forward(self, bag):
scores = self.attention(bag) # one score per instance
weights = torch.softmax(scores, dim=0) # how much each instance matters
pooled = (weights * bag).sum(dim=0) # weighted average -> one bag vector
return self.classifier(pooled), weights
train_bags, train_labels = make_bags(200)
test_bags, test_labels = make_bags(50)
model = AttentionMIL()
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
loss_fn = nn.BCEWithLogitsLoss()
for epoch in range(30):
total_loss = 0.0
for bag, label in zip(train_bags, train_labels):
optimizer.zero_grad()
logit, _ = model(bag)
loss = loss_fn(logit.squeeze(), label)
loss.backward()
optimizer.step()
total_loss += loss.item()
if epoch % 10 == 0 or epoch == 29:
print(f"epoch {epoch:2d} avg loss {total_loss / len(train_bags):.4f}")
correct = 0
for bag, label in zip(test_bags, test_labels):
logit, _ = model(bag)
pred = (torch.sigmoid(logit) > 0.5).float().item()
correct += (pred == label.item())
print(f"test bag-level accuracy: {correct / len(test_bags):.2f}")
# Did attention actually find the planted instance in a positive bag?
for bag, label in zip(test_bags, test_labels):
if label.item() == 1.0:
_, weights = model(bag)
print(f"a positive bag's attention weights: {[round(w, 3) for w in weights.squeeze().tolist()]}")
print(f"most-attended instance index: {weights.argmax().item()}")
breakepoch 0 avg loss 0.4178 epoch 10 avg loss 0.0165 epoch 20 avg loss 0.0021 epoch 29 avg loss 0.0007 test bag-level accuracy: 0.98 a positive bag's attention weights: [0.0, 0.014, 0.0, 0.961, 0.0, 0.0, 0.008, 0.017] most-attended instance index: 3
What actually happened
Training loss drops from 0.4178 to 0.0007 over 30 epochs, and the model reaches 98% bag-level accuracy on unseen test bags — despite never being told which instance inside any bag was the planted, "cancerous" one.
Look at the attention weights on the final line: one instance, index 3, got 96.1% of the model's attention. Every other instance in that bag got close to zero. The model found the planted signal entirely on its own, purely by learning which pattern of weights explained the bag-level label correctly.
Line by line, the parts that are not obvious:
self.attentionscores each instance independently, thentorch.softmax(scores, dim=0)turns those scores into weights that sum to 1 across the whole bag — this is what lets the model focus on a small number of instances.pooled = (weights * bag).sum(dim=0)is the actual aggregation step: a weighted average across instances, turning a variable-length bag into one fixed-size vector for the classifier.- The final loop specifically inspects a positive bag's attention weights — this check is the closest thing MIL offers to interpretability, and it is worth running on every trained model, not only this one.
Common mistakes
Using simple average pooling instead of attention pooling. Averaging every instance equally drowns out a signal that lives in only one or two instances out of many — exactly the situation a real pathology slide presents.
Assuming high bag-level accuracy means the model is looking at the right instance. It is possible for a model to reach the correct bag-level answer for the wrong reason. Checking attention weights against known-relevant regions is a necessary sanity check, not an optional one.
Training with very large, unbalanced bags without care. A bag with thousands of instances and only one relevant one is a much harder optimisation problem than this toy example's eight-instance bags — real pathology MIL systems need larger models and more careful training than this illustration shows.
Try it yourself
Change bag_size=8 to bag_size=64, making the "needle in a haystack" problem harder. Retrain and see how many more epochs the model needs before the loss drops as sharply as it did here.
What to learn next
- Gigapixel pathology slides — the tiling process that produces the instances a MIL model consumes.
- What is a neural network? — the underlying building block behind the attention and classifier modules above.
- Grad-CAM and saliency maps — a different family of tools for checking what a vision model is actually looking at.
Researcher — Mathematics and papers.
Formal definition
In multiple instance learning, training data consists of bags X_i = {x_i1, ..., x_iK}, each with a single label Y_i, where individual instance labels y_ik are never observed. The standard MIL assumption is:
Y_i = 1 if at least one y_ik = 1 for x_ik in X_i
Y_i = 0 if every y_ik = 0This is the standard MIL assumption, directly matching the "one rotten grape" framing: a bag is positive if and only if it contains at least one positive instance. Some domains use relaxed variants (e.g. requiring a threshold count or proportion of positive instances), but the standard assumption is the most common starting point in computational pathology.
Attention-based deep MIL
Ilse, Tomczak & Welling (2018) introduced a learnable, permutation-invariant pooling operator for MIL, replacing fixed pooling (max, mean) with attention weights learned end to end:
a_k = exp( w^T tanh(V h_k) ) / SUM_j exp( w^T tanh(V h_j) )
z = SUM_k a_k * h_kh_k— the feature embedding of instancek(e.g. a CNN's output for one tile)V,w— learned parameters of the attention mechanisma_k— the resulting attention weight for instancek, summing to 1 across the bagz— the bag-level representation, a weighted sum of instance embeddings, fed to a final classifier
Because z is a sum over instances with learned, data-dependent weights, it is permutation-invariant: shuffling the order of instances in the bag does not change the output, a property essential for a set of unordered tile embeddings.
CLAM and clustering-constrained attention
CLAM (Lu et al., 2021) extends attention MIL for whole-slide pathology specifically, adding an auxiliary clustering objective that encourages the learned instance embeddings to separate meaningfully by class, and supporting multi-class slide classification. Its attention weights are explicitly designed to be interpretable, producing a heatmap over the original slide that highlights which tissue regions most influenced the diagnosis.
Cost
For a bag of K instances with d-dimensional embeddings, attention pooling costs O(K*d) for the attention scores and O(K*d) for the weighted sum — linear in bag size, which matters given that real pathology bags can contain K on the order of 10^4 to 10^5 instances after tiling and background removal, as covered in gigapixel pathology slides.
Key references
- Ilse, M., Tomczak, J. & Welling, M. (2018). Attention-based Deep Multiple Instance Learning. ICML.
- Lu, M. et al. (2021). Data-efficient and weakly supervised computational pathology on whole-slide images. Nature Biomedical Engineering 5.
- Dietterich, T., Lathrop, R. & Lozano-Pérez, T. (1997). Solving the Multiple Instance Problem with Axis-Parallel Rectangles. Artificial Intelligence 89(1-2). The original formalisation of the MIL problem.
- Campanella, G. et al. (2019). Clinical-grade computational pathology using weakly supervised deep learning on whole slide images. Nature Medicine 25.
Current state
Attention-based MIL, and CLAM specifically, remain the dominant approach for whole-slide classification with slide-level labels as of this writing. Transformer-based set architectures for MIL are an active research direction, aiming to model interactions between instances rather than scoring each instance independently, though attention-MIL's simplicity and interpretability keep it a strong default baseline in practice.
What to learn next
- Gigapixel pathology slides — the preprocessing pipeline that produces this lesson's input bags.
- Grad-CAM and saliency maps — complementary interpretability tools outside the MIL setting.
- Why your model fails at the next hospital — a generalisation risk that applies directly to MIL pathology models.