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)
SocialNet#
Source code
Check the source code at polavieja_lab/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.
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.#
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:
Install SocialNet in its own environment
To keep SocialNet apart from idtracker.ai, create a new environment with Python 3.10 or later:
Install PyTorch with the command for your computer from the PyTorch website, then install SocialNet:
Basic usage of SocialNet#
Train a model with
model_train()and test it withmodel_test():model_trainTrain SocialNet model from a set of trajectory files.
model_testTest the model on the given trajectory files.
Then analyze the trained model with these plotting functions:
plot.plot_aggregation_subnetworkVisualizes the aggregation subnetwork of the model by plotting how weights vary as a function of the specified variables.
plot.plot_interaction_subnetworkPlots 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_scoresPlots processed interaction scores (attraction, alignment, and repulsion) as a function of kinematic variables of the focal and neighbour animals.
plot.plot_productPlots the combination (product) of
plot.plot_interaction_scores()andplot.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.
Alternative SocialNet CLI
SocialNet can also be used from the command line, without writing Python code:
Trajectory formats#
model_train()andmodel_test()(and thetrain,testandsamplecommands) accept:idtracker.ai outputs: a session folder, or a
.npy,.h5or.pickletrajectories file..npy,.npzor.pklfiles with a positions array of shape(frames, individuals, 2), or with a dictionary like the idtracker.ai one (a"trajectories"key plus metadata)..csv,.tsvor.txttables, with one row per frame (x1, y1, x2, y2, ..., or columns likefish1_x, fish1_y) or one row per frame and individual (frame, id, x, ycolumns). Missing positions are empty cells ornan.Positions are in pixels. The optional metadata keys are
frames_per_second,body_length,arena_radius,arena_center,setup_pointsandfa_ignore_individuals_as_focal(identities never used as focal). Files that cannot hold this metadata (plain arrays,.npzand tables) read it from a JSON file with the same name, for examplefish.jsonnext tofish.csv. If the arena radius and center are not given, SocialNet estimates them from the trajectories.References