Mask2Former and universal segmentation
Instead of labelling every pixel, predict a fixed set of masks and give each one a label, which turns semantic, instance and panoptic segmentation into one model.
- 16 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Stop labelling pixels one by one. Predict a fixed stack of shaded regions, and write one word on each.
The analogy
Take a photo and a stack of one hundred sheets of tracing paper. On each sheet you may shade one region and write one word beside it.
Sheet one: shade the sky, write "sky". Sheet two: shade the left car, write "car". Sheet three: the right car. Sheets you do not need get stamped "nothing".
Lay the stack on the photo and you have described the scene completely. You never went pixel by pixel, and yet every pixel is covered.
That is the whole idea. It sounds like a small change in bookkeeping. It replaced three separate families of models with one.
Why it exists
Look at what the three older approaches had in common. Each committed to an output shape early, and that shape decided what the model could never say.
Shading every pixel with one category cannot count two touching cars. Outlining boxes and drawing inside them cannot describe sky. Stapling both together needs a merge rule that throws away good regions.
The tracing-sheet idea makes none of those commitments. A sheet is a region with a name. Sky is a region with a name. A single car is a region with a name.
The task then decides how you read the stack, not how you build it.
How one stack answers three questions
one stack of shaded sheets, each with a word
|
+-------------+--------------+
| | |
semantic instance panoptic
| | |
blend every keep only let the most
sheet by its the sheets confident sheet
word, per naming win each pixel
pixel countable
thingsNothing about the model changes between those three columns. Only the reading changes.
The two ideas that make it work
Nobody tells a sheet what to shade. During training, the model produces its stack and the correct regions are laid beside it. Then a matching step pairs them up. This sheet is closest to the sky, that one to the left car.
It is the same job as pairing dance partners so the total mismatch is smallest. Sheets left unpaired are told to say "nothing" next time. Which sheet handles which region is never decided by a human.
Each sheet only looks where it is working. When a sheet is refining the left car, staring at the whole photo is a distraction.
Think of reading a dense page with a card covering everything except the line you are on. You read that line better. The model does the same thing. The paper calls it masked attention: attention meaning where to look, masked meaning only inside my region.
That one restriction separates the first version of this idea from the version that beat every specialised model.
Where you have already seen this
- Photo apps that offer "select sky", "select subject" and "select background" from one tap, using one model.
- Video editing tools that separate a scene into its objects for recolouring.
- Robotics and driving stacks that need the road area and each pedestrian from a single pass.
What is honestly hard here
The stack has a fixed size. If a photo contains more regions than there are sheets, some region goes undescribed.
In practice the fixed size is comfortably larger than any real scene, so this rarely bites. It is still a genuine limit. Remember it before pointing one of these models at an aerial photo of a crowd.
Remember this
- Predict regions with names, not a name for every pixel.
- A matching step decides which prediction is responsible for which real region.
- One model, read three ways, covers all three segmentation tasks.
What to learn next
- Segment Anything — segmentation you steer with a click instead of a class list.
- Vision transformers — the backbone underneath most of these models.
- Attention — where the queries, keys and values in that decoder come from.
Developer — Code and libraries.
Setup
pip install numpy scipyThe full model needs a large download. The mechanism does not, and the mechanism is what people get wrong.
Below, six queries produce six masks, and those same six outputs are read three different ways.
The mechanism, end to end
import numpy as np
from scipy.optimize import linear_sum_assignment
H, W, D, Q = 8, 14, 4, 6 # image, embedding size, number of queries
CLASSES = ["sky", "road", "car", "no object"]
NO_OBJECT = 3
# ---- the pixel decoder's output: one embedding vector per pixel ----
region = np.zeros((H, W), int) # 0 sky, 1 road, 2 car A, 3 car B
region[4:] = 1
region[4:7, 1:5] = 2
region[4:7, 8:12] = 3
pixel_emb = np.eye(D)[region].transpose(2, 0, 1) # (D, H, W)
# ---- the transformer decoder's output: one row per query ----
# each query holds a mask embedding and a distribution over classes
mask_emb = np.array([
[1, 0, 0, 0], # q0 -> sky
[0, 1, 0, 0], # q1 -> road
[0, 0, 1, 0], # q2 -> car A
[0, 0, 0, 1], # q3 -> car B
[0, 0, 1, 0], # q4 -> car A again, a duplicate
[0, 0, 0, 0], # q5 -> nothing in particular
], float)
class_prob = np.array([
[0.97, 0.01, 0.01, 0.01],
[0.01, 0.96, 0.02, 0.01],
[0.02, 0.02, 0.94, 0.02],
[0.02, 0.03, 0.88, 0.07],
[0.03, 0.05, 0.52, 0.40],
[0.05, 0.05, 0.05, 0.85],
])
# ---- masks come from a dot product, nothing more ----
mask_logits = np.einsum("qd,dhw->qhw", mask_emb, pixel_emb) * 8 - 4
mask_prob = 1 / (1 + np.exp(-mask_logits))
print("one dot product turns", Q, "query vectors into", Q, "masks:", mask_logits.shape)
GLYPH = " .:-=+*#@"
def strip(m, row=5):
return "".join(GLYPH[int(v * 8)] for v in m[row])
print("\nrows 1 and 5 of each query mask (@ = confident foreground):")
for q in range(Q):
print(f" q{q} |{strip(mask_prob[q], 1)}| |{strip(mask_prob[q], 5)}| "
f"best class: {CLASSES[class_prob[q].argmax()]}")
# ---- ONE set of outputs, THREE different tasks ----
def semantic(mask_prob, class_prob):
thing = class_prob[:, :NO_OBJECT] # drop the no-object column
return np.einsum("qc,qhw->chw", thing, mask_prob).argmax(0)
THINGS = {2} # only cars are countable here
def instances(mask_prob, class_prob, thr=0.6):
out = []
for q in range(len(class_prob)):
c = class_prob[q].argmax()
if c in THINGS and class_prob[q, c] > thr:
out.append((q, CLASSES[c], class_prob[q, c], (mask_prob[q] > 0.5).sum()))
return out
def panoptic(mask_prob, class_prob):
score = class_prob[:, :NO_OBJECT].max(1)[:, None, None] * mask_prob
winner = score.argmax(0) # one query wins each pixel
keep = score.max(0) > 0.5
return np.where(keep, winner, -1)
sem = semantic(mask_prob, class_prob)
print("\nsemantic map from those queries (s sky, r road, c car):")
for row in sem:
print(" " + "".join("src"[v] for v in row))
print("\ninstances: countable classes only, score above 0.6:")
for q, name, s, area in instances(mask_prob, class_prob):
print(f" q{q} {name:5s} score {s:.2f} {area:>3} pixels")
pan = panoptic(mask_prob, class_prob)
print("\npanoptic map, the winning query id per pixel (. = nothing above 0.5):")
for row in pan:
print(" " + "".join("." if v < 0 else str(v) for v in row))
print(" note q4 wins nothing: q2 predicts the same mask with a higher class score")
# ---- training: which query is responsible for which ground-truth mask? ----
gt_masks = np.stack([region == 0, region == 1, region == 2, region == 3]).astype(float)
gt_class = [0, 1, 2, 2]
def dice_cost(p, g):
return 1 - (2 * (p * g).sum() + 1) / (p.sum() + g.sum() + 1)
cost = np.zeros((Q, len(gt_masks)))
for q in range(Q):
for g in range(len(gt_masks)):
cost[q, g] = dice_cost(mask_prob[q], gt_masks[g]) - class_prob[q, gt_class[g]]
print("\nmatching cost, queries down, ground-truth masks across:")
print(" " + "".join(f"gt{g} " for g in range(len(gt_masks))))
for q in range(Q):
print(f" q{q} " + "".join(f"{cost[q, g]:+7.3f} " for g in range(len(gt_masks))))
rows, cols = linear_sum_assignment(cost)
print("\nHungarian assignment (each ground-truth mask gets exactly one query):")
for q, g in zip(rows, cols):
print(f" q{q} -> gt{g} ({CLASSES[gt_class[g]]}), cost {cost[q, g]:+.3f}")
print(" unmatched queries are trained towards 'no object':",
[f"q{q}" for q in range(Q) if q not in rows])one dot product turns 6 query vectors into 6 masks: (6, 8, 14)
rows 1 and 5 of each query mask (@ = confident foreground):
q0 |##############| | | best class: sky
q1 | | |# ### ##| best class: road
q2 | | | #### | best class: car
q3 | | | #### | best class: car
q4 | | | #### | best class: car
q5 | | | | best class: no object
semantic map from those queries (s sky, r road, c car):
ssssssssssssss
ssssssssssssss
ssssssssssssss
ssssssssssssss
rccccrrrccccrr
rccccrrrccccrr
rccccrrrccccrr
rrrrrrrrrrrrrr
instances: countable classes only, score above 0.6:
q2 car score 0.94 12 pixels
q3 car score 0.88 12 pixels
panoptic map, the winning query id per pixel (. = nothing above 0.5):
00000000000000
00000000000000
00000000000000
00000000000000
12222111333311
12222111333311
12222111333311
11111111111111
note q4 wins nothing: q2 predicts the same mask with a higher class score
matching cost, queries down, ground-truth masks across:
gt0 gt1 gt2 gt3
q0 -0.952 +0.966 +0.969 +0.969
q1 +0.956 -0.929 +0.949 +0.949
q2 +0.937 +0.934 -0.864 +0.006
q3 +0.937 +0.924 +0.066 -0.804
q4 +0.927 +0.904 -0.444 +0.426
q5 +0.899 +0.889 +0.855 +0.855
Hungarian assignment (each ground-truth mask gets exactly one query):
q0 -> gt0 (sky), cost -0.952
q1 -> gt1 (road), cost -0.929
q2 -> gt2 (car), cost -0.864
q3 -> gt3 (car), cost -0.804
unmatched queries are trained towards 'no object': ['q4', 'q5']Reading the output
A mask is a dot product. np.einsum("qd,dhw->qhw", ...) multiplies each query's embedding against every pixel's embedding. That single line is the whole mask head: no crop, no fixed 28x28 grid, no upsampling from a box. Mask resolution is the resolution of the pixel embedding map.
One set of outputs, three readings, no retraining. The semantic reading blends every query weighted by its class probability, so duplicates are harmless. Here q2 and q4 predict the same car, and the blend does not care. The instance reading filters to countable classes and keeps confident queries. The panoptic reading takes an argmax over queries, so exactly one query owns each pixel.
Look at what happened to the duplicate. In the panoptic map, q4 wins nothing at all. It predicts the same mask as q2 but with a lower class score, so it loses every pixel by argmax. Compare that with panoptic segmentation, where the same duplicate needed an explicit survival rule and a hand-set threshold. Here the argmax handles it.
The cost matrix is where training decides responsibility. Row q2 reads +0.937, +0.934, -0.864, +0.006. Strongly negative against car A, near zero against car B. Negative means good, because the cost is mask dissimilarity minus class probability.
The assignment is a global optimum, not a greedy pick. linear_sum_assignment solves the Hungarian algorithm in cubic time. Greedy assignment would let q4, which also fits car A well, steal the pairing and leave q2 unmatched. The global solution avoids that.
q4 and q5 are pushed towards "no object". With 100 queries and 5 real regions, 95 queries learn to say nothing on every image. That is why the no-object class is usually down-weighted in the loss, typically by a factor of about 0.1.
Running the real thing
from transformers import AutoImageProcessor, Mask2FormerForUniversalSegmentation
from PIL import Image
import torch
name = "facebook/mask2former-swin-tiny-coco-instance" # about 190 MB of weights
processor = AutoImageProcessor.from_pretrained(name)
model = Mask2FormerForUniversalSegmentation.from_pretrained(name)
image = Image.open("your_photo.jpg").convert("RGB")
inputs = processor(image, return_tensors="pt")
with torch.no_grad():
outputs = model(**inputs)
# the two tensors the toy script above imitates
print(outputs.class_queries_logits.shape) # (batch, num_queries, num_classes + 1)
print(outputs.masks_queries_logits.shape) # (batch, num_queries, h, w)
result = processor.post_process_instance_segmentation(
outputs, target_sizes=[(image.height, image.width)]
)[0]
print(result["segmentation"].shape)No output block for that snippet, deliberately. It needs a network download and an image file, and the printed values depend on both. The shapes are the point: class_queries_logits and masks_queries_logits are exactly the class_prob and mask_logits of the toy script.
Written against transformers 5.6.2. The three post-processing calls are post_process_semantic_segmentation, post_process_instance_segmentation and post_process_panoptic_segmentation. All three accept the same outputs object. Checkpoints are trained per task, so ...-coco-instance and ...-coco-panoptic are different downloads. Sizes run 190 MB for Swin-tiny instance, 432 MB for Swin-base panoptic, and 866 MB for Swin-large ADE semantic.
Common mistakes
Calling the wrong post-processor for the checkpoint. A checkpoint fine-tuned for instance segmentation will produce a semantic map when asked, and it will be poor. Match the post-processor to what the checkpoint was trained on.
Expecting the number of queries to be the number of objects. With the Hugging Face default of 100 queries, most predict "no object". Filter by class probability before counting anything.
Reading masks_queries_logits as final. They come out at a reduced resolution and need interpolating to the original image size. The post-processing helpers do this, using target_sizes.
Applying a softmax over queries. Queries do not compete at a pixel during training; each mask is an independent sigmoid. Competition happens at inference, through the argmax in the panoptic reading.
Assuming masks are connected. A query may claim two disconnected blobs. That is correct for an occluded object and wrong for a duplicate, and only your data can tell you which.
Try it yourself
Raise q4's car probability from 0.52 to 0.99 so it beats q2. Re-run. The semantic map is unchanged, the instance list swaps one query for another, and the panoptic map hands the whole car to q4. Three readings, three different sensitivities to the same change.
What to learn next
- Segment Anything — segmentation you steer with a click instead of a class list.
- Vision transformers — the backbone underneath most of these models.
- Attention — where the queries, keys and values in that decoder come from.
Researcher — Mathematics and papers.
Mask classification, formally
Predict a fixed-size set $z = {(p_i, m_i)}_{i=1}^{N}$, where $p_i \in \Delta^{K}$ is a distribution over $K$ classes plus a no-object symbol $\varnothing$, and $m_i \in [0,1]^{H \times W}$ is a binary mask probability map. $N$ is a constant, typically 100, and is independent of the number of objects.
The mask comes from a dot product between per-query embeddings $\mathcal{E}{\text{mask}} \in \mathbb{R}^{N \times C}$ and a per-pixel embedding map $\mathcal{E}{\text{pixel}} \in \mathbb{R}^{C \times H \times W}$:
$$ m_i = \sigma!\left(\sum_{c} \mathcal{E}{\text{mask}}[i, c] \cdot \mathcal{E}{\text{pixel}}[c, \cdot, \cdot]\right) $$
Semantic inference marginalises over queries:
$$ \text{argmax}{c \in {1..K}} \sum{i=1}^{N} p_i(c) \cdot m_i[h, w] $$
Panoptic inference takes $\text{argmax}_i \; \max_c p_i(c) \cdot m_i[h,w]$ instead, then filters low-area segments.
Training
Ground truth is ${(c^{gt}_j, m^{gt}j)}{j=1}^{M}$ with $M \le N$, padded with $\varnothing$. A bipartite matching $\sigma$ is found by the Hungarian algorithm over the cost
$$ \mathcal{C}(i, j) = -p_i(c^{gt}j) + \mathcal{L}{\text{mask}}(m_i, m^{gt}_j) $$
with $\mathcal{L}_{\text{mask}}$ a weighted sum of binary cross-entropy and Dice. The training loss then applies cross-entropy on classes for all $N$ queries, and the mask loss only on matched pairs. The $\varnothing$ class is down-weighted, typically by 0.1, because it dominates by count.
This is DETR's set-prediction recipe (Carion et al., 2020) applied to masks. The matching, not the architecture, is what removes the need for non-maximum suppression: no two queries are ever rewarded for the same ground-truth region.
What Mask2Former added over MaskFormer
Cheng et al. (2022) list three changes, and the paper's ablations attribute most of the gain to the first.
- Masked attention. Standard cross-attention lets a query attend to all $HW$ locations. Masked attention restricts it to the foreground of that query's own prediction from the previous decoder layer: $$ X_{l} = \text{softmax}(\mathcal{M}_{l-1} + Q_l K_l^{\top}) V_l + X_{l-1} $$ where $\mathcal{M}_{l-1}[x, y] = 0$ if the previous-layer mask prediction at that location exceeds 0.5, and $-\infty$ otherwise. It converges faster and localises better. Note the subtlety. This is a hard, non-differentiable gate computed from the model's own previous output. It is iterative refinement, not learned attention sparsity.
- Multi-scale features round-robin. Decoder layers consume feature maps at strides 32, 16 and 8 in rotation, so small objects are not starved.
- Training efficiency. The mask loss is computed on a sampled subset of points rather than every pixel, using PointRend-style importance sampling. This cuts training memory by roughly a factor of three and enables the higher resolutions the method needs.
Reported results: 57.8 PQ on COCO panoptic, 50.1 AP on COCO instance, 57.7 mIoU on ADE20K semantic, all with one architecture.
Honest caveats
- One architecture, not one checkpoint. Mask2Former is trained separately per task and per dataset. OneFormer (Jain et al., 2023) closed that gap with a task-conditioned joint training scheme, which is the stronger claim to universality.
- Small objects remain the weak point. Fixed query count plus low-resolution mask embeddings means very small instances are under-served relative to a well-tuned detector.
- Convergence is slow. Set-prediction training inherits DETR's long schedules. Fifty epochs on COCO is normal, against twelve for a Mask R-CNN baseline.
- Query behaviour is not interpretable. Queries do not correspond to fixed classes or positions. Analyses show they specialise loosely by scale and location, and the specialisation shifts across training runs.
The wider pattern
Per-pixel classification, box-and-crop, and merge heuristics were all representational commitments made for engineering convenience. Set prediction removes them, at the cost of a matching step and slower convergence.
The same move recurs across vision: DETR for detection, TrackFormer for tracking, and the promptable formulation in Segment Anything. Whenever a task's output is a set, predicting a set directly beats predicting a grid and recovering the set afterwards.
Papers
- Carion et al., End-to-End Object Detection with Transformers (DETR), ECCV 2020 — arxiv.org/abs/2005.12872
- Cheng, Schwing, Kirillov, Per-Pixel Classification is Not All You Need (MaskFormer), NeurIPS 2021 — arxiv.org/abs/2107.06278
- Cheng, Misra, Schwing, Kirillov, Girdhar, Masked-attention Mask Transformer for Universal Image Segmentation, CVPR 2022 — arxiv.org/abs/2112.01527
- Jain et al., OneFormer: One Transformer to Rule Universal Image Segmentation, CVPR 2023 — arxiv.org/abs/2211.06220
- Xie et al., SegFormer, NeurIPS 2021 — arxiv.org/abs/2105.15203
What to learn next
- Segment Anything — segmentation you steer with a click instead of a class list.
- Vision transformers — the backbone underneath most of these models.
- Attention — where the queries, keys and values in that decoder come from.