SocialNet#

SocialNet, presented in [1], is a deep learning model of collective behavior that learns the rules governing how animals influence each other. It is organized into two interpretable modules: a pair-interaction subnetwork that maps how one individual affects another, and an aggregation subnetwork that describes how each individual weighs and combines the influences of all its neighbors.

SocialNet is trained on trajectory files from idtracker.ai, or from other sources in one of the supported formats.

SocialNet diagram

The SocialNet architecture (extracted from Figure 1 in [1]). (A) Variables used to predict future turns. Asocial variables, those only involving the focal, in red. Social variables, those involving both the focal and a neighbour, in blue. (B) Pair-interaction subnetwork of SocialNet, receiving asocial variables \(\alpha\) and social variables \(\sigma_i\) from a single neighbour \(i\), and outputting a single scalar value. All pair-interaction networks share the same weights. (C) Aggregation subnetwork of SocialNet. Same structure as B, but the input is a restricted symmetric subset of the variables and the output is passed through an exponential function to make it positive. (D) General SocialNet architecture, showing how the inputs of the pair-interaction and aggregation subnetworks are integrated to produce a single logit \(z\) for the focal fish turning right after 1 s.#

https://gitlab.com/polavieja_lab/socialnet/-/raw/master/examples/interaction_subnetwork.png https://gitlab.com/polavieja_lab/socialnet/-/raw/master/examples/aggregation_subnetwork_vars_fv_nbv_nbx_nby.png

Left: Example interaction map showing how SocialNet predicts the influence of a neighbor on the focal animal’s probability of turning right after 1 second (obtained with plot.plot_interaction_subnetwork()).

Right: Example aggregation map showing how SocialNet weighs different neighbor variables when aggregating social information (obtained with plot.plot_aggregation_subnetwork()).

Install SocialNet#

Warning

SocialNet is not included in idtracker.ai and must be installed separately.

SocialNet runs on PyTorch, like idtracker.ai, so you can install it in the environment where you installed idtracker.ai (see Installation). Activate that environment and install SocialNet from our repository:

conda activate idtrackerai
pip install git+https://gitlab.com/polavieja_lab/socialnet
Install SocialNet in its own environment

To keep SocialNet apart from idtracker.ai, create a new environment with Python 3.10 or later:

conda create -n socialnet python=3.13
conda activate socialnet

Install PyTorch with the command for your computer from the PyTorch website, then install SocialNet:

pip install git+https://gitlab.com/polavieja_lab/socialnet

Basic usage of SocialNet#

Train a model with model_train() and test it with model_test():

model_train

Train SocialNet model from a set of trajectory files.

model_test

Test the model on the given trajectory files.

Then analyze the trained model with these plotting functions:

plot.plot_aggregation_subnetwork

Visualizes the aggregation subnetwork of the model by plotting how weights vary as a function of the specified variables.

plot.plot_interaction_subnetwork

Plots the probability (logit) of the focal animal turning right resulting from the pair-interaction subnetwork of the attention network, as a function of the orientation of the neighbour with respect to the focal (θi) and the speed of the neighbour (vi).

plot.plot_interaction_scores

Plots processed interaction scores (attraction, alignment, and repulsion) as a function of kinematic variables of the focal and neighbour animals.

plot.plot_product

Plots the combination (product) of plot.plot_interaction_scores() and plot.plot_aggregation_subnetwork(), visualizing how the interaction and attention subnetworks jointly affect the predicted interaction scores.

The example below downloads sample trajectory files from our data repository, then trains, tests and plots a model.

Example of training, testing and plotting SocialNet Open in Colab #
from pathlib import Path
import gdown  # pip install gdown
from socialnet import model_test, model_train
from socialnet.plot import (
    plot_aggregation_subnetwork,
    plot_interaction_scores,
    plot_product,
    plot_interaction_subnetwork,
)

# data from https://drive.google.com/drive/folders/1VH97_bNFz09Ke_kBL1oV2HbTS25BwnKC

gdown.download(
    id="1y1ZhNr3eWbhYwA_ZPfIs9UjsinszKCwp", output="zebrafish_60_1.npy", resume=True
)
gdown.download(
    id="1aJb2pgzJE8dhDkWVGKWdBgYgFRhpWkj7", output="zebrafish_60_2.npy", resume=True
)
gdown.download(
    id="1LaNIqFGD5N9SUIq0yfIGEZUnrzF3TgLy", output="zebrafish_60_3.npy", resume=True
)

results_dict = model_train(
    ["zebrafish_60_1.npy", "zebrafish_60_2.npy", "zebrafish_60_3.npy"],
    session_name="example",
)

expected_output_folder = Path.cwd() / "socialnet_session_example"


test_results = model_test(
    trajectory_files=["zebrafish_60_1.npy", "zebrafish_60_2.npy", "zebrafish_60_3.npy"],
    model_folder=expected_output_folder,
)

# plot slices of the neighbor x and y coordinates for different values of the focal velocity and the neighbor velocity
# fv = focal velocity
# nba = neighbour acceleration
# nbv = neighbour velocity
# nbx = neighbour x position
# nby = neighbour y position
fig_vars = ("fv", "nbv", "nbx", "nby")  # row_var, col_var, x_var, y_var

plot_aggregation_subnetwork(expected_output_folder, fig_vars=fig_vars)
plot_interaction_subnetwork(expected_output_folder)
plot_interaction_scores(expected_output_folder, fig_vars=fig_vars)
plot_product(expected_output_folder, fig_vars=fig_vars)
Alternative SocialNet CLI

SocialNet can also be used from the command line, without writing Python code:

Main CLI commands#
socialnet train --help
socialnet test --help
socialnet sample --help
Plotting CLI commands#
socialnet plot --help
socialnet plot aggregation_subnetwork --help
socialnet plot interaction_subnetwork --help
socialnet plot interaction_scores --help
socialnet plot product --help

Trajectory formats#

model_train() and model_test() (and the train, test and sample commands) accept:

  • idtracker.ai outputs: a session folder, or a .npy, .h5 or .pickle trajectories file.

  • .npy, .npz or .pkl files with a positions array of shape (frames, individuals, 2), or with a dictionary like the idtracker.ai one (a "trajectories" key plus metadata).

  • .csv, .tsv or .txt tables, with one row per frame (x1, y1, x2, y2, ..., or columns like fish1_x, fish1_y) or one row per frame and individual (frame, id, x, y columns). Missing positions are empty cells or nan.

Positions are in pixels. The optional metadata keys are frames_per_second, body_length, arena_radius, arena_center, setup_points and fa_ignore_individuals_as_focal (identities never used as focal). Files that cannot hold this metadata (plain arrays, .npz and tables) read it from a JSON file with the same name, for example fish.json next to fish.csv. If the arena radius and center are not given, SocialNet estimates them from the trajectories.

References