Skip to content

Commit a5d2b25

Browse files
committed
Merge branch 'main' into dtw_clean
2 parents eeda0c3 + edb532f commit a5d2b25

62 files changed

Lines changed: 591 additions & 496 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

.readthedocs.yaml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,4 +25,4 @@ python:
2525

2626
# Build documentation in the "docs/" directory with Sphinx
2727
sphinx:
28-
configuration: docs/conf.py
28+
configuration: docs/conf.py

applications/benchmarking/DynaCLR/ImageNet/imagenet_embeddings.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@ def __init__(
5050
self.model = timm.create_model(model_name, pretrained=True)
5151
self.model.eval()
5252
except ImportError:
53-
raise ImportError("Please install the timm library: " "pip install timm")
53+
raise ImportError("Please install the timm library: pip install timm")
5454

5555
def _reduce_5d_input(self, x: torch.Tensor) -> torch.Tensor:
5656
"""Reduce 5D input (B, C, D, H, W) to 4D (B, C, H, W) using specified methods.
@@ -303,7 +303,7 @@ def main(config, model):
303303
phate_kwargs = cfg["embedding"]["phate_kwargs"]
304304

305305
if "umap_kwargs" in cfg["embedding"]:
306-
umap_kwargs = cfg["embedding"]["umap_kwargs"]
306+
cfg["embedding"]["umap_kwargs"]
307307

308308
if "pca_kwargs" in cfg["embedding"]:
309309
pca_kwargs = cfg["embedding"]["pca_kwargs"]

applications/benchmarking/DynaCLR/OpenPhenom/config_template.yml

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -13,11 +13,11 @@ model:
1313
# Default is "middle_slice" if not specified
1414
channel_reduction_methods:
1515
"Phase3D": "middle_slice" # For phase contrast, middle slice often works well
16-
"raw GFP EX488 EM525-45": "max"
16+
"raw GFP EX488 EM525-45": "max"
1717

1818
# Data module configuration
1919
datamodule:
20-
source_channel:
20+
source_channel:
2121
- Phase3D
2222
- "raw GFP EX488 EM525-45"
2323
z_range: [25, 40]
@@ -63,4 +63,4 @@ embedding:
6363
execution:
6464
overwrite: false
6565
save_config: true
66-
show_config: true
66+
show_config: true

applications/benchmarking/DynaCLR/OpenPhenom/openphenom_embeddings.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -57,8 +57,7 @@ def __init__(
5757
self.model.eval()
5858
except ImportError:
5959
raise ImportError(
60-
"Please install the OpenPhenom dependencies: "
61-
"pip install transformers"
60+
"Please install the OpenPhenom dependencies: pip install transformers"
6261
)
6362

6463
def on_predict_start(self):

applications/contrastive_phenotyping/evaluation/ALFI_MSD_v2.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
# %%
22
from pathlib import Path
3+
34
import matplotlib.pyplot as plt
45
import numpy as np
6+
57
from viscy.representation.embedding_writer import read_embedding_dataset
68
from viscy.representation.evaluation.distance import (
79
compute_displacement,

applications/contrastive_phenotyping/evaluation/PC_vs_computed_features.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,14 @@
1-
""" Script to compute the correlation between PCA and UMAP features and computed features
1+
"""Script to compute the correlation between PCA and UMAP features and computed features
22
* finds the computed features best representing the PCA and UMAP components
33
* outputs a heatmap of the correlation between PCA and UMAP features and computed features
44
"""
55

66
# %%
77
from pathlib import Path
8+
89
import matplotlib.pyplot as plt
910
import seaborn as sns
10-
from compute_pca_features import compute_features, compute_correlation_and_save_png
11+
from compute_pca_features import compute_correlation_and_save_png, compute_features
1112

1213
# %% for sensor features
1314

@@ -44,7 +45,7 @@
4445
# features_sensor = pd.read_csv("/hpc/projects/comp.micro/infected_cell_imaging/Single_cell_phenotyping/ContrastiveLearning/Figure_panels/cell_division/features_allset_sensor.csv")
4546

4647
# take a subset without the 768 features
47-
feature_columns = [f"feature_{i+1}" for i in range(768)]
48+
feature_columns = [f"feature_{i + 1}" for i in range(768)]
4849
features_subset_sensor = features_sensor.drop(columns=feature_columns)
4950
correlation_sensor = compute_correlation_and_save_png(
5051
features_subset_sensor,

applications/contrastive_phenotyping/evaluation/compute_pca_features.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -107,7 +107,7 @@ def compute_features(
107107
# convert the xarray to dataframe structure and add columns for computed features
108108
embedding_df = embedding_dataset["sample"].to_dataframe().reset_index(drop=True)
109109
feature_columns = pd.DataFrame(
110-
features_npy, columns=[f"feature_{i+1}" for i in range(768)]
110+
features_npy, columns=[f"feature_{i + 1}" for i in range(768)]
111111
)
112112

113113
embedding_df = pd.concat([embedding_df, feature_columns], axis=1)
@@ -167,7 +167,6 @@ def compute_features(
167167
unique_fov_names = sorted(list(set(fov_names_list)))
168168

169169
for fov_name in unique_fov_names:
170-
171170
unique_track_ids = embedding_df[embedding_df["fov_name"] == fov_name][
172171
"track_id"
173172
].unique()
@@ -180,7 +179,6 @@ def compute_features(
180179
(embedding_df["fov_name"] == fov_name)
181180
& (embedding_df["track_id"] == track_id)
182181
].empty:
183-
184182
prediction_dataset = dataset_of_tracks(
185183
data_path,
186184
tracks_path,
@@ -214,7 +212,6 @@ def compute_features(
214212
& (embedding_df["track_id"] == track_id)
215213
]["t"]
216214
):
217-
218215
# Basic statistical features for both channels
219216
phase_features = CellFeatures(phase[i], nucl_mask[i])
220217
PF = phase_features.compute_all_features()

applications/contrastive_phenotyping/evaluation/cosine_dissimilarity_dataset.py

Lines changed: 6 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -3,9 +3,14 @@
33
from typing import Optional
44

55
import matplotlib.pyplot as plt
6+
import numpy as np
7+
import pandas as pd
68
import seaborn as sns
7-
from sklearn.preprocessing import StandardScaler
89
from numpy.typing import NDArray
10+
from scipy.optimize import minimize_scalar
11+
from scipy.stats import gaussian_kde
12+
from sklearn.preprocessing import StandardScaler
13+
from tqdm import tqdm
914

1015
from viscy.representation.embedding_writer import read_embedding_dataset
1116
from viscy.representation.evaluation.clustering import (
@@ -14,13 +19,6 @@
1419
rank_nearest_neighbors,
1520
select_block,
1621
)
17-
import numpy as np
18-
from tqdm import tqdm
19-
import pandas as pd
20-
21-
from scipy.stats import gaussian_kde
22-
from scipy.optimize import minimize_scalar
23-
2422

2523
plt.style.use("../evaluation/figure.mplstyle")
2624

applications/contrastive_phenotyping/evaluation/displacement.py

Lines changed: 94 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -6,74 +6,112 @@
66

77
from viscy.representation.embedding_writer import read_embedding_dataset
88
from viscy.representation.evaluation.distance import (
9-
calculate_normalized_euclidean_distance_cell,
10-
compute_displacement_mean_std_full,
9+
calculate_normalized_euclidean_distance_cell,
10+
compute_displacement_mean_std_full,
1111
)
1212

13-
# %% paths
13+
# %% paths
1414

1515
features_path_30_min = Path(
16-
"/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/time_interval/predict/feb_test_time_interval_1_epoch_178.zarr"
16+
"/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/time_interval/predict/feb_test_time_interval_1_epoch_178.zarr"
1717
)
1818

19-
feature_path_no_track = Path("/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/negpair_random_sampling2/feb_fixed_test_predict.zarr")
19+
feature_path_no_track = Path(
20+
"/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/negpair_random_sampling2/feb_fixed_test_predict.zarr"
21+
)
2022

21-
features_path_any_time = Path("/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/negpair_difcell_randomtime_sampling/Ver2_updateTracking_refineModel/predictions/Feb_2chan_128patch_32projDim/2chan_128patch_56ckpt_FebTest.zarr")
23+
features_path_any_time = Path(
24+
"/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/negpair_difcell_randomtime_sampling/Ver2_updateTracking_refineModel/predictions/Feb_2chan_128patch_32projDim/2chan_128patch_56ckpt_FebTest.zarr"
25+
)
2226

2327
data_path = Path(
24-
"/hpc/projects/intracellular_dashboard/viral-sensor/2024_02_04_A549_DENV_ZIKV_timelapse/8-train-test-split/registered_test.zarr"
28+
"/hpc/projects/intracellular_dashboard/viral-sensor/2024_02_04_A549_DENV_ZIKV_timelapse/8-train-test-split/registered_test.zarr"
2529
)
2630

2731
tracks_path = Path(
28-
"/hpc/projects/intracellular_dashboard/viral-sensor/2024_02_04_A549_DENV_ZIKV_timelapse/8-train-test-split/track_test.zarr"
32+
"/hpc/projects/intracellular_dashboard/viral-sensor/2024_02_04_A549_DENV_ZIKV_timelapse/8-train-test-split/track_test.zarr"
2933
)
3034

3135
# %% Load embedding datasets for all three sampling
32-
fov_name = '/B/4/6'
36+
fov_name = "/B/4/6"
3337
track_id = 52
3438

3539
embedding_dataset_30_min = read_embedding_dataset(features_path_30_min)
3640
embedding_dataset_no_track = read_embedding_dataset(feature_path_no_track)
3741
embedding_dataset_any_time = read_embedding_dataset(features_path_any_time)
3842

39-
#%%
43+
# %%
4044
# Calculate displacement for each sampling
41-
time_points_30_min, cosine_similarities_30_min = calculate_normalized_euclidean_distance_cell(embedding_dataset_30_min, fov_name, track_id)
42-
time_points_no_track, cosine_similarities_no_track = calculate_normalized_euclidean_distance_cell(embedding_dataset_no_track, fov_name, track_id)
43-
time_points_any_time, cosine_similarities_any_time = calculate_normalized_euclidean_distance_cell(embedding_dataset_any_time, fov_name, track_id)
45+
time_points_30_min, cosine_similarities_30_min = (
46+
calculate_normalized_euclidean_distance_cell(
47+
embedding_dataset_30_min, fov_name, track_id
48+
)
49+
)
50+
time_points_no_track, cosine_similarities_no_track = (
51+
calculate_normalized_euclidean_distance_cell(
52+
embedding_dataset_no_track, fov_name, track_id
53+
)
54+
)
55+
time_points_any_time, cosine_similarities_any_time = (
56+
calculate_normalized_euclidean_distance_cell(
57+
embedding_dataset_any_time, fov_name, track_id
58+
)
59+
)
4460

4561
# %% Plot displacement over time for all three conditions
4662

4763
plt.figure(figsize=(10, 6))
4864

49-
plt.plot(time_points_no_track, cosine_similarities_no_track, marker='o', label='classical contrastive (no tracking)')
50-
plt.plot(time_points_any_time, cosine_similarities_any_time, marker='o', label='cell aware')
51-
plt.plot(time_points_30_min, cosine_similarities_30_min, marker='o', label='cell & time aware (interval 30 min)')
65+
plt.plot(
66+
time_points_no_track,
67+
cosine_similarities_no_track,
68+
marker="o",
69+
label="classical contrastive (no tracking)",
70+
)
71+
plt.plot(
72+
time_points_any_time, cosine_similarities_any_time, marker="o", label="cell aware"
73+
)
74+
plt.plot(
75+
time_points_30_min,
76+
cosine_similarities_30_min,
77+
marker="o",
78+
label="cell & time aware (interval 30 min)",
79+
)
5280

5381
plt.xlabel("Time Delay (t)", fontsize=10)
5482
plt.ylabel("Normalized Euclidean Distance with First Time Point", fontsize=10)
55-
plt.title("Normalized Euclidean Distance (Features) Over Time for Infected Cell", fontsize=12)
83+
plt.title(
84+
"Normalized Euclidean Distance (Features) Over Time for Infected Cell", fontsize=12
85+
)
5686

5787
plt.grid(True)
5888
plt.legend(fontsize=10)
5989

60-
#plt.savefig('4_euc_dist_full.svg', format='svg')
90+
# plt.savefig('4_euc_dist_full.svg', format='svg')
6191
plt.show()
6292

6393

6494
# %% Paths to datasets
65-
features_path_30_min = Path("/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/time_interval/predict/feb_test_time_interval_1_epoch_178.zarr")
66-
feature_path_no_track = Path("/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/negpair_random_sampling2/feb_fixed_test_predict.zarr")
95+
features_path_30_min = Path(
96+
"/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/time_interval/predict/feb_test_time_interval_1_epoch_178.zarr"
97+
)
98+
feature_path_no_track = Path(
99+
"/hpc/projects/intracellular_dashboard/viral-sensor/infection_classification/models/time_sampling_strategies/negpair_random_sampling2/feb_fixed_test_predict.zarr"
100+
)
67101

68102
embedding_dataset_30_min = read_embedding_dataset(features_path_30_min)
69103
embedding_dataset_no_track = read_embedding_dataset(feature_path_no_track)
70104

71105

72106
# %%
73-
max_tau = 10
107+
max_tau = 10
74108

75-
mean_displacement_30_min_euc, std_displacement_30_min_euc = compute_displacement_mean_std_full(embedding_dataset_30_min, max_tau)
76-
mean_displacement_no_track_euc, std_displacement_no_track_euc = compute_displacement_mean_std_full(embedding_dataset_no_track, max_tau)
109+
mean_displacement_30_min_euc, std_displacement_30_min_euc = (
110+
compute_displacement_mean_std_full(embedding_dataset_30_min, max_tau)
111+
)
112+
mean_displacement_no_track_euc, std_displacement_no_track_euc = (
113+
compute_displacement_mean_std_full(embedding_dataset_no_track, max_tau)
114+
)
77115

78116
# %% Plot 2: Cosine Displacements
79117
plt.figure(figsize=(10, 6))
@@ -83,24 +121,44 @@
83121
mean_values_30_min_euc = list(mean_displacement_30_min_euc.values())
84122
std_values_30_min_euc = list(std_displacement_30_min_euc.values())
85123

86-
plt.plot(taus, mean_values_30_min_euc, marker='o', label='Cell & Time Aware (30 min interval)', color='green')
87-
plt.fill_between(taus,
88-
np.array(mean_values_30_min_euc) - np.array(std_values_30_min_euc),
89-
np.array(mean_values_30_min_euc) + np.array(std_values_30_min_euc),
90-
color='green', alpha=0.3, label='Std Dev (30 min interval)')
124+
plt.plot(
125+
taus,
126+
mean_values_30_min_euc,
127+
marker="o",
128+
label="Cell & Time Aware (30 min interval)",
129+
color="green",
130+
)
131+
plt.fill_between(
132+
taus,
133+
np.array(mean_values_30_min_euc) - np.array(std_values_30_min_euc),
134+
np.array(mean_values_30_min_euc) + np.array(std_values_30_min_euc),
135+
color="green",
136+
alpha=0.3,
137+
label="Std Dev (30 min interval)",
138+
)
91139

92140
mean_values_no_track_euc = list(mean_displacement_no_track_euc.values())
93141
std_values_no_track_euc = list(std_displacement_no_track_euc.values())
94142

95-
plt.plot(taus, mean_values_no_track_euc, marker='o', label='Classical Contrastive (No Tracking)', color='blue')
96-
plt.fill_between(taus,
97-
np.array(mean_values_no_track_euc) - np.array(std_values_no_track_euc),
98-
np.array(mean_values_no_track_euc) + np.array(std_values_no_track_euc),
99-
color='blue', alpha=0.3, label='Std Dev (No Tracking)')
143+
plt.plot(
144+
taus,
145+
mean_values_no_track_euc,
146+
marker="o",
147+
label="Classical Contrastive (No Tracking)",
148+
color="blue",
149+
)
150+
plt.fill_between(
151+
taus,
152+
np.array(mean_values_no_track_euc) - np.array(std_values_no_track_euc),
153+
np.array(mean_values_no_track_euc) + np.array(std_values_no_track_euc),
154+
color="blue",
155+
alpha=0.3,
156+
label="Std Dev (No Tracking)",
157+
)
100158

101-
plt.xlabel('Time Shift (τ)')
102-
plt.ylabel('Euclidean Distance')
103-
plt.title('Embedding Displacement Over Time (Features)')
159+
plt.xlabel("Time Shift (τ)")
160+
plt.ylabel("Euclidean Distance")
161+
plt.title("Embedding Displacement Over Time (Features)")
104162

105163
plt.grid(True)
106164
plt.legend()

applications/contrastive_phenotyping/evaluation/imagenet_pretrained_features.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -111,7 +111,7 @@
111111

112112
x_train = data_train_val.drop(
113113
columns=[
114-
"division",
114+
"division",
115115
"fov_name",
116116
"t",
117117
"track_id",

0 commit comments

Comments
 (0)