Medical Imaging AI

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.

On this page 6
  1. Why it exists
  2. How it works
  3. Where you have already seen it
  4. An honest warning
  5. Remember this
  6. What to learn next

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

Developer — Code and libraries.

Setup

bash
pip install torch

Minimal runnable code

attention_mil.py
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()}")
        break
Output
epoch  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.attention scores each instance independently, then torch.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

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 = 0

This 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_k
  • h_k — the feature embedding of instance k (e.g. a CNN's output for one tile)
  • V, w — learned parameters of the attention mechanism
  • a_k — the resulting attention weight for instance k, summing to 1 across the bag
  • z — 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