Skip to content

OrthogonalProcrustesAlignment computes the full pairwise distance matrix before subsampling, leading to excessive memory usage #307

Description

@onford

Hi CEBRA team,

First, thank you for making CEBRA open source! While reproducing the multi-animal hippocampus example, I encountered a memory issue in OrthogonalProcrustesAlignment.

After inspecting the implementation, I noticed that the current workflow is:

  1. compute the full pairwise label distance matrix: distance = self._distance(label, ref_label);
  2. sort every row: target_idx = np.argsort(distance, axis=1)[:, :self.top_k]
  3. construct X and Y
  4. apply subsample (if provided)
  5. solve the Orthogonal Procrustes problem

This means that subsample only reduces the size of the final Procrustes problem, but does not reduce the memory required for the pairwise distance computation.

For example, aligning datasets of sizes ref_label(10178, 3) and label(47431, 3) creates a distance matrix of shape (47431, 10178), which already occupies several GB of memory before subsampling is applied. On my machine(running code on WSL), the process is terminated by the Linux OOM killer before reaching the subsampling step.

From the documentation, I initially expected subsample to reduce the overall memory footprint, but in the current implementation it only affects the final Procrustes optimization.

Is this intended behavior?

If so, it might be worth mentioning in the documentation that subsample does not reduce the memory required for nearest-neighbor matching.

Thanks!

Activity

  1. stes commented on Jul 17, 2026

    @stes
    Member

    Dear @onford , thanks a lot for using CEBRA. Could you post the full code snippet you are using here to help me reproducing the issue and discussing what we can do here? Thanks!

  2. onford commented on Jul 18, 2026

    @onford
    Author

    Dear @onford , thanks a lot for using CEBRA. Could you post the full code snippet you are using here to help me reproducing the issue and discussing what we can do here? Thanks!

    Hello! I followed the doc Technical: Training models across animals and wrote a python file with cebra==0.6.1.

    import cebra.data
    import cebra.datasets
    from cebra import CEBRA
    import pickle
    import os
    
    from pathlib import Path
    
    PROJECT_ROOT = Path(__file__).parent.parent
    MODEL_ROOT = PROJECT_ROOT / 'models' / 'rat_hippocampus'
    FIGURE_ROOT = PROJECT_ROOT / 'figures' / 'exp6'
    PICKLE_ROOT = PROJECT_ROOT / 'pickles' / 'exp6'
    os.makedirs(MODEL_ROOT, exist_ok=True)
    os.makedirs(FIGURE_ROOT, exist_ok=True)
    os.makedirs(PICKLE_ROOT, exist_ok=True)
    
    hippocampus_a = cebra.datasets.init('rat-hippocampus-single-achilles')
    hippocampus_b = cebra.datasets.init('rat-hippocampus-single-buddy')
    hippocampus_c = cebra.datasets.init('rat-hippocampus-single-cicero')
    hippocampus_g = cebra.datasets.init('rat-hippocampus-single-gatsby')
    
    names = ["achilles", "buddy", "cicero", "gatsby"]
    datas = [hippocampus_a.neural.numpy(), hippocampus_b.neural.numpy(), hippocampus_c.neural.numpy(), hippocampus_g.neural.numpy()]
    labels = [hippocampus_a.continuous_index.numpy(), hippocampus_b.continuous_index.numpy(), hippocampus_c.continuous_index.numpy(), hippocampus_g.continuous_index.numpy()]
    
    max_iterations = 10 # for a quick debug
    
    embeddings = dict()
    
    # Single session training
    for name, X, y in zip(names, datas, labels):
        # Fit one CEBRA model per session (i.e., per rat)
        print(f"Fitting CEBRA for {name}")
        cebra_model = CEBRA(model_architecture='offset10-model',
                            batch_size=512,
                            learning_rate=3e-4,
                            temperature=1,
                            output_dimension=3,
                            max_iterations=max_iterations,
                            distance='cosine',
                            conditional='time_delta',
                            device='cuda_if_available',
                            verbose=True,
                            time_offsets=10)
    
        cebra_model.fit(X, y)
        embeddings[name] = cebra_model.transform(X)
    
    # Align the single session embeddings to the first rat
    alignment = cebra.data.helper.OrthogonalProcrustesAlignment()
    first_rat = list(embeddings.keys())[0]
    
    for j, rat_name in enumerate(list(embeddings.keys())[1:]):
        embeddings[f"{rat_name}"] = alignment.fit_transform(
            embeddings[first_rat], embeddings[rat_name], labels[0], labels[j+1])
    
    # Save embeddings in current folder
    with open(PICKLE_ROOT / 'embeddings.pkl', 'wb') as f:
        pickle.dump(embeddings, f)
    
    multi_embeddings = dict()
    
    # Multisession training
    multi_cebra_model = CEBRA(model_architecture='offset10-model',
                        batch_size=512,
                        learning_rate=3e-4,
                        temperature=1,
                        output_dimension=3,
                        max_iterations=max_iterations,
                        distance='cosine',
                        conditional='time_delta',
                        device='cuda_if_available',
                        verbose=True,
                        time_offsets=10)
    
    # Provide a list of data, i.e. datas = [data_a, data_b, ...]
    multi_cebra_model.fit(datas, labels)
    
    # Transform each session with the right model, by providing the corresponding session ID
    for i, (name, X) in enumerate(zip(names, datas)):
        multi_embeddings[name] = multi_cebra_model.transform(X, session_id=i)
    
    # Save embeddings in current folder
    with open(PICKLE_ROOT / 'multi_embeddings.pkl', 'wb') as f:
        pickle.dump(multi_embeddings, f)

    And the execution result was:

    Fitting CEBRA for achilles
    pos: -0.8009 neg:  7.0235 total:  6.2226 temperature:  1.0000: 100%|█████████████████████████████████████████████████████████████| 10/10 [00:01<00:00,  8.01it/s]
    Fitting CEBRA for buddy
    pos: -0.9098 neg:  7.1451 total:  6.2352 temperature:  1.0000: 100%|█████████████████████████████████████████████████████████████| 10/10 [00:00<00:00, 10.05it/s]
    Fitting CEBRA for cicero
    pos: -0.9430 neg:  7.1863 total:  6.2432 temperature:  1.0000: 100%|█████████████████████████████████████████████████████████████| 10/10 [00:02<00:00,  4.42it/s]
    Fitting CEBRA for gatsby
    pos: -0.9003 neg:  7.1311 total:  6.2308 temperature:  1.0000: 100%|█████████████████████████████████████████████████████████████| 10/10 [00:01<00:00,  7.85it/s]
    Killed                     python /home/ruiwang18/research/neuroscience/cebra-reproduction/experiments/6-1-multi_animal_joint_training.py
    

    I checked the log by dmesg and found it was a memory issue.

    [265564.096738] [   8752]  1000  8752     1567       38        0       38         0    61440      416             0 bash
    [265564.097258] [   8832]  1000  8832  1079503   444107   444018       89         0  6459392   340960             0 python
    [265564.097833] [   9184]  1000  9184  2068620  1287428  1287370       58         0 11538432     8480             0 python
    [265564.098561] oom-kill:constraint=CONSTRAINT_NONE,nodemask=(null),cpuset=/,mems_allowed=0,global_oom,task_memcg=/,task=python,pid=9184,uid=1000
    [265564.099365] Out of memory: Killed process 9184 (python) total-vm:8274480kB, anon-rss:5149480kB, file-rss:232kB, shmem-rss:0kB, UID:1000 pgtables:11268kB oom_score_adj:0
    
  3. onford commented on Jul 18, 2026

    @onford
    Author

    I thought the parameter subsample of OrthogonalProcrustesAlignment could help reducing memory cost, but it turned out not, and that's why I came up with this issue. I try the code today and find it runs normally. By using /usr/bin/time -v python ... I found the memory cost being 6,274,080 kbytes, which is maybe a little expensive for WSL users.

    Also I notice that OrthogonalProcrustesAlignment._distance creates a $n_1\times n_2$ array, with the scale of label and reference label denoted as $n_1$ and $n_2$, which hinders CEBRA from scaling to larger datasets.

  4. AmitSubhash commented on Jul 19, 2026

    @AmitSubhash

    Confirmed. subsample only shrinks the final solve: it's applied after the full (n_label × n_ref) distance matrix and its argsort are already built, so it can't reduce peak memory. At your sizes the matrix is ~3.9 GB (float64), and peak is roughly 3× that from the einsum temporaries plus the argsort index array.

    Fix: compute top_k in row chunks so the full matrix is never materialized. Since argsort(axis=1) is per-row independent, this is bit-identical to the current output, keeps compute the same, needs no new dependency, and just adds a chunk_size parameter. Peak drops from O(n_label × n_ref) to O(chunk_size × n_ref). I verified it reproduces the current result exactly on the cicero→achilles hippocampus alignment, for both continuous and discrete behavior labels (identical aligned embeddings).

    Happy to open a PR with a regression test. Would you like chunk_size exposed on the constructor? (Two unrelated nits I can fold in or skip: np.random.choice for subsampling draws with replacement, and transform() guards self.transform instead of self._transform.)

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions