Skip to content

Commit 30aafd8

Browse files
committed
Log the similarity between anchor and negatives (#332)
* log the negative * add tests for the triplet anchor and positive * add also the euclidean
1 parent a5d2b25 commit 30aafd8

3 files changed

Lines changed: 341 additions & 3 deletions

File tree

tests/conftest.py

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -114,3 +114,58 @@ def tracks_hcs_dataset(tmp_path_factory: TempPathFactory) -> Path:
114114
)
115115
fake_tracks.to_csv(dataset_path / fov_name / "tracks.csv", index=False)
116116
return dataset_path
117+
118+
119+
@fixture(scope="function")
120+
def tracks_with_gaps_dataset(tmp_path_factory: TempPathFactory) -> Path:
121+
"""Provides a HCS OME-Zarr dataset with tracking results with gaps in time."""
122+
dataset_path = tmp_path_factory.mktemp("tracks_gaps.zarr")
123+
_build_hcs(dataset_path, ["nuclei_labels"], (1, 256, 256), np.uint16, 3)
124+
125+
# Define different track patterns for different FOVs
126+
track_patterns = {
127+
"A/1/0": [
128+
# Track 0: complete sequence t=[0,1,2,3]
129+
{"track_id": 0, "t": 0, "y": 128, "x": 128, "id": 0},
130+
{"track_id": 0, "t": 1, "y": 128, "x": 128, "id": 1},
131+
{"track_id": 0, "t": 2, "y": 128, "x": 128, "id": 2},
132+
{"track_id": 0, "t": 3, "y": 128, "x": 128, "id": 3},
133+
# Track 1: ends early t=[0,1]
134+
{"track_id": 1, "t": 0, "y": 100, "x": 100, "id": 4},
135+
{"track_id": 1, "t": 1, "y": 100, "x": 100, "id": 5},
136+
],
137+
"A/1/1": [
138+
# Track 0: gap at t=2, has t=[0,1,3]
139+
{"track_id": 0, "t": 0, "y": 128, "x": 128, "id": 0},
140+
{"track_id": 0, "t": 1, "y": 128, "x": 128, "id": 1},
141+
{"track_id": 0, "t": 3, "y": 128, "x": 128, "id": 2},
142+
# Track 1: even timepoints only t=[0,2,4]
143+
{"track_id": 1, "t": 0, "y": 100, "x": 100, "id": 3},
144+
{"track_id": 1, "t": 2, "y": 100, "x": 100, "id": 4},
145+
{"track_id": 1, "t": 4, "y": 100, "x": 100, "id": 5},
146+
],
147+
"A/2/0": [
148+
# Track 0: single timepoint t=[0]
149+
{"track_id": 0, "t": 0, "y": 128, "x": 128, "id": 0},
150+
# Track 1: complete short sequence t=[0,1,2]
151+
{"track_id": 1, "t": 0, "y": 100, "x": 100, "id": 1},
152+
{"track_id": 1, "t": 1, "y": 100, "x": 100, "id": 2},
153+
{"track_id": 1, "t": 2, "y": 100, "x": 100, "id": 3},
154+
],
155+
}
156+
157+
for fov_name, _ in open_ome_zarr(dataset_path).positions():
158+
if fov_name in track_patterns:
159+
tracks_data = track_patterns[fov_name]
160+
else:
161+
# Default tracks for other FOVs
162+
tracks_data = [
163+
{"track_id": 0, "t": 0, "y": 128, "x": 128, "id": 0},
164+
]
165+
166+
tracks_df = pd.DataFrame(tracks_data)
167+
tracks_df["parent_track_id"] = -1
168+
tracks_df["parent_id"] = -1
169+
tracks_df.to_csv(dataset_path / fov_name / "tracks.csv", index=False)
170+
171+
return dataset_path

tests/data/test_triplet.py

Lines changed: 242 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
from iohub import open_ome_zarr
33
from pytest import mark
44

5-
from viscy.data.triplet import TripletDataModule
5+
from viscy.data.triplet import TripletDataModule, TripletDataset
66

77

88
@mark.parametrize("include_wells", [None, ["A/1", "A/2", "B/1"]])
@@ -109,3 +109,244 @@ def test_datamodule_z_window_size(
109109
expected_z_shape,
110110
*yx_patch_size,
111111
)
112+
113+
114+
def test_filter_anchors_time_interval_any(
115+
preprocessed_hcs_dataset, tracks_with_gaps_dataset
116+
):
117+
"""Test that time_interval='any' returns all tracks unchanged."""
118+
with open_ome_zarr(preprocessed_hcs_dataset) as dataset:
119+
channel_names = dataset.channel_names
120+
positions = list(dataset.positions())
121+
122+
# Create dataset with time_interval="any"
123+
tracks_tables = []
124+
for fov_name, _ in positions:
125+
tracks_df = pd.read_csv(
126+
next((tracks_with_gaps_dataset / fov_name).glob("*.csv"))
127+
).astype(int)
128+
tracks_tables.append(tracks_df)
129+
130+
total_tracks = sum(len(df) for df in tracks_tables)
131+
132+
ds = TripletDataset(
133+
positions=[pos for _, pos in positions],
134+
tracks_tables=tracks_tables,
135+
channel_names=channel_names,
136+
initial_yx_patch_size=(64, 64),
137+
z_range=slice(4, 9),
138+
fit=True,
139+
time_interval="any",
140+
)
141+
142+
# Should return all tracks
143+
assert len(ds.valid_anchors) == total_tracks
144+
145+
146+
def test_filter_anchors_time_interval_1(
147+
preprocessed_hcs_dataset, tracks_with_gaps_dataset
148+
):
149+
"""Test filtering with time_interval=1."""
150+
with open_ome_zarr(preprocessed_hcs_dataset) as dataset:
151+
channel_names = dataset.channel_names
152+
positions = list(dataset.positions())
153+
154+
tracks_tables = []
155+
for fov_name, _ in positions:
156+
tracks_df = pd.read_csv(
157+
next((tracks_with_gaps_dataset / fov_name).glob("*.csv"))
158+
).astype(int)
159+
tracks_tables.append(tracks_df)
160+
161+
ds = TripletDataset(
162+
positions=[pos for _, pos in positions],
163+
tracks_tables=tracks_tables,
164+
channel_names=channel_names,
165+
initial_yx_patch_size=(64, 64),
166+
z_range=slice(4, 9),
167+
fit=True,
168+
time_interval=1,
169+
)
170+
171+
# Check expected anchors per FOV/track
172+
valid_anchors = ds.valid_anchors
173+
174+
# FOV A/1/0, Track 0: t=[0,1,2,3] -> valid anchors at t=[0,1,2]
175+
fov_a10_track0 = valid_anchors[
176+
(valid_anchors["fov_name"] == "A/1/0") & (valid_anchors["track_id"] == 0)
177+
]
178+
assert set(fov_a10_track0["t"]) == {0, 1, 2}
179+
180+
# FOV A/1/0, Track 1: t=[0,1] -> valid anchor at t=[0]
181+
fov_a10_track1 = valid_anchors[
182+
(valid_anchors["fov_name"] == "A/1/0") & (valid_anchors["track_id"] == 1)
183+
]
184+
assert set(fov_a10_track1["t"]) == {0}
185+
186+
# FOV A/1/1, Track 0: t=[0,1,3] -> valid anchor at t=[0] only (t=1 has no t+1=2)
187+
fov_a11_track0 = valid_anchors[
188+
(valid_anchors["fov_name"] == "A/1/1") & (valid_anchors["track_id"] == 0)
189+
]
190+
assert set(fov_a11_track0["t"]) == {0}
191+
192+
# FOV A/1/1, Track 1: t=[0,2,4] -> no valid anchors (gaps of 2, no consecutive t+1)
193+
fov_a11_track1 = valid_anchors[
194+
(valid_anchors["fov_name"] == "A/1/1") & (valid_anchors["track_id"] == 1)
195+
]
196+
assert len(fov_a11_track1) == 0
197+
198+
# FOV A/2/0, Track 0: t=[0] -> no valid anchors (no t+1)
199+
fov_a20_track0 = valid_anchors[
200+
(valid_anchors["fov_name"] == "A/2/0") & (valid_anchors["track_id"] == 0)
201+
]
202+
assert len(fov_a20_track0) == 0
203+
204+
# FOV A/2/0, Track 1: t=[0,1,2] -> valid anchors at t=[0,1]
205+
fov_a20_track1 = valid_anchors[
206+
(valid_anchors["fov_name"] == "A/2/0") & (valid_anchors["track_id"] == 1)
207+
]
208+
assert set(fov_a20_track1["t"]) == {0, 1}
209+
210+
211+
def test_filter_anchors_time_interval_2(
212+
preprocessed_hcs_dataset, tracks_with_gaps_dataset
213+
):
214+
"""Test filtering with time_interval=2."""
215+
with open_ome_zarr(preprocessed_hcs_dataset) as dataset:
216+
channel_names = dataset.channel_names
217+
positions = list(dataset.positions())
218+
219+
tracks_tables = []
220+
for fov_name, _ in positions:
221+
tracks_df = pd.read_csv(
222+
next((tracks_with_gaps_dataset / fov_name).glob("*.csv"))
223+
).astype(int)
224+
tracks_tables.append(tracks_df)
225+
226+
ds = TripletDataset(
227+
positions=[pos for _, pos in positions],
228+
tracks_tables=tracks_tables,
229+
channel_names=channel_names,
230+
initial_yx_patch_size=(64, 64),
231+
z_range=slice(4, 9),
232+
fit=True,
233+
time_interval=2,
234+
)
235+
236+
valid_anchors = ds.valid_anchors
237+
238+
# FOV A/1/0, Track 0: t=[0,1,2,3] -> valid anchors at t=[0,1] (t+2 available)
239+
fov_a10_track0 = valid_anchors[
240+
(valid_anchors["fov_name"] == "A/1/0") & (valid_anchors["track_id"] == 0)
241+
]
242+
assert set(fov_a10_track0["t"]) == {0, 1}
243+
244+
# FOV A/1/0, Track 1: t=[0,1] -> no valid anchors (no t+2)
245+
fov_a10_track1 = valid_anchors[
246+
(valid_anchors["fov_name"] == "A/1/0") & (valid_anchors["track_id"] == 1)
247+
]
248+
assert len(fov_a10_track1) == 0
249+
250+
# FOV A/1/1, Track 0: t=[0,1,3] -> valid anchor at t=[1] (t=1+2=3 exists)
251+
fov_a11_track0 = valid_anchors[
252+
(valid_anchors["fov_name"] == "A/1/1") & (valid_anchors["track_id"] == 0)
253+
]
254+
assert set(fov_a11_track0["t"]) == {1}
255+
256+
# FOV A/1/1, Track 1: t=[0,2,4] -> valid anchors at t=[0,2]
257+
fov_a11_track1 = valid_anchors[
258+
(valid_anchors["fov_name"] == "A/1/1") & (valid_anchors["track_id"] == 1)
259+
]
260+
assert set(fov_a11_track1["t"]) == {0, 2}
261+
262+
# FOV A/2/0, Track 1: t=[0,1,2] -> valid anchor at t=[0]
263+
fov_a20_track1 = valid_anchors[
264+
(valid_anchors["fov_name"] == "A/2/0") & (valid_anchors["track_id"] == 1)
265+
]
266+
assert set(fov_a20_track1["t"]) == {0}
267+
268+
269+
def test_filter_anchors_cross_fov_independence(
270+
preprocessed_hcs_dataset, tracks_with_gaps_dataset
271+
):
272+
"""Test that same track_id in different FOVs are treated independently."""
273+
with open_ome_zarr(preprocessed_hcs_dataset) as dataset:
274+
channel_names = dataset.channel_names
275+
positions = list(dataset.positions())
276+
277+
tracks_tables = []
278+
for fov_name, _ in positions:
279+
tracks_df = pd.read_csv(
280+
next((tracks_with_gaps_dataset / fov_name).glob("*.csv"))
281+
).astype(int)
282+
tracks_tables.append(tracks_df)
283+
284+
ds = TripletDataset(
285+
positions=[pos for _, pos in positions],
286+
tracks_tables=tracks_tables,
287+
channel_names=channel_names,
288+
initial_yx_patch_size=(64, 64),
289+
z_range=slice(4, 9),
290+
fit=True,
291+
time_interval=1,
292+
)
293+
294+
# Check global_track_id format and uniqueness
295+
assert "global_track_id" in ds.tracks.columns
296+
global_track_ids = ds.tracks["global_track_id"].unique()
297+
298+
# Verify format: should be "fov_name_track_id"
299+
for gid in global_track_ids:
300+
assert "_" in gid
301+
fov_part, track_id_part = gid.rsplit("_", 1)
302+
assert "/" in fov_part # FOV names contain slashes like "A/1/0"
303+
304+
# Track 0 exists in multiple FOVs (A/1/0, A/1/1, A/2/0) but should have different global_track_ids
305+
track0_global_ids = ds.tracks[ds.tracks["track_id"] == 0][
306+
"global_track_id"
307+
].unique()
308+
assert len(track0_global_ids) >= 3 # At least 3 different FOVs with track_id=0
309+
310+
# Verify that filtering is independent per FOV
311+
# A/1/0 Track 0 (continuous) should have more valid anchors than A/1/1 Track 0 (with gap)
312+
valid_a10_track0 = ds.valid_anchors[
313+
(ds.valid_anchors["fov_name"] == "A/1/0") & (ds.valid_anchors["track_id"] == 0)
314+
]
315+
valid_a11_track0 = ds.valid_anchors[
316+
(ds.valid_anchors["fov_name"] == "A/1/1") & (ds.valid_anchors["track_id"] == 0)
317+
]
318+
# A/1/0 Track 0 has t=[0,1,2] valid (3 anchors)
319+
# A/1/1 Track 0 has t=[0] valid (1 anchor, gap at t=2)
320+
assert len(valid_a10_track0) == 3
321+
assert len(valid_a11_track0) == 1
322+
323+
324+
def test_filter_anchors_predict_mode(
325+
preprocessed_hcs_dataset, tracks_with_gaps_dataset
326+
):
327+
"""Test that predict mode (fit=False) returns all tracks regardless of time_interval."""
328+
with open_ome_zarr(preprocessed_hcs_dataset) as dataset:
329+
channel_names = dataset.channel_names
330+
positions = list(dataset.positions())
331+
332+
tracks_tables = []
333+
for fov_name, _ in positions:
334+
tracks_df = pd.read_csv(
335+
next((tracks_with_gaps_dataset / fov_name).glob("*.csv"))
336+
).astype(int)
337+
tracks_tables.append(tracks_df)
338+
339+
total_tracks = sum(len(df) for df in tracks_tables)
340+
341+
ds = TripletDataset(
342+
positions=[pos for _, pos in positions],
343+
tracks_tables=tracks_tables,
344+
channel_names=channel_names,
345+
initial_yx_patch_size=(64, 64),
346+
z_range=slice(4, 9),
347+
fit=False, # Predict mode
348+
time_interval=1,
349+
)
350+
351+
# Should return all tracks even with time_interval=1
352+
assert len(ds.valid_anchors) == total_tracks

viscy/representation/engine.py

Lines changed: 44 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@ def __init__(
3636
log_batches_per_epoch: int = 8,
3737
log_samples_per_batch: int = 1,
3838
log_embeddings: bool = False,
39+
log_negative_metrics_every_n_epochs: int = 2,
3940
example_input_array_shape: Sequence[int] = (1, 2, 15, 256, 256),
4041
) -> None:
4142
super().__init__()
@@ -49,6 +50,7 @@ def __init__(
4950
self.training_step_outputs = []
5051
self.validation_step_outputs = []
5152
self.log_embeddings = log_embeddings
53+
self.log_negative_metrics_every_n_epochs = log_negative_metrics_every_n_epochs
5254

5355
def forward(self, x: Tensor) -> tuple[Tensor, Tensor]:
5456
"""Return both features and projections.
@@ -94,8 +96,8 @@ def _log_metrics(
9496
cosine_sim_pos = F.cosine_similarity(anchor, positive, dim=1).mean()
9597
euclidean_dist_pos = F.pairwise_distance(anchor, positive).mean()
9698
log_metric_dict = {
97-
f"metrics/cosine_similarity_positive/{stage}": cosine_sim_pos,
98-
f"metrics/euclidean_distance_positive/{stage}": euclidean_dist_pos,
99+
f"metrics/cosine_similarity/positive/{stage}": cosine_sim_pos,
100+
f"metrics/euclidean_distance/positive/{stage}": euclidean_dist_pos,
99101
}
100102

101103
if negative is not None:
@@ -107,6 +109,46 @@ def _log_metrics(
107109
log_metric_dict[f"metrics/euclidean_distance_negative/{stage}"] = (
108110
euclidean_dist_neg
109111
)
112+
elif isinstance(self.loss_function, NTXentLoss):
113+
if self.current_epoch % self.log_negative_metrics_every_n_epochs == 0:
114+
batch_size = anchor.size(0)
115+
116+
# Cosine similarity metrics
117+
anchor_norm = F.normalize(anchor, dim=1)
118+
positive_norm = F.normalize(positive, dim=1)
119+
all_embeddings_norm = torch.cat([anchor_norm, positive_norm], dim=0)
120+
sim_matrix = torch.mm(anchor_norm, all_embeddings_norm.t())
121+
122+
mask = torch.ones_like(sim_matrix, dtype=torch.bool)
123+
mask[range(batch_size), range(batch_size)] = False # Exclude self
124+
mask[range(batch_size), range(batch_size, 2 * batch_size)] = (
125+
False # Exclude positive
126+
)
127+
128+
negative_sims = sim_matrix[mask].view(batch_size, -1)
129+
130+
mean_neg_sim = negative_sims.mean()
131+
sum_neg_sim = negative_sims.sum(dim=1).mean()
132+
margin_cosine = cosine_sim_pos - mean_neg_sim
133+
134+
all_embeddings = torch.cat([anchor, positive], dim=0)
135+
dist_matrix = torch.cdist(anchor, all_embeddings, p=2)
136+
negative_dists = dist_matrix[mask].view(batch_size, -1)
137+
138+
mean_neg_dist = negative_dists.mean()
139+
sum_neg_dist = negative_dists.sum(dim=1).mean()
140+
margin_euclidean = mean_neg_dist - euclidean_dist_pos
141+
142+
log_metric_dict.update(
143+
{
144+
f"metrics/cosine_similarity/negative_mean/{stage}": mean_neg_sim,
145+
f"metrics/cosine_similarity/negative_sum/{stage}": sum_neg_sim,
146+
f"metrics/margin_positive/negative/{stage}": margin_cosine,
147+
f"metrics/euclidean_distance/negative_mean/{stage}": mean_neg_dist,
148+
f"metrics/euclidean_distance/negative_sum/{stage}": sum_neg_dist,
149+
f"metrics/margin_euclidean_positive/negative/{stage}": margin_euclidean,
150+
}
151+
)
110152

111153
self.log_dict(
112154
log_metric_dict,

0 commit comments

Comments
 (0)