Skip to content

Repository files navigation

Tracing Flow

Tracing Flow reconstructs particle trajectories from discrete time snapshots. It combines optimal transport between adjacent time points with neural networks that learn initial velocity and continuous-time acceleration.

The current entry point is Train.py.

Files

  • Train.py: example training script for SimLineage
  • tracingflow.py: OT matching, trajectory sampling, training, visualization
  • dataset.py: dataset wrappers for CSV files in datasets/
  • models.py: InitialVelocityNet and AccelerationNet
  • config.py: hyperparameters
  • utils.py: keypoint export helpers

Setup

Create the Conda environment:

conda env create -f enviroment.yaml
conda activate tracingflow

Create the checkpoint directory before training because the code saves weights to ./checkpoints/ but does not create it automatically:

mkdir checkpoints

If your machine does not have CUDA, change device = "cuda" to device = "cpu" in config.py.

Train

Run the default example:

python Train.py

Train.py currently trains SimLineage with dim = 2.

Training flow:

  1. Load a dataset.
  2. Build InitialVelocityNet and AccelerationNet.
  3. Train with TracingFlow.train_tracing_flow(...).
  4. Generate trajectories with generate_trajectory(...).
  5. Plot results and save keypoints to output_data/.

Switch Datasets

To use another dataset, edit Train.py and change the dataset class and dim.

Example:

from dataset import SimLineage3D

my_dataset = SimLineage3D()
dim = 3
my_ini_vel_net = InitialVelocityNet(dim=dim)
my_acc_net = AccelerationNet(dim=dim)

dim must match the feature dimension of the selected dataset.

Data Format

This repository does not include a standalone data generation script. To use your own data, place a CSV file in datasets/ and add or reuse a dataset class in dataset.py.

Most datasets use:

samples,x1,x2,...,xD,barcodes
0,...
1,...
  • samples: discrete time index
  • x1 ... xD: coordinates or features
  • barcodes: optional label column

SimLineage3D uses this format instead:

time,x1,x2,x3,barcodes
0,...
1,...

For a new dataset class, define:

  • self.csv_path
  • self.time_steps
  • feature column slicing
  • whether the time column is samples or time

Included Dataset Classes

Main dataset wrappers in dataset.py:

  • SimLineage
  • SimLineage3D
  • Cite5D
  • Cite100D
  • RealLineage
  • Cite5D_wo_1, Cite5D_wo_2, Cite5D_wo_3
  • EB5D, EB5D_wo_1, EB5D_wo_2, EB5D_wo_3, EB5D_wo_4
  • Gulf_data
  • Molecular_data

High-dimensional data are projected to 2D with PCA for visualization.

Outputs

Training saves model weights to checkpoints/:

  • ini_vel_<save_name>.pth
  • acc_<save_name>.pth

save_keypoints(...) saves:

  • <prefix>_position_time_1.npy
  • <prefix>_weight_time_1.npy
  • <prefix>_barcode_time_1.npy when labels are provided

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages