Writing

Letting the atlas be unbalanced

A manifold-shaped representation only started to pay off at scale once we stopped forcing every chart to be used equally.

The manifold hypothesis is one of those ideas that almost everyone in machine learning accepts and few methods use directly. An Atari frame has tens of thousands of pixels, but what changes from one frame to the next is a short list of factors: where the player is, where the ball is, what the score says. If the data really lies near a low-dimensional manifold, the natural target for a learned representation is that manifold. Most self-supervised methods still map everything into one flat embedding space and leave the downstream task to sort out the geometry.

Differential geometry has a precise tool for describing curved spaces, the atlas. A manifold that cannot be flattened in one piece can still be covered by several overlapping charts, each a local coordinate system that is flat enough to work with. MSimCLR (Korman, 2021) brought this idea into self-supervised learning. Its encoder produces several chart embeddings together with a membership probability that says which chart an input belongs to. The construction is elegant, but it has a practical catch. It only outperforms plain SimCLR when the target encoding dimension is kept extremely small, and as soon as the model is given a realistic amount of capacity the advantage disappears.

Our ICLR 2024 paper started from the question of why. Our reading was that the problem lies in the prior. MSimCLR pulls the membership distribution toward uniform, so that inputs are spread evenly over the charts. That sounds like a sensible way to make every chart useful, but it adds uncertainty to every prediction and it tends to produce heads that look alike. We had seen the same failure in reinforcement learning, where the heads of a bootstrapped ensemble drift toward one another unless something keeps them apart. A uniform prior over charts encourages exactly that kind of redundancy.

The unbalanced atlas removes it. Instead of pulling the membership distribution toward uniform, we add a small maximum mean discrepancy term that pushes it away from uniform, so each input is encouraged to commit to one chart. During training the output is the membership-weighted sum of the chart outputs. At inference the model takes the most probable chart and uses that head alone. The effect shows up directly in the entropy of the membership output, which is much lower than under a uniform prior. Some charts end up covering a lot of the data and others very little, which is where the name comes from.

Two smaller choices turned out to matter as much as the prior. Each chart’s coordinate map is the identity rather than a learned linear layer, so the chart outputs are the encoder’s own features. And when a prediction is scored against its target, we use the average of the chart outputs rather than a single chart. We call these dilated targets. Geometrically the average corresponds to a Minkowski sum of the charts, and under a convexity assumption it guarantees that points lying where charts overlap are not lost from the representation. Nothing forces a trained network to satisfy that assumption, and the paper says so, but the ablations suggest the dilated targets are doing real work.

We built the method, DIM-UA, on Spatiotemporal DeepInfomax (ST-DIM), which learns state representations from consecutive Atari frames with two InfoNCE objectives. The global-local term asks the representation of the frame at time t to pick out local feature patches of the frame at t+1, and the local-local term matches patches with patches at the same position. DIM-UA keeps both terms, changes the global-local score so that it uses the averaged chart outputs, and adds the membership term with a weight of 0.1. Everything else in the training setup is unchanged, which keeps the comparison with ST-DIM clean.

The evaluation uses AtariARI, a benchmark built on the Atari 2600 that reads the game’s RAM, so every frame comes with ground-truth state variables. They fall into five groups: the agent’s position, the positions of small objects such as a ball, the positions of other objects, the score, clock and lives counters, and a miscellaneous group. Frames are collected with a random policy. The encoder is trained without labels and then frozen, and a linear probe is fitted for each variable. If the probe can read a variable off the representation, the representation has kept it.

The headline comparison is at 16,384 hidden units. DIM-UA, with four heads of 4,096 units each, reaches a mean F1 score of 0.75 averaged over categories. ST-DIM in its original configuration reaches 0.72. The same ST-DIM widened to 16,384 units in a single head does worse, at 0.70, and varies more between runs. Probe accuracy follows the same order, 0.76 for DIM-UA against 0.73 and 0.71. Across the 19 games in the study, DIM-UA was at least as good as both versions of ST-DIM in every game. Freeway shows the difference most sharply. The widened ST-DIM sometimes collapses there, averaging an F1 of 0.30 with a standard deviation larger than the mean, while DIM-UA reaches 0.86.

The way the two methods scale says more than the headline. At small sizes plain ST-DIM is better, because the structure of an atlas costs something when each head has only a few dimensions to work with. Around 2,048 units the curves meet. Beyond that point ST-DIM gets worse as it gets wider, while DIM-UA keeps improving. The number of heads interacts with this. Two heads work best below 2,048 units and worst at 16,384. Eight heads are the weakest at small sizes but show no sign of levelling off at the largest size we tried. The simplest account is that each chart needs enough dimensions to be useful, and once it has them, more charts let the model describe more of the manifold.

The ablations separate the ingredients. Restoring the uniform prior, which is essentially the MSimCLR recipe, gives the worst results at every size. Removing the dilated targets helps at 512 units but falls behind as the representation grows. The two changes are therefore not interchangeable. Unbalanced membership keeps the heads from becoming copies of each other, and dilated targets keep the regions where charts overlap inside the representation.

We also tried the idea outside reinforcement learning, with SimCLR on CIFAR-10, a ResNet-50 and 1,000 epochs of training. There the gain is small and should not be oversold. SimCLR reaches 88.3% accuracy with 512 units, the unbalanced version 88.6% with eight heads of 512, and the best MSimCLR configuration 87.8%. The purpose of that experiment was mainly to check that the paradigm carries over to a different contrastive objective without breaking it.

The limitations are real. The convexity assumption behind the dilated targets is not guaranteed. Training several heads is slower, and diversity between them takes more epochs to appear. And 16,384 units is a lot of capacity for a small encoder. It is where the method does best, but not necessarily what anyone wants to deploy. The most interesting open question is how the number of charts should relate to the intrinsic dimension of the data. If that relationship could be estimated, the atlas could be sized from the data rather than found by grid search.

The idea also reaches beyond Atari. Neural recordings, medical images and multimodal clinical records all contain regions of very different local complexity, across subjects, channels and conditions. A representation that can spend its capacity unevenly, instead of forcing every region into the same template, looks well suited to data like that.

The paper is on OpenReview and arXiv, and the code for DIM-UA is on GitHub.