Library / SDK
lucidrains/slot-attention avatar
lucidrains/slot-attention

lucidrains/slot-attention: Object-Centric Clustering in PyTorch

Implementation of Slot Attention from GoogleAI

499 stars34 forksPythonMIT

At a glance

What is it?
A PyTorch reimplementation of the Slot Attention mechanism from GoogleAI, plus an adaptive-slot wrapper from a 2024 follow-up paper. It gives you a differentiable module that turns a set of feature vectors into a fixed number of slot vectors, and it is small enough to read in one sitting.
Who is it for?
Adopt this if you want the Slot Attention mechanism as a plain PyTorch module you can drop into a larger model, and you already have a training loop and a dataset of feature vectors to feed it. Do not adopt it expecting a complete object-discovery pipeline: the repository ships the module and a test suite, and the README does not document a training script, a pretrained checkpoint or a dataset loader.
Can I use it commercially?
Yes. MIT is a permissive licence: you can use, modify and sell software built on it, as long as you keep its copyright and licence notices.
Is it still maintained?
Yes. The repository last received commits 115 days ago.
What is it written in?
Mainly Python, according to GitHub's language statistics.

Answers come from the project's GitHub data, last synced on September 29, 2026, and from our analysis. They are not legal advice.

Editorial analysis

What Slot Attention solves, and who this repository is for

Slot Attention addresses a specific problem in representation learning: given a set of input feature vectors, produce a small set of vectors that each correspond to one object or entity in the scene, without being told how many objects there are or which input vectors belong to which. The paper this repository implements is 'Object-Centric Learning with Slot Attention' by Locatello and co-authors, cited in the README with its arXiv identifier 2006.15055. The README also links a video that describes what the network can do, and notes that Google released an official implementation under google-research.

The audience is narrow and identifiable. You are an engineer or researcher who already has a backbone that produces per-position features, and you want the clustering step that sits on top of them. The README's own example reflects this: the input is a tensor of shape (2, 1024, 512), meaning two scenes, 1024 feature vectors per scene, 512 dimensions each, and the output is (2, 5, 512), meaning five slots per scene. Nothing in the repository converts images or video into those 1024 vectors. That part is yours.

This is a module, not an application. There is no CLI, no server, no configuration file for a training run. The top-level entries are a package directory, a tests directory, a single training script named train_duplicate_slots.py, a pyproject.toml, a README, a licence and a diagram image. If you want an end-to-end object discovery system, this is one component of it.

How the slot attention mechanism is wired in this implementation

The mechanism follows the paper's iterative scheme. Slots are initialized, then refined over a fixed number of attention iterations. In the README example, iters defaults to 3, and the constructor takes num_slots and dim alongside it. Each iteration computes attention from the slots over the input features, normalizes it so that input positions compete for slots (rather than slots competing for inputs, which is the inverted direction used in standard transformer attention), and updates the slots with the resulting weighted average.

The package exposes at least two classes, SlotAttention and MultiHeadSlotAttention, both importable from the top-level slot_attention package. The multi-head variant takes dim, num_slots and iters in the README's example. The README does not spell out the internal differences between the two beyond the name, so if you need the head count semantics you have to read slot_attention/ directly.

One property the README states explicitly is worth pausing on: after training, the network is reported to generalize to a slightly different number of slots, and you can override the count at call time with the num_slots keyword in forward. That means the slot count is not baked into the weights. It is a real convenience when you do not know the object count in advance, but the README qualifies it as 'slightly different', which is a hedge rather than a guarantee. Treat the override as a small perturbation, not as a licence to train with 5 slots and deploy with 50.

Dependencies listed in pyproject.toml are einops (>=0.8.0), einx (>=0.3.0), torch, and vector-quantize-pytorch. The presence of einx and vector-quantize-pytorch is a hint that the package has grown past the original paper's scope; the README's citation list includes a 2026 NeurIPS paper on MetaSlot and an ICLR 2026 paper on orthogonality in object-centric representations, so the package tracks the surrounding literature rather than freezing at the 2020 method.

Installing slot_attention and running the README example

The README gives a single install command. It pulls the package from PyPI along with the dependencies declared in pyproject.toml, which are einops, einx, torch and vector-quantize-pytorch. Note that the distribution name on PyPI uses an underscore, while the import name in code also uses an underscore.

bash
pip install slot_attention

