Sequential Attend, Infer, Repeat: Generative Modelling of Moving Objects

Adam R. Kosiorek, Hyunjik Kim, Ingmar Posner, Yee Whye Teh

Introduction

The ability to identify objects in their environments and to understand relations between them is a cornerstone of human intelligence . Arguably, in doing so we rely on a notion of spatial and temporal consistency which gives rise to an expectation that objects do not appear out of thin air, nor do they spontaneously vanish, and that they can be described by properties such as location, appearance and some dynamic behaviour that explains their evolution over time. We argue that this notion of consistency can be seen as an inductive bias that improves the efficiency of our learning. Equally, we posit that introducing such a bias towards spatio-temporal consistency into our models should greatly reduce the amount of supervision required for learning.

One way of achieving such inductive biases is through model structure. While recent successes in deep learning demonstrate that progress is possible without explicitly imbuing models with interpretable structure , recent works show that introducing such structure into deep models can indeed lead to favourable inductive biases improving performance e.g. in convolutional networks or in tasks requiring relational reasoning . Structure can also make neural networks useful in new contexts by significantly improving generalization, data efficiency or extending their capabilities to unstructured inputs .

Attend, Infer, Repeat (air), introduced by , is a notable example of such a structured probabilistic model that relies on deep learning and admits efficient amortized inference. Trained without any supervision, air is able to decompose a visual scene into its constituent components and to generate a (learned) number of latent variables that explicitly encode the location and appearance of each object. While this approach is inspiring, its focus on modelling individual (and thereby inherently static) scenes leads to a number of limitations. For example, it often merges two objects that are close together into one since no temporal context is available to distinguish between them. Similarly, we demonstrate that air struggles to identify partially occluded objects, e.g. when they extend beyond the boundaries of the scene frame (see Figure 7 in Section 4.1).

Our contribution is to mitigate the shortcomings of air by introducing a sequential version that models sequences of frames, enabling it to discover and track objects over time as well as to generate convincing extrapolations of frames into the future. We achieve this by leveraging temporal information to learn a richer, more capable generative model. Specifically, we extend air into a spatio-temporal state-space model and train it on unlabelled image sequences of dynamic objects. We show that the resulting model, which we name Sequential air (sqair), retains the strengths of the original AIR formulation while outperforming it on moving mnist digits.

The rest of this work is organised as follows. In Section 2, we describe the generative model and inference of air. In Section 3, we discuss its limitations and how it can be improved, thereby introducing Sequential Attend, Infer, Repeat (sqair), our extension of air to image sequences. In Section 4, we demonstrate the model on a dataset of multiple moving MNIST digits (Section 4.1) and compare it against air trained on each frame and Variational Recurrent Neural Network (vrnn) of with convolutional architectures, and show the superior performance of sqair in terms of log marginal likelihood and interpretability of latent variables. We also investigate the utility of inferred latent variables of sqair in downstream tasks. In Section 4.2 we apply sqair on real-world pedestrian CCTV data, where sqair learns to reliably detect, track and generate walking pedestrians without any supervision. Code for the implementation on the mnist datasetcode: github.com/akosiorek/sqair and the results videovideo: youtu.be/-IUNQgSLE0c are available online.

Attend, Infer, Repeat (AIR)

Generative Model The generative model of air is defined as follows

Sequential Attend-Infer-Repeat

While capable of decomposing a scene into objects, air only describes single images. Should we want a similar decomposition of an image sequence, it would be desirable to do so in a temporally consistent manner. For example, we might want to detect objects of the scene as well as infer dynamics and track identities of any persistent objects. Thus, we introduce Sequential Attend, Infer, Repeat (sqair), whereby air is augmented with a state-space model (ssm) to achieve temporal consistency in the generated images of the sequence. The resulting probabilistic model is composed of two parts: Discovery (disc), which is responsible for detecting (or introducing, in the case of the generation) new objects at every time-step (essentially equivalent to air), and Propagation (prop), responsible for updating (or forgetting) latent variables from the previous time-step given the new observation (image), effectively implementing the temporal ssm. We now formally introduce sqair by first describing its generative model and then the inference network.