The first real use is the constructor plus a forward pass. The README's example builds five slots of dimension 512 over 1024 input vectors per scene, with three attention iterations, and passes a random tensor to check the shapes before you wire in real data.

python
import torch
from slot_attention import SlotAttention

slot_attn = SlotAttention(
    num_slots = 5,
    dim = 512,
    iters = 3   # iterations of attention, defaults to 3
)

inputs = torch.randn(2, 1024, 512)
slot_attn(inputs) # (2, 5, 512)

What you should see is a tensor of shape (2, 5, 512): the batch dimension is preserved, the 1024 input positions are collapsed into 5 slots, and the feature dimension is unchanged. If your output is not that shape, the mismatch is almost always in dim, which must match the last dimension of your input.

The second pattern the README documents is overriding the slot count at call time, which is the feature that makes the module usable when the object count varies between samples.

python
slot_attn(inputs, num_slots = 8) # (2, 8, 512)

The third pattern is the adaptive slot wrapper, which wraps a slot attention module and returns both the slots and a per-slot keep mask. The README attributes this to the adaptive slot paper, arXiv 2406.09196, and describes the mask as a differentiable one-hot decision about whether to use a slot.

python
import torch
from slot_attention import MultiHeadSlotAttention, AdaptiveSlotWrapper

slot_attn = MultiHeadSlotAttention(
    dim = 512,
    num_slots = 5,
    iters = 3,
)

adaptive_slots = AdaptiveSlotWrapper(
    slot_attn,
    temperature = 0.5 # gumbel softmax temperature
)

inputs = torch.randn(2, 1024, 512)
slots, keep_slots = adaptive_slots(inputs) # (2, 5, 512), (2, 5)

The README then gives the auxiliary loss directly: keep_aux_loss = keep_slots.sum(), to be added to your main loss with some weight. The temperature argument is described as the gumbel softmax temperature. The README does not state a recommended weight for the auxiliary term or how temperature should be scheduled, so that is a hyperparameter you tune yourself.

Where slot attention is the wrong tool, and what the README leaves out

The clearest limitation is scope. The repository does not include a training pipeline for object discovery. There is a train_duplicate_slots.py at the top level, and the name suggests it exercises a specific duplication scenario rather than reproducing the paper's experiments on CLEVR-style data. The README does not document a dataset, a pretrained checkpoint, an evaluation metric or a reproduction command. If your goal is to reproduce the paper's numbers, you are building that harness yourself.

The second limitation is the slot-count override. The README says the trained network generalizes to a slightly different number of slots. It does not quantify 'slightly', and it does not say what happens outside that range. Attention in this mechanism is normalized across slots, so changing the count changes the competition among slots; a large change is a different computation, not a free parameter change.

The third is the adaptive slot path. It is documented as a wrapper with a temperature and an auxiliary sum loss, and that is the extent of it. There is no guidance on what weight to give keep_aux_loss, no note on whether the temperature is annealed during training, and no statement about how the mask behaves at initialization. The README also does not document rollback or version pinning behaviour, so if you depend on a specific signature you should pin the version in your own requirements rather than tracking the latest release.

The fourth is dependency surface. The package imports einops, einx and vector-quantize-pytorch. Those are small libraries, but they are additional code in your environment for what is conceptually a single attention block, and vector-quantize-pytorch in particular is a much larger package than this use case requires. If you are vendoring the mechanism into a constrained environment, that dependency list is the thing you will end up editing.

The official Google implementation and what changes between them

The README itself points to the alternative: Google's official repository, released under google-research/slot_attention, which the README announces as an update. The difference in approach is not cosmetic. The official release is the authors' code accompanying the paper, and its purpose is to reproduce the paper's experiments, which means it comes with the surrounding training and evaluation machinery that this package does not ship. This package is a single PyTorch module designed to be imported into someone else's model.

That split determines which one you want. If you are trying to understand the mechanism, or you want to attach slots to an existing architecture, this package is the shorter path: one pip install, one constructor, one forward call, and the shapes are documented in the README. If you want the paper's results on the paper's data, the official repository is the reference point, because that is what it was written to do. Choosing this package for reproduction work means writing the data pipeline, the training loop and the evaluation yourself, and then discovering whether your numbers are comparable to the paper's at all.

There is also a middle option worth naming: reading the module and copying the attention block into your own codebase. The mechanism is not large, and the package's own dependency list suggests it has accumulated features beyond the original paper. If you only need the core iteration, the official paper and the official repository are both available as references, and neither imposes the einx or vector-quantize-pytorch dependencies.

Licence, maintenance and the cost of upgrading

The licence is MIT, declared both in the LICENSE file at the top level and in the pyproject.toml licence field. MIT is permissive: it allows use, modification and redistribution, including in closed-source products, provided the copyright notice and licence text are retained. This is a summary of what the licence identifier means, not legal advice; if your organisation has a policy on third-party licences, run the LICENSE file past whoever owns that policy, and note that the dependencies (torch, einops, einx, vector-quantize-pytorch) carry their own licences, which the README does not discuss.

The repository is not archived, and the last push was on 2026-06-06. The most recent release listed is 1.4.0 from 2024-08-20, while pyproject.toml declares version 1.5.2, so the version in the source tree is ahead of the newest release entry. If you install from PyPI you get a released version; if you install from the repository you get 1.5.2. Those are not necessarily the same code, and the README does not document a changelog, so the difference between them is something you would have to diff yourself.

Upgrade cost is the practical concern. The package depends on einops >= 0.8.0 and einx >= 0.3.0 with lower bounds but no upper bounds, and on torch with no version constraint at all. That means a fresh install resolves against whatever is current, and a torch upgrade in your environment can change behaviour under this package without any change to the package itself. The README does not document a compatibility matrix. Pinning slot_attention and its dependencies together in your own lockfile is the only mechanism the repository gives you for reproducible installs. The declared Python support in pyproject.toml is 3.6, which is old enough that it tells you little about what is actually tested today; the tests directory is the place to look for what the maintainer runs.

Editorial conclusion

Adopt this if you want the Slot Attention mechanism as a plain PyTorch module you can drop into a larger model, and you already have a training loop and a dataset of feature vectors to feed it. Do not adopt it expecting a complete object-discovery pipeline: the repository ships the module and a test suite, and the README does not document a training script, a pretrained checkpoint or a dataset loader. Before you build on it, check the pyproject.toml dependency list against your environment, since it pins einops and einx and pulls in vector-quantize-pytorch, and read the forward signature in slot_attention/ to confirm how the iters and num_slots arguments behave in the version you install.

Frequently asked questions

What is slot attention in lucidrains/slot-attention?

It is a PyTorch implementation of the Slot Attention mechanism from the paper 'Object-Centric Learning with Slot Attention' by Locatello and co-authors. The module takes a batch of feature vectors and returns a fixed number of slot vectors, which the paper uses to represent individual objects in a scene.

How do I install slot-attention?

The README gives one command: pip install slot_attention. That pulls in the dependencies declared in pyproject.toml, which are einops, einx, torch and vector-quantize-pytorch.

Can I change the number of slots after training in slot-attention?

Yes. The README states that after training the network is reported to generalize to a slightly different number of slots, and you can override the count with the num_slots keyword in forward, for example slot_attn(inputs, num_slots = 8). The README does not say how far outside the trained count this holds.

Does lucidrains/slot-attention include a training script or pretrained model?

The repository contains a train_duplicate_slots.py at the top level and a tests directory, but the README does not document a dataset, a pretrained checkpoint or a reproduction command. The README's usage section shows the module being called directly on a tensor.

What is the adaptive slot wrapper in slot-attention?

The README describes AdaptiveSlotWrapper as an implementation of the adaptive slot method from arXiv 2406.09196, which produces a differentiable one-hot mask for whether to use each slot. It wraps a slot attention module, takes a gumbel softmax temperature, and returns the slots plus a keep mask whose sum the README suggests adding to your main loss.

Is there an official slot attention implementation besides this one?

Yes. The README notes that Google released the official repository under google-research/slot_attention. That release accompanies the paper, while this package is a standalone PyTorch module intended to be imported into another model.

Official sources

  1. Issues
  2. License: MIT
  3. lucidrains/slot-attention on GitHub
  4. README
  5. Releases
Add this badge to your README

If you maintain this project, the badge below links readers to this analysis and shows its maintenance status from the daily GitHub snapshot. Paste the markdown into your README; add ?metric=license or ?metric=stars to the image URL for a different field.

Add this badge to your README

markdown
[![Hysen Labs](https://hysenlabs.com/badge/lucidrains-slot-attention.svg)](https://hysenlabs.com/projects/lucidrains-slot-attention)