The discovery prior pD(Dt,ztDt∣ztPt)p^{D}(D_{t},\mathbf{z}_{t}^{\mathcal{D}_{t}}|\mathbf{z}_{t}^{\mathcal{P}_{t}}) samples latent variables for new objects that enter the frame. The propagation prior pP(ztPt∣zt−1)p^{P}(\mathbf{z}_{t}^{\mathcal{P}_{t}}|\mathbf{z}_{t-1}) samples latent variables for objects that persist in the frame and removes latents of objects that disappear from the frame, thereby modelling dynamics and appearance changes. Both priors are learned during training. The exact forms of the priors are given in Appendix B.

Inference Similarly to air, inference in sqair can capture the number of objects and the representation describing the location and appearance of each object that is necessary to explain every image in a sequence. As with generation, inference is divided into prop and disc. During prop, the inference network achieves two tasks. Firstly, the latent variables from the previous time step are used to infer the current ones, modelling the change in location and appearance of the corresponding objects, thereby attaining temporal consistency. This is implemented by the temporal rnn {\color[rgb]{0,0,1}\definecolor[named]{pgfstrokecolor}{rgb}{0,0,1}\operatorname{R}_{\phi}^{T}}, with hidden states {\color[rgb]{0,0,1}\definecolor[named]{pgfstrokecolor}{rgb}{0,0,1}\bm{h}_{t}^{T}} (recurs in tt). Crucially, it does not access the current image directly, but uses the output of the relation rnn (cf. ). The relation rnn takes relations between objects into account, thereby implementing the explaining away phenomenon; it is essential for capturing any interactions between objects as well as occlusion (or overlap, if one object is occluded by another). See Figure 7 for an example. These two rnn s together decide whether to retain or to forget objects that have been propagated from the previous time step. During disc, the network infers further latent variables that are needed to describe any new objects that have entered the frame. All latent variables remaining after prop and disc are passed on to the next time step.

See Figures 2 and 3 for the inference network structure . The full variational posterior is defined as

Discovery, described by qϕDq^{D}_{\phi}, is very similar to the full posterior of air, cf. Equation 2. The only difference is the conditioning on ztPt\mathbf{z}_{t}^{\mathcal{P}_{t}}, which allows for a different number of discovered objects at each time-step and also for objects explained by prop not to be explained again. The second term, or qϕPq^{P}_{\phi}, describes propagation. The detailed structures of qϕDq^{D}_{\phi} and qϕPq^{P}_{\phi} are shown in Figure 3, while all the pertinent algorithms and equations can be found in Appendices A and C, respectively.

Learning We train sqair as an importance-weighted auto-encoder (iwae) of . Specifically, we maximise the importance-weighted evidence lower-bound L\textscIWAE\mathcal{L}_{\textsc{IWAE}}, namely

Experiments

We evaluate sqair on two datasets. Firstly, we perform an extensive evaluation on moving mnist digits, where we show that it can learn to reliably detect, track and generate moving digits (Section 4.1). Moreover, we show that sqair can simulate moving objects into the future — an outcome it has not been trained for. We also study the utility of learned representations for a downstream task. Secondly, we apply sqair to real-world pedestrian CCTV data from static cameras (DukeMTMC, ), where we perform background subtraction as pre-processing. In this experiment, we show that sqair learns to detect, track, predict and generate walking pedestrians without human supervision.

The dataset consists of sequences of length 10 of multiple moving mnist digits. All images are of size 50×5050\times 50 and there are zero, one or two digits in every frame (with equal probability). Sequences are generated such that no objects overlap in the first frame, and all objects are present through the sequence; the digits can move out of the frame, but always come back. See Appendix F for an experiment on a harder version of this dataset. There are 60,000 training and 10,000 testing sequences created from the respective mnist datasets. We train two variants of sqair: the mlp-sqair uses only fully-connected networks, while the conv-sqair replaces the networks used to encode images and glimpses with convolutional ones; it also uses a subpixel-convolution network as the glimpse decoder . See Appendix D for details of the model architectures and the training procedure.

We use air and vrnn as baselines for comparison. vrnn can be thought of as a sequential vae with an rnn as its deterministic backbone. Being similar to a vae, its latent variables are not structured, nor easily interpretable. For a fair comparison, we control the latent dimensionality of vrnn and the number of learnable parameters. We provide implementation details in Section D.3.

2 Generative Modelling of Walking Pedestrians

To evaluate the model in a more challenging, real-world setting, we turn to data from static CCTV cameras of the DukeMTMC dataset . As part of pre-precessing, we use standard background subtraction algorithms . In this experiment, we use 31503150 training and 350350 validation sequences of length 55. For details of model architectures, training and data pre-processing, see Appendix E. We evaluate the model qualitatively by examining reconstructions, conditional samples (conditioned on the first four frames) and samples from the prior (Figure 8 and Appendix I). We see that the model learns to reliably detect and track walking pedestrians, even when they are close to each other.

There are some spurious detections and re-detections of the same objects, which is mostly caused by imperfections of the background subtraction pipeline — backgrounds are often noisy and there are sudden appearance changes when a part of a person is treated as background in the pre-processing pipeline. The object counting accuracy in this experiment is 0.57120.5712 on the validation dataset, and we noticed that it does increase with the size of the training set. We also had to use early stopping to prevent overfitting, and the model was trained for only 315315k iterations (>1>1M for mnist experiments). Hence, we conjecture that accuracy and marginal likelihood can be further improved by using a bigger dataset.

Related Work

There have been many approaches to modelling objects in images and videos. Object detection and tracking are typically learned in a supervised manner, where object bounding boxes and often additional labels are part of the training data. Single-object tracking commonly use Siamese networks, which can be seen as an rnn unrolled over two time-steps . Recently, used an rnn with an attention mechanism in the hart model to predict bounding boxes for single objects, while robustly modelling their motion and appearance. Multi-object tracking is typically attained by detecting objects and performing data association on bounding-boxes . used an end-to-end supervised approach that detects objects and performs data association. In the unsupervised setting, where the training data consists of only images or videos, the dominant approach is to distill the inductive bias of spatial consistency into a discriminative model. detect single objects and their parts in images, and incorporate temporal consistency to better track single objects. Sqair is unsupervised and hence it does not rely on bounding boxes nor additional labels for training, while being able to learn arbitrary motion and appearance models similarly to hart . At the same time, is inherently multi-object and performs data association implicitly (cf. Appendix A). Unlike the other unsupervised approaches, temporal consistency is baked into the model structure of sqair and further enforced by lower kl divergence when an object is tracked.

Many works on video prediction learn a deterministic model conditioned on the current frame to predict the future ones . Since these models do not model uncertainty in the prediction, they can suffer from the multiple futures problem — since perfect prediction is impossible, the model produces blurry predictions which are a mean of possible outcomes. This is addressed in stochastic latent variable models trained using variational inference to generate multiple plausible videos given a sequence of images . Unlike sqair, these approaches do not model objects or their positions explicitly, thus the representations they learn are of limited interpretability.

Learning decomposed representations of object appearance and position lies at the heart of our model. This problem can be also seen as perceptual grouping, which involves modelling pixels as spatial mixtures of entities. and learn to decompose images into separate entities by iterative refinement of spatial clusters using either learned updates or the Expectation Maximization algorithm; and extend these approaches to videos, achieving very similar results to sqair. Perhaps the most similar work to ours is the concurrently developed model of . The above approaches rely on iterative inference procedures, but do not exhibit the object-counting behaviour of sqair. For this reason, their computational complexities are proportional to the predefined maximum number of objects, while sqair can be more computationally efficient by adapting to the number of objects currently present in an image.

Another interesting line of work is the gan-based unsupervised video generation that decomposes motion and content . These methods learn interpretable features of content and motion, but deal only with single objects and do not explicitly model their locations. Nonetheless, adversarial approaches to learning structured probabilistic models of objects offer a plausible alternative direction of research.

To the best of our knowledge, is the only known approach that models pixels belonging to a variable number of objects in a video together with their locations in the generative sense. This work uses a Bayesian nonparametric (BNP) model, which relies on mixtures of Dirichlet processes to cluster pixels belonging to an object. However, the choice of the model necessitates complex inference algorithms involving Gibbs sampling and Sequential Monte Carlo, to the extent that any sensible approximation of the marginal likelihood is infeasible. It also uses a fixed likelihood function, while ours is learnable.

The object appearance-persistence-disappearance model in sqair is reminiscent of the Markov Indian buffet process (MIBP) of , another BNP model. MIBP was used as a model for blind source separation, where multiple sources contribute toward an audio signal, and can appear, persist, disappear and reappear independently. The prior in sqair is similar, but the crucial differences are that sqair combines the BNP prior with flexible neural network models for the dynamics and likelihood, as well as variational learning via amortized inference. The interface between deep learning and BNP, and graphical models in general, remains a fertile area of research.

Discussion

In this paper we proposed sqair, a probabilistic model that extends air to image sequences, and thereby achieves temporally consistent reconstructions and samples. In doing so, we enhanced air’s capability of disentangling overlapping objects and identifying partially observed objects.

This work continues the thread of , and, together with , presents unsupervised object detection & tracking with learnable likelihoods by the means of generative modelling of objects. In particular, our work is the first one to explicitly model object presence, appearance and location through time. Being a generative model, sqair can be used for conditional generation, where it can extrapolate sequences into the future. As such, it would be interesting to use it in a reinforcement learning setting in conjunction with Imagination-Augmented Agents or more generally as a world model , especially for settings with simple backgrounds, e. g., games like Montezuma’s Revenge or Pacman.

The framework offers various avenues of further research; Sqair leads to interpretable representations, but the interpretability of what variables can be further enhanced by using alternative objectives that disentangle factors of variation in the objects . Moreover, in its current state, sqair can work only with simple backgrounds and static cameras. In future work, we would like to address this shortcoming, as well as speed up the sequential inference process whose complexity is linear in the number of objects. The generative model, which currently assumes additive image composition, can be further improved by e. g., autoregressive modelling . It can lead to higher fidelity of the model and improved handling of occluded objects. Finally, the sqair model is very complex, and it would be useful to perform a series of ablation studies to further investigate the roles of different components.

Acknowledgements

We would like to thank Ali Eslami for his help in implementing air, Alex Bewley and Martin Engelcke for discussions and valuable insights and anonymous reviewers for their constructive feedback. Additionally, we acknowledge that HK and YWT’s research leading to these results has received funding from the European Research Council under the European Union’s Seventh Framework Programme (FP7/2007-2013) ERC grant agreement no. 617071.

References

Appendix A Algorithms

Image generation, described by Algorithm 1, is exactly the same for sqair and air. Algorithms 2 and 3 describe inference in sqair. Note that disc is equivalent to air if no latent variables are present in the inputs.

To condition disc on propagated latent variables (Algorithm 3 in Algorithm 3), we encode the latter by using a two-layer mlp similarly to ,

Note that other encoding schemes are possible, though we have experimented only with this one.

Appendix B Details for the Generative Model of SQAIR

In implementation, we upper bound the number of objects at any given time by NN. In detail, the discovery prior is given by

Appendix C Details for the Inference of SQAIR

The propagation inference network qϕPq^{P}_{\phi} is given as below,

with {\color[rgb]{1,0.6015625,0}\definecolor[named]{pgfstrokecolor}{rgb}{1,0.6015625,0}\bm{h}_{t}^{R,i}} the hidden state of the relation rnn (see Equation 14). Its role is to capture information from the observation xt\mathbf{x}_{t} as well as to model dependencies between different objects. The propagation posterior for a single object can be expanded as follows,

where the second term is the delta distribution centered on the presence of this object at the previous time-step. If it was not there, it cannot be propagated. Let j∈{0,…,i−1}j\in\{0,\dots,i-1\} be the index of the most recent present object before object ii. Hidden states are updated as follows,

Appendix D Details of the moving-mnist Experiments

All models are trained by maximising the evidence lower bound (elbo) LIWAE\mathcal{L}_{IWAE} (Equation 5) with the rmsprop optimizer with momentum equal to 0.90.9. We use the learning rate of 10−510^{-5} and decrease it to 13⋅10−5\frac{1}{3}\cdot 10^{-5} after 400k and to 10−610^{-6} after 1000k training iterations. Models are trained for the maximum of 2⋅1062\cdot 10^{6} training iterations; we apply early stopping in case of overfitting. Sqair models are trained with a curriculum of sequences of increasing length: we start with three time-steps, and increase by one time-step every 10510^{5} training steps until reaching the maximum length of 10. When training air, we treated all time-steps of a sequence as independent, and we trained it on all data (sequences of length ten, split into ten independent sequences of length one).

D.2 Sqair and air Model Architectures

All models use glimpse size of 20×2020\times 20 and exponential linear unit (elu) non-linearities for all layers except RNNs and output layers. mlp-sqair uses fully-connected layers for all networks. In both variants of sqair, the {\color[rgb]{1,0,1}\definecolor[named]{pgfstrokecolor}{rgb}{1,0,1}\operatorname{R}_{\phi}^{D}} and {\color[rgb]{1,0.6015625,0}\definecolor[named]{pgfstrokecolor}{rgb}{1,0.6015625,0}\operatorname{R}_{\phi}^{R}} RNNs are the vanilla RNNs. The propagation prior rnn and the temporal rnn {\color[rgb]{0,0,1}\definecolor[named]{pgfstrokecolor}{rgb}{0,0,1}\operatorname{R}_{\phi}^{T}} use gated recurrent unit (gru). air follows the same architecture as mlp-sqair. All fully-connected layers and RNNs in mlp-sqair and air have 256 units; they have 2.92.9M and 1.71.7M trainable parameters, respectively.

Conv-sqair differs from the mlp version in that it uses cnns for the glimpse and image encoders and a subpixel-cnn for the glimpse decoder. All fully connected layers and RNNs have 128 units. The encoders share the cnn, which is followed by a single fully-connected layer (different for each encoder). The cnn has four convolutional layers with featuresmapsandstridesoffeatures maps and strides of. The glimpse decoder is composed of two fully-connected layers with $hiddenunits,whoseoutputsarereshapedintohidden units, whose outputs are reshaped into32featuresmapsofsizefeatures maps of size5\times 5,followedbyasubpixel−cnnwiththreelayersof<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><mi>f</mi><mi>e</mi><mi>a</mi><mi>t</mi><mi>u</mi><mi>r</mi><mi>e</mi><mi>m</mi><mi>a</mi><mi>p</mi><mi>s</mi><mi>a</mi><mi>n</mi><mi>d</mi><mi>s</mi><mi>t</mi><mi>r</mi><mi>i</mi><mi>d</mi><mi>e</mi><mi>s</mi><mi>o</mi><mi>f</mi></mrow><annotationencoding="application/x−tex">featuremapsandstridesof</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.8889em;vertical−align:−0.1944em;"></span><spanclass="mordmathnormal"style="margin−right:0.1076em;">f</span><spanclass="mordmathnormal">e</span><spanclass="mordmathnormal">a</span><spanclass="mordmathnormal">t</span><spanclass="mordmathnormal">u</span><spanclass="mordmathnormal"style="margin−right:0.0278em;">r</span><spanclass="mordmathnormal">e</span><spanclass="mordmathnormal">ma</span><spanclass="mordmathnormal">p</span><spanclass="mordmathnormal">s</span><spanclass="mordmathnormal">an</span><spanclass="mordmathnormal">d</span><spanclass="mordmathnormal">s</span><spanclass="mordmathnormal">t</span><spanclass="mordmathnormal"style="margin−right:0.0278em;">r</span><spanclass="mordmathnormal">i</span><spanclass="mordmathnormal">d</span><spanclass="mordmathnormal">eso</span><spanclass="mordmathnormal"style="margin−right:0.1076em;">f</span></span></span></span></span>.Allfiltersareofsize, followed by a subpixel-cnn with three layers of <span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mi>f</mi><mi>e</mi><mi>a</mi><mi>t</mi><mi>u</mi><mi>r</mi><mi>e</mi><mi>m</mi><mi>a</mi><mi>p</mi><mi>s</mi><mi>a</mi><mi>n</mi><mi>d</mi><mi>s</mi><mi>t</mi><mi>r</mi><mi>i</mi><mi>d</mi><mi>e</mi><mi>s</mi><mi>o</mi><mi>f</mi></mrow><annotation encoding="application/x-tex">feature maps and strides of</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.8889em;vertical-align:-0.1944em;"></span><span class="mord mathnormal" style="margin-right:0.1076em;">f</span><span class="mord mathnormal">e</span><span class="mord mathnormal">a</span><span class="mord mathnormal">t</span><span class="mord mathnormal">u</span><span class="mord mathnormal" style="margin-right:0.0278em;">r</span><span class="mord mathnormal">e</span><span class="mord mathnormal">ma</span><span class="mord mathnormal">p</span><span class="mord mathnormal">s</span><span class="mord mathnormal">an</span><span class="mord mathnormal">d</span><span class="mord mathnormal">s</span><span class="mord mathnormal">t</span><span class="mord mathnormal" style="margin-right:0.0278em;">r</span><span class="mord mathnormal">i</span><span class="mord mathnormal">d</span><span class="mord mathnormal">eso</span><span class="mord mathnormal" style="margin-right:0.1076em;">f</span></span></span></span></span>. All filters are of size3\times 3.Conv−sqairhas. Conv-sqair has2.6$M trainable parameters.

We have experimented with different sizes of fully-connected layers and RNNs; we kept the size of all layers the same and altered it in increments of 32 units. Values greater than 256 for mlp-sqair and 128 for conv-sqair resulted in overfitting. Models with as few as 32 units per layer (<0.9<0.9M trainable parameters for mlp-sqair) displayed the same qualitative behaviour as reported models, but showed lower quantitative performance.

D.3 Vrnn Implementation and Training Details

For each of mlp-vrnn and conv-vrnn, we experimented with three architectures: small/medium/large. We used HH=H′H^{\prime}=JJ=128/256/512 and LL=L′L^{\prime}=2/3/4 for mlp-vrnn, giving number of parameters of 1.21.2M/2.12.1M/9.89.8M. For conv-vrnn, the number of features maps we used was ,, and ,withstridesof, with strides of, andand, all with 3×33\times 3 filters, HH=JJ=128128/256256/512512 and LL=1, giving number of parameters of 0.80.8M/2.62.6M/6.16.1M. The largest convolutional encoder architecture is very similar to that in applied to mnist.

We have chosen the medium-sized models for comparison with sqair due to overfitting encountered in larger models.

D.4 Addition Experiment

Appendix E Details of the DukeMTMC Experiments

We take videos from cameras one, two, five, six and eight from the DukeMTMC dataset . As pre-processing, we invert colors and subtract backgrounds using standard OpenCV tools , downsample to the resolution of 240×175240\times 175, convert to gray-scale and randomly crop fragments of size 64×6464\times 64. Finally, we generate 35003500 sequences of length five such that the maximum number of objects present in any single frame is three and we split them into training and validation sets with the ratio of 9:19:1.

We use the same training procedure as for the mnist experiments. The only exception is the learning curriculum, which goes from three to five time-steps, since this is the maximum length of the sequences.

The reported model is similar to conv-sqair. We set the glimpse size to 28×1228\times 12 to account for the expected aspect ratio of pedestrians. Glimpse and image encoders share a cnn with featuremapsandstridesoffeature maps and strides of followed by a fully-connected layer (different for each encoder). The glimpse decoder is implemented as a two-layer fully-connected network with 128 and 1344 units, whose outputs are reshaped into 64 feature maps of size 7×37\times 3, followed by a subpixel-cnn with two layers of featuremapsandstridesoffeature maps and strides of. All remaining fully-connected layers in the model have 128 units. The total number of trainable parameters is 3.53.5M.

Appendix F Harder multi-mnist Experiment

We created a version of the multi-mnist dataset, where objects can appear or disappear at an arbitrary point in time. It differs from the dataset described in Section 4.1, where all digits are present throughout the sequence. All other dataset parameters are the same as in Section 4.1. Figure 9 shows an example sequence and mlp-sqair reconstructions with marked glimpse locations. The model has no trouble detecting new digits in the middle of the sequence and rediscovering a digit that was previously present.

Appendix G Failure cases of sqair

Appendix H Reconstruction and Samples from the Moving-MNIST Dataset

H.2 Samples

H.3 Conditional Generation

Appendix I Reconstruction and Samples from the DukeMTMC Dataset