ledtrack 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
ledtrack/__init__.py ADDED
File without changes
@@ -0,0 +1,692 @@
1
+ import numpy as np
2
+ import tifffile
3
+ import pandas as pd
4
+ import time
5
+ from multiprocessing import Process as multirun
6
+ from multiprocessing import Manager #, Queue, JoinableQueue
7
+ from itertools import chain
8
+ import datetime
9
+ now = datetime.datetime.now()
10
+ from .utils import get_cfg
11
+ import hydra
12
+ from hydra.utils import to_absolute_path as abs_path
13
+ from .linear_solver import linear_solver
14
+ from os.path import join
15
+ from scipy.spatial.distance import cdist
16
+ from omegaconf import DictConfig
17
+ # from GUI import LineageTree as LT
18
+
19
+
20
+ def merging_and_pruning(cfg, track, centroid,
21
+ lieange_tree=False):
22
+ """对细胞轨迹进行合并和修剪处理。
23
+
24
+ 通过合并邻近帧细胞距离较近的轨迹,以及修剪过短的轨迹,
25
+ 来优化细胞追踪结果,减少由碎片和过度分割导致的错误轨迹。
26
+
27
+ Args:
28
+ cfg: 配置对象,包含处理参数
29
+ track: 轨迹数组,形状为(n_tracks, n_frames)的np.ndarray
30
+ centroid: 细胞质心坐标数组
31
+ lieange_tree: 是否显示家系树结构,默认为False
32
+
33
+ Returns:
34
+ 合并和修剪后的轨迹数组,形状为(n_tracks, n_frames)的np.ndarray,
35
+ 转置后返回以使维度为(n_frames, n_tracks)
36
+ """
37
+ # if lieange_tree:
38
+ # lt = LT.LineageTree(track, if_scene=False).tree
39
+ # lt.show()
40
+
41
+ # # merging the track whose neighboring-frame cells are close
42
+ if cfg.track.merge:
43
+ track = merge_track(cfg, track, centroid)
44
+ # # # # prune too short tracks (caused by debris)
45
+ mm = track[:, -1] != -1
46
+ len_ = track.shape[1] - np.sum(track == -1, axis=1)
47
+ ll = len_ > cfg.track.min_length
48
+ track = track[np.logical_or(mm, ll)]
49
+
50
+ # prune too short leaf tracks (caused by debris and over-segmentation)
51
+ if cfg.track.prune_leaf:
52
+ track = prune(track, cfg.track.min_length)
53
+
54
+ return track.T
55
+
56
+
57
+ def leaf_length(track, tr):
58
+ cell_id = np.where(tr != -1)[0]
59
+ if cell_id.size == 0:
60
+ return 0, 0, 0, -1
61
+ branch = np.array([np.sum(track[:, id] == tr[id]) for id in cell_id])
62
+ leaf = cell_id[np.where(branch == 1)[0]]
63
+ if leaf.size == 0:
64
+ return 0, cell_id[-1], cell_id[0], cell_id[-1]
65
+ return leaf.max() - leaf.min() + 1, leaf.max(), leaf.min(), cell_id[-1]
66
+
67
+
68
+ def check_track(track):
69
+ for frame, cells in enumerate(track[:-1]):
70
+ cell_list = np.unique(cells)
71
+ cell_list = cell_list[cell_list != -1]
72
+ lnk_num = [len(np.unique(track[frame + 1, np.where(cells == c)[0]])) for c in cell_list]
73
+ if np.any(np.array(lnk_num) > 2):
74
+ return False
75
+ return True
76
+
77
+
78
+ def prune(track, min_length=10):
79
+ """
80
+ 裁剪短轨迹
81
+
82
+ 根据最小长度阈值删除过短的轨迹。算法会迭代查找并裁剪短轨迹,
83
+ 对于每个轨迹,找出共享相同父细胞的分支,保留较长的分支而删除较短的。
84
+
85
+ Args:
86
+ track: 轨迹数据,shape=(n_tracks, n_frames),np.ndarray
87
+ min_length: 轨迹最小长度阈值
88
+
89
+ Returns:
90
+ 裁剪后的轨迹,shape=(n_tracks, n_frames),np.ndarray
91
+ """
92
+ not_all_pruned = True
93
+ while not_all_pruned:
94
+ all_leaf_len = [leaf_length(track, tr) for tr in track]
95
+ prune_trs = []
96
+ eq_len = []
97
+ for ith, tr in enumerate(track):
98
+
99
+ not_max_leaf = False
100
+ leaf_len, leaf_max, leaf_min, track_max = all_leaf_len[ith]
101
+ branch = np.where(track[:, leaf_min-1] == tr[leaf_min-1])[0] \
102
+ if leaf_min > 0 or tr[leaf_min-1] != -1 else np.array([])
103
+ for i in branch:
104
+ if i == ith:
105
+ continue
106
+ leaf_len_, leaf_max_, leaf_min_, track_max_ = all_leaf_len[i]
107
+ if track_max_ > track_max:
108
+ not_max_leaf = True
109
+ break
110
+ elif track_max_ == track_max and leaf_len < min_length:
111
+ prune_trs.append(max(ith, i))
112
+ eq_len.append([min(ith, i), leaf_min])
113
+ if leaf_len < min_length and not_max_leaf and leaf_max != track.shape[1] - 1:
114
+ prune_trs.append(ith)
115
+ track = np.delete(track, prune_trs, axis=0)
116
+ if len(prune_trs) == 0:
117
+ not_all_pruned = False
118
+
119
+ return track
120
+
121
+
122
+ def merge_track(cfg, track, centroid,
123
+ itv=1,
124
+ merge_branch=True):
125
+ """
126
+ 合并细胞轨迹,包括分裂轨迹的合并和跨间隔轨迹的合并
127
+
128
+ Args:
129
+ cfg: 配置对象,包含轨迹相关设置
130
+ track: 轨迹数组,形状为[轨迹数量, 帧数],存储每条轨迹在各帧中的细胞ID
131
+ centroid: 质心信息,包含各帧中细胞的位置、大小等信息
132
+ itv: 跨帧间隔,默认为1
133
+ merge_branch: 是否合并分支轨迹,默认为True
134
+
135
+ Returns:
136
+ track: 合并后的轨迹数组,形状可能比输入小
137
+ """
138
+ nearest = cfg.track.nearest
139
+ ## itv>1 时后续分析还有bug, 请将itv设置为1
140
+ merge_track = []
141
+ for ith, tr in enumerate(track):
142
+ cell_id0 = np.where(tr != -1)[0]
143
+ frame_min, frame_max = np.min(cell_id0), np.max(cell_id0)
144
+
145
+ for forward in range(frame_max + 1, frame_max + itv + 1):
146
+ if forward < track.shape[1]:
147
+ cndd_idx = np.where(track[:, frame_max] == -1)[0]
148
+ cndd_idx1 = np.where(track[cndd_idx, forward - 1] == -1)[0]
149
+ merge_idx = np.where(track[cndd_idx[cndd_idx1], forward] != -1)[0]
150
+ merge_idx = cndd_idx[cndd_idx1[merge_idx]]
151
+ if len(merge_idx) > 0:
152
+ # ### not consider of m
153
+ dist = np.linalg.norm(centroid[frame_max][tr[frame_max], 0, :2] -
154
+ centroid[forward][track[merge_idx, forward], 0, :2], 2, axis=1)
155
+ size = centroid[frame_max][tr[frame_max], 0, 3]
156
+ size_list = centroid[forward][track[merge_idx, forward], 0, 3]
157
+ # dist1 = [np.linalg.norm(centroid[frame_max][tr[frame_max], 0, :2] -
158
+ # centroid[forward][track[idx, forward], 0, :2], 2) for idx in merge_idx]
159
+ # ### consider of m
160
+ # dist = [np.linalg.norm(centroid[frame_max][tr[frame_max], 0] -
161
+ # np.sum(centroid[forward][track[idx, forward]], axis=0), 2)
162
+ # for idx in merge_idx]
163
+
164
+ if_merge = np.argsort(dist)[0]
165
+ if nearest or (dist[if_merge] < cfg.track.thr_dist
166
+ and size_list[if_merge] > 0.6 * size
167
+ and size_list[if_merge] < 1.4 * size):
168
+ cell = track[merge_idx[if_merge], forward]
169
+ for idx in np.where(track[:, forward] == cell)[0]:
170
+ track[idx, :forward] = track[ith, :forward]
171
+ # print('merge', idx, 'to', ith, 'with distance', dist[if_merge], 'frame', forward)
172
+ merge_track.append(ith)
173
+ break
174
+ track = np.delete(track, merge_track, axis=0)
175
+
176
+ if merge_branch:
177
+ lnk_num_track = np.zeros_like(track, dtype=np.int8) - 1
178
+ for fth, cellss in enumerate(track.T[:-1]):
179
+ candd_mother = np.unique(cellss)
180
+ candd_mother = candd_mother[candd_mother != -1]
181
+ mother_id = [np.where(track[:, fth] == m)[0] for m in candd_mother]
182
+ for idxs in mother_id:
183
+ if len(np.unique(track[idxs, fth + 1])) > 2:
184
+ print(len(idxs))
185
+ lnk_num_track[idxs, fth] = len(np.unique(track[idxs, fth + 1]))
186
+ assert np.all(lnk_num_track < 3), 'wrong linking !!!'
187
+ # merge branching
188
+ for ith, tr in enumerate(track):
189
+ cell_id0 = np.where(tr != -1)[0]
190
+ frame_min = np.min(cell_id0)
191
+ if frame_min == 0:
192
+ continue
193
+ centx, centy, label, size = centroid[frame_min][tr[frame_min]][0]
194
+ dist_edge = np.min((centx, centy, cfg.shape[0] - centx, cfg.shape[1] - centy))
195
+ if dist_edge > cfg.track.thr_dist:
196
+ dist = np.linalg.norm(centroid[frame_min - 1][:, 0, :2] - np.array((centx, centy)), 2, axis=1)
197
+ size_list = centroid[frame_min - 1][:, 0, 3]
198
+ # size_diff = size - size_list
199
+ # size_diff_rel = np.abs(size_diff) / size
200
+ candd = np.argsort(dist)
201
+ ## debug: 排除已有两个连接的细胞
202
+ for id in candd:
203
+ # if id in candd_mother_id:
204
+ trackid = np.where(track[:, frame_min - 1] == id)[0]
205
+ cell_cycle = np.all(lnk_num_track[trackid,
206
+ frame_min - cfg.track.div_interval: frame_min + cfg.track.div_interval] == 1)
207
+ if nearest or (cell_cycle
208
+ and len(trackid) > 0
209
+ and dist[id] < cfg.track.thr_dist
210
+ and size_list[id] < 0.7*size
211
+ and size_list[id] > 0.3*size):
212
+ # trackid = np.where(track[:, frame_min - 1] == candd)[0]
213
+ mergeid = track[:, frame_min] == track[ith, frame_min]
214
+ track[mergeid, :frame_min] = track[trackid[0], :frame_min]
215
+ lnk_num_track[trackid, frame_min - 1] += 1
216
+ lnk_num_track[mergeid, frame_min - 1] = lnk_num_track[trackid[0], frame_min - 1]
217
+ break
218
+ assert np.all(lnk_num_track < 3), 'wrong linking !!!'
219
+
220
+ assert check_track(track.T)
221
+ return track
222
+
223
+
224
+ def get_interference(num_cell, centroid, center,
225
+ radius=200, a=0.4, b=1, c=0):
226
+ """
227
+ 计算细胞间的干扰概率
228
+
229
+ 基于细胞总数和邻居数量计算细胞从一个位置移动到另一个位置的干扰概率。
230
+ 公式: P(x_(t+1,i)↛x_(t,j)) = 1/(α×total cell count + neighbor count + c)
231
+
232
+ Args:
233
+ num_cell: 细胞总数
234
+ centroid: 细胞的质心坐标数组,形状为(n, 2),包含每个细胞的x, y坐标
235
+ center: 中心点坐标,用于计算距离的参考中心
236
+ radius: 邻居搜索半径,默认为200,仅在此范围内的细胞计入邻居数
237
+ a: 总细胞数权重系数,默认为0.4
238
+ b: 邻居数权重系数,默认为1
239
+ c: 偏置常数,默认为0
240
+
241
+ Returns:
242
+ float: 干扰概率值,值越小表示干扰越强
243
+ """
244
+ points_array = np.array(centroid[:, :, :2])
245
+ center_array = np.array(center)
246
+ # the distance from each point to the center
247
+ distances = np.linalg.norm(points_array - center_array, axis=1)
248
+ # Count the number of cells within the radius
249
+ num_neighbor = np.sum(distances <= radius)
250
+
251
+ # P(x_(t+1)↛x_(t) )=1/(α×total cell count+neighbor count+1)
252
+ intf = 1 / (a * num_cell + b * num_neighbor + c)
253
+ return intf
254
+
255
+
256
+ def expand_virtual_node(theta, matrix, centroid, shape):
257
+ """
258
+ 为新进入视野的细胞添加虚拟节点到转移矩阵。
259
+
260
+ Args:
261
+ theta: 高斯函数参数,控制距离衰减速度。
262
+ matrix: 转移矩阵。
263
+ centroid: 细胞的质心坐标,形状为 (N, 2, 4)。
264
+ shape: 图像的形状 (height, width)。
265
+
266
+ Returns:
267
+ 扩展并转置后的转移矩阵,最后一列为虚拟节点的值。
268
+ """
269
+ # trans_matrix = np.hstack((matrix, np.zeros((matrix.shape[0], 1), dtype=np.float32)))
270
+
271
+ # for i, cent in enumerate(centroid[:, 0, :2]):
272
+ #
273
+ # dist_edge = np.min((cent[0], cent[1], shape[0]-cent[0], shape[1]-cent[1]))
274
+ # dist_edge_gau = np.exp(-dist_edge ** 2 / theta)
275
+ #
276
+ # dist_neighbor = np.max(matrix[i])
277
+ # # # # if dist_edge < 100 else -1 这会增加错误链接
278
+ # trans_matrix[i][-1] = np.max((dist_edge_gau, 1 - dist_neighbor)) # if dist_edge < 100 else -1
279
+
280
+ cent = centroid[:, 0, :2]
281
+ dist_edge = np.min([cent[:, 0], cent[:, 1], shape[0]-cent[:, 0], shape[1]-cent[:, 1]], axis=0)
282
+ dist_edge_gau = np.exp(-dist_edge ** 2 / theta)
283
+ dist_neighbor = np.max(matrix, axis=1)
284
+ virtual_node = np.max((dist_edge_gau, 1 - dist_neighbor), axis=0)
285
+ trans_matrix = np.hstack((matrix, virtual_node.reshape(-1, 1)))
286
+
287
+ return trans_matrix.T
288
+
289
+
290
+ def cal_trans_matrix(cfg,
291
+ centroid: list,
292
+ mask,
293
+ flow,
294
+ num_cell_f,
295
+ theta: float=800.,
296
+ jitter_thr: float=0.6,
297
+ global_norm: bool=True,
298
+ ):
299
+ """
300
+ 计算两帧之间的转移矩阵,用于细胞跟踪。
301
+
302
+ 根据光流信息和质心位置,计算从前一帧到当前帧的细胞转移概率矩阵。
303
+ 支持全局归一化和局部归一化两种计算模式。
304
+
305
+ Args:
306
+ cfg: 配置对象,包含跟踪参数(如max_movenment)
307
+ centroid: 质心列表,centroid[-1]为当前帧,centroid[-2]为前一帧
308
+ mask: 当前帧的分割掩码
309
+ flow: 光流数据,包含速度分量
310
+ num_cell_f: 每帧的细胞数量列表
311
+ theta: 距离衰减参数,控制转移概率的衰减速度,默认800
312
+ jitter_thr: 抖动阈值,低于此比例的位移视为抖动并置零,默认0.6
313
+ global_norm: 是否使用全局归一化模式,默认True
314
+
315
+ Returns:
316
+ 转移概率矩阵,表示从当前帧细胞到前一帧细胞的转移概率
317
+ """
318
+ num_pix = np.sum(mask > 0)
319
+ vel = np.sum(np.sum(flow[0:2], axis=1), axis=1).astype(np.float32) / num_pix
320
+ mag = np.sum(np.sqrt((flow[0] ** 2 + flow[1] ** 2))) / num_pix
321
+
322
+ if global_norm:
323
+ if np.sqrt(np.sum(vel ** 2)) / mag < jitter_thr:
324
+ vel = 0
325
+ vj, mo = centroid[-1][:, 0, :2], centroid[-1][:, 1, :2]
326
+ vk = centroid[-2][:, 0, :2]
327
+ dist_lh = cdist(vj + vel, vk, metric='sqeuclidean')
328
+ prior_lh = cdist(vj + mo, vk, metric='sqeuclidean')
329
+ post = np.exp(-dist_lh / theta) * np.exp(-prior_lh / theta * 4)
330
+ sum_norm = 1 / (np.sum(post, axis=1, keepdims=True) + 1e-8)
331
+ matrix = post * sum_norm
332
+ else:
333
+ matrix = []
334
+ for j, (vj, m) in enumerate(centroid[-1]):
335
+ vj, m = vj[:2], m[:2]
336
+ # 计算干扰 越小越均匀, 越大差异越大
337
+ intf = get_interference(num_cell_f[-1], centroid[-1], vj, radius=cfg.shape[0] / 6)
338
+ for k, (vk, _) in enumerate(centroid[-2]):
339
+ vk = vk[:2]
340
+ if np.sqrt(np.sum(vel ** 2)) / mag < jitter_thr:
341
+ vel = 0
342
+
343
+ dt = np.sum((vj - vk) ** 2)
344
+ ## 距离太大的直接不考虑, 给极小的负分数
345
+ if dt > cfg.track.max_movenment ** 2:
346
+ matrix.append(-1)
347
+ else:
348
+ dist = np.sum((vj + vel - vk) ** 2)
349
+ dist_lh = np.exp(-dist / theta) # / (theta * np.sqrt(2 * np.pi)) ## 快速计算,不算常数因子,不影响结果
350
+ dist = np.sum((vj + m - vk) ** 2)
351
+ prior = np.exp(-dist / theta * 2) # / (theta * np.sqrt(2 * np.pi)) ## 快速计算,不算常数因子,不影响结果
352
+ post = prior * dist_lh / (prior * dist_lh + (1 - prior) * intf)
353
+ matrix.append(post)
354
+ matrix = np.array(matrix, dtype=np.float32).reshape(len(centroid[-1]), len(centroid[-2]))
355
+ matrix = expand_virtual_node(theta, matrix, centroid[-1], cfg.shape)
356
+
357
+ return matrix
358
+
359
+
360
+ def get_track(cfg, q,
361
+ mask_dir: list,
362
+ flow_dir: list,
363
+ run_id: int,
364
+ max_run: int,
365
+ jitter_thr=0.6,
366
+ max_division:int=None,
367
+ new_detect:int=None,
368
+ ) -> None:
369
+ """
370
+ 获取跟踪过渡矩阵
371
+
372
+ 读取掩码和光流数据,计算细胞质心,过渡矩阵,并通过线性求解器进行细胞匹配。
373
+
374
+ Args:
375
+ cfg: 配置对象
376
+ q: 多进程队列
377
+ mask_dir: 掩码文件路径列表
378
+ flow_dir: 光流文件路径列表
379
+ run_id: 运行 ID
380
+ max_run: 最大运行次数
381
+ jitter_thr: 抖动阈值,默认为 0.6
382
+ max_division: 最大分裂数,可选
383
+ new_detect: 新检测参数,可选
384
+
385
+ Returns:
386
+ None,结果通过队列返回
387
+ """
388
+ print(f'------{run_id}th running start-------\n')
389
+ num_cell_f = []
390
+ centroid = []
391
+ trans_matrix = []
392
+ match = []
393
+
394
+ assert len(mask_dir) >= 2, 'mask dir should be longer than 2'
395
+ assert len(flow_dir) >= 1, 'flow dir should be longer than 1'
396
+
397
+ mask = tifffile.imread(mask_dir[0])
398
+ labels = np.unique(mask)[1:]
399
+
400
+ centroidt0 = []
401
+ for label in labels:
402
+ cell_p = np.where(mask == label)
403
+ pix_ = len(cell_p[0])
404
+ centroidt0.append((np.append(np.mean(np.asarray(cell_p).T, axis=0),
405
+ (label, pix_)), (-1, -1, -1, -1)))
406
+ assert len(labels) > 0, f'no cell detected in 0th frame!!! cell num = {len(labels)}'
407
+ centroid.append(np.asarray(centroidt0))
408
+ num_cell_f.append(len(centroid[-1]))
409
+ theta = cfg.track.max_movenment ** 2 / 16
410
+
411
+ for i in range(len(flow_dir)):
412
+ mask = tifffile.imread(mask_dir[i+1])
413
+ labels = np.unique(mask)[1:]
414
+ centroidt = []
415
+ assert len(labels) > 0, f'no cell detected in {i}th frame!!! cell num = {len(labels)}'
416
+ # assert len(labels) == labels.max(), f"number and id of cells doesn't match"
417
+ flow = tifffile.imread(flow_dir[i])
418
+ flow = flow * (mask > 0)
419
+ for label in labels:
420
+ pix_ = len(np.where(mask == label)[0])
421
+ cell = np.where(mask == label)
422
+ m = np.append(np.mean(flow[:2, cell[0], cell[1]], axis=1), (0, 0))
423
+ centroidt.append((np.append(np.mean(np.asarray(cell).T, axis=0), (label, pix_)), m))
424
+ centroid.append(np.asarray(centroidt))
425
+ num_cell_f.append(len(centroid[-1]))
426
+ matrix = cal_trans_matrix(cfg, centroid, mask, flow, num_cell_f, theta, jitter_thr)
427
+ final_match = linear_solver(cfg, matrix, max_division=max_division, new_detect=new_detect)
428
+ match.append(final_match)
429
+ if run_id < max_run - 1:
430
+ q.put([run_id, num_cell_f[:-1], centroid[:-1], trans_matrix, match])
431
+ else:
432
+ q.put([run_id, num_cell_f, centroid, trans_matrix, match])
433
+ print(f'queue size: {q.qsize()}/{max_run}', f'------{run_id}th running end-------\n')
434
+
435
+
436
+ def run(cfg, q, mask_dir, flow_dir,
437
+ run_num, start_frame, num_f,
438
+ last_itv=10, jitter_thr=0.6, type='notequal'):
439
+ """
440
+ 多线程/多进程运行细胞追踪任务。
441
+
442
+ 根据 run_num 参数将任务分配到多个线程并行处理,支持两种分配方式:
443
+ - equal 分配:每个线程处理相近数量的文件
444
+ - 加权分配:按递增间隔分配,后续线程处理更多文件
445
+
446
+ Args:
447
+ cfg: 配置文件对象
448
+ q: 队列或数据队列
449
+ mask_dir: 掩码文件列表
450
+ flow_dir: 光流文件列表
451
+ run_num: 并行线程数量
452
+ start_frame: 起始帧编号
453
+ num_f: 总帧数
454
+ last_itv: 加权分配时第一个线程处理的帧数,默认10
455
+ jitter_thr: 抖动阈值,默认0.6
456
+ type: 分配类型,'equal' 为均分分配,其他值为加权分配,默认 'not equal'
457
+
458
+ Returns:
459
+ None
460
+ """
461
+ rs = []
462
+ if run_num <= 1:
463
+ get_track(cfg, q, mask_dir, flow_dir, 0, 1, jitter_thr) # 单线程处理
464
+ elif type == 'equal' or run_num < 10:# or run_num < 4:
465
+ # 计算每个线程需要处理的文件数量
466
+ interval = num_f // run_num + 1
467
+ if run_num > num_f:
468
+ run_num = num_f - 1
469
+ interval = 1
470
+ for i in range(run_num):
471
+ if int(i * interval) < len(mask_dir)-1:
472
+ args = (cfg, q,
473
+ mask_dir[int(i * interval):int(i * interval + interval + 1)],
474
+ flow_dir[int(i * interval):int(i * interval + interval)],
475
+ i,
476
+ run_num,
477
+ jitter_thr
478
+ )
479
+ r = multirun(target=get_track, args=args)
480
+ r.start()
481
+ rs.append(r)
482
+ else:
483
+ if run_num > num_f:
484
+ run_num = num_f - 1
485
+ itvs = np.array([1] * run_num).astype(int)
486
+ else:
487
+ itv = (2 * num_f / run_num - 2 * last_itv) / (run_num - 1)
488
+ itvs = [int(np.round(last_itv + i * itv)) for i in range(run_num)][::-1]
489
+
490
+ itvs[0] = int(itvs[0] + num_f - np.sum(itvs) - 1)
491
+ itvs = np.asarray(itvs)
492
+ print('file intervals:', itvs, f'sum: {np.sum(itvs)} + 1')
493
+ assert np.sum(itvs) == num_f-1, f'intervals sum {np.sum(itvs)} is not equal to cell num {num_f} - 1'
494
+ assert np.all(np.array(itvs) > 0), f'there is zero-interval '
495
+
496
+ for i in range(run_num):
497
+ print(f'file interval {i}:',
498
+ [start_frame+np.sum(itvs[:i]), start_frame+np.sum(itvs[:i+1])],
499
+ len(mask_dir[np.sum(itvs[:i]): np.sum(itvs[:i+1])]))
500
+ args = (cfg, q,
501
+ mask_dir[np.sum(itvs[:i]): np.sum(itvs[:i+1]) + 1],
502
+ flow_dir[np.sum(itvs[:i]): np.sum(itvs[:i+1])],
503
+ i,
504
+ run_num,
505
+ jitter_thr,
506
+ )
507
+ r = multirun(target=get_track, args=args)
508
+ r.start()
509
+ rs.append(r)
510
+ [r.join() for r in rs]
511
+ print(f'------all running ended-------\n')
512
+
513
+
514
+ def get_lineage(cfg, tracks, num_cell_f):
515
+ """
516
+ 根据细胞轨迹数据构建细胞谱系。
517
+
518
+ 遍历所有帧,从初始帧开始逐帧构建每个细胞的谱系链。
519
+ 支持细胞的持续存在、消失(标记为-1)和新出现的轨迹。
520
+
521
+ Args:
522
+ cfg: 配置对象
523
+ tracks: 轨迹列表,每帧包含该帧的细胞ID列表
524
+ num_cell_f: 每帧的细胞数量列表
525
+
526
+ Returns:
527
+ numpy.ndarray: 转置后的细胞谱系数组
528
+ """
529
+
530
+ for f, tr in enumerate(tracks):
531
+
532
+ # 提取当前轨迹中的细胞id
533
+ cells_t = set(tr)
534
+
535
+ assert len(tr) == num_cell_f[f + 1], \
536
+ f"number {len(tr)} and id {num_cell_f[f + 1]} don't match"
537
+
538
+ current_tracks = []
539
+ # 初始化当前迭代的轨迹列表
540
+ if f == 0:
541
+ all_tracks = [[i] for i in range(num_cell_f[0])]
542
+
543
+ # 遍历当前已有的轨迹
544
+ for track in all_tracks:
545
+
546
+ # 提取最后一个id
547
+ last_id = track[-1]
548
+
549
+ next_cells = np.where(np.array(tr) == last_id)[0]
550
+
551
+ # 如果id不在当前细胞集合中,则标记为-1(死亡、出视野)
552
+ if last_id not in cells_t:
553
+ current_tracks.append(track + [-1])
554
+ assert next_cells.size == 0
555
+ # continue
556
+ else:
557
+ # 为每个可能的下一个轨迹点连接到当前轨迹上
558
+ for next_cell in next_cells:
559
+ current_tracks.append(track + [next_cell])
560
+
561
+ for next_cell in np.where(np.array(tr) == num_cell_f[f])[0]:
562
+ current_tracks.append([-1] * (f + 1) + [next_cell])
563
+
564
+ # 注意:这里将current_tracks赋值回all_tracks,累积所有轨迹
565
+ all_tracks = current_tracks
566
+ track_list = np.array(all_tracks, dtype=np.int32)[:-1].T
567
+
568
+ assert check_track(track_list)
569
+
570
+ return np.array(all_tracks, dtype=np.int32).T
571
+
572
+
573
+ def cell_tracker(cfg,
574
+ mask_file,
575
+ flow_file,
576
+ run_num=1,
577
+ start_frame=0,
578
+ num_f=100,
579
+ last_itv=2,
580
+ jitter_thr=0.6,
581
+ ):
582
+ """
583
+ 细胞跟踪主程序入口。
584
+
585
+ 使用多线程方式处理细胞掩膜和光流数据,执行细胞跟踪并返回轨迹和质心结果。
586
+
587
+ Args:
588
+ cfg: 配置对象,包含跟踪算法的相关配置参数。
589
+ mask_file: 细胞掩膜文件路径。
590
+ flow_file: 光流文件路径。
591
+ run_num: 运行次数,默认为1。
592
+ start_frame: 起始帧索引,默认为0。
593
+ num_f: 处理的总帧数,默认为100。
594
+ last_itv: 最后时间间隔,默认为2。
595
+ jitter_thr: 抖动阈值,默认为0.6。
596
+
597
+ Returns:
598
+ tuple: 包含轨迹数组和质心数组的元组。
599
+ """
600
+ run_id, num_cell_f, centroid, trans_matrix, track = [], [], [], [], []
601
+ with Manager() as manager:
602
+ q = manager.Queue()
603
+ run(cfg, q, mask_file, flow_file, run_num, start_frame, num_f, last_itv, jitter_thr, type='not equal')
604
+ while not q.empty():
605
+ n, nc, cn, tm, tr = q.get()
606
+ run_id.append(n)
607
+ num_cell_f.append(nc)
608
+ centroid.append(cn)
609
+ # trans_matrix.append(tm)
610
+ track.append(tr)
611
+ num_cell_f = [num_cell_f[i] for i in sorted(range(len(run_id)), key=lambda k: run_id[k])]
612
+ num_cell_f = list(chain(*num_cell_f))
613
+ centroid = [centroid[i] for i in sorted(range(len(run_id)), key=lambda k: run_id[k])]
614
+ centroid = list(chain(*centroid))
615
+ # trans_matrix = [trans_matrix[i] for i in sorted(range(len(run_id)), key=lambda k: run_id[k])]
616
+ # trans_matrix = list(chain(*trans_matrix))
617
+ track_match = [track[i] for i in sorted(range(len(run_id)), key=lambda k: run_id[k])]
618
+ track_match = list(chain(*track_match))
619
+
620
+ # if method == 'greedy':
621
+ # import greedy
622
+ # track_cell = greedy.track_cell
623
+ # elif method == 'linear_solver':
624
+ # import linear_solver
625
+ # track_cell = linear_solver.track_cell
626
+ #
627
+ # else:
628
+ # raise ValueError("unsupported method")
629
+ np.save(join(cfg.result, 'centroid.npy'), np.array(centroid, dtype=object))
630
+ track = get_lineage(cfg, track_match, num_cell_f)
631
+
632
+ return track, centroid
633
+
634
+
635
+ def convert2id(cfg, track, centroid):
636
+ track_copy = np.zeros_like(track, dtype=np.int16) - 1
637
+ for ith, tr in enumerate(track):
638
+ for jth, t in enumerate(tr):
639
+ if track[ith, jth] != -1:
640
+ track_copy[ith, jth] = int(centroid[jth][track[ith, jth], 0, 2])
641
+ df = pd.DataFrame(data=track_copy.T, columns=None)
642
+ df.to_csv(join(cfg.result, 'track_id.csv'), index=False)
643
+
644
+
645
+ @hydra.main(config_path=abs_path('config'), version_base='1.3', config_name='tracker')
646
+ def tracker(cfg: DictConfig):
647
+ """
648
+ 主程序入口,用于启动多线程处理过程。
649
+ Args:
650
+ method: 选择的跟踪方法,可以是'greedy', 'linear_solver'之一
651
+ mask_dir: 细胞掩膜文件夹路径,
652
+ """
653
+ cfg = get_cfg(cfg, type='track')
654
+ start = time.time()
655
+
656
+ run_num = cfg.track.run_num
657
+ last_itv = cfg.track.last_itv
658
+
659
+ start_frame = cfg.dataloader.start_frame
660
+ num_f = cfg.dataloader.num_frame
661
+
662
+ if cfg.track.load_track:
663
+ track = pd.read_csv(abs_path(cfg.track.track_file)).to_numpy()
664
+ centroid = np.load(abs_path(cfg.track.centroid_file), allow_pickle=True)
665
+ else:
666
+ track, centroid = cell_tracker(cfg,
667
+ mask_file=cfg.mask_dir,
668
+ flow_file=cfg.flow_dir,
669
+ run_num=run_num,
670
+ start_frame=start_frame,
671
+ num_f=num_f,
672
+ last_itv=last_itv,
673
+ jitter_thr=cfg.track.jitter_thr,
674
+ )
675
+ if num_f > 15 and cfg.track.post_pro:
676
+ track = merging_and_pruning(cfg, track.T, centroid)
677
+ df = pd.DataFrame(data=track, columns=None)
678
+ df.to_csv(join(cfg.result, 'track_results.csv'), index=False)
679
+ convert2id(cfg, track.T, centroid)
680
+
681
+ print(now.strftime("%Y-%m-%d %H:%M:%S") + f'{num_f}-frame time cost:', int(time.time() - start), 's')
682
+
683
+
684
+ if __name__ == "__main__":
685
+
686
+ # tracker()
687
+ track = pd.read_csv(abs_path(r"D:\Data\Haishan\experiment\ipsc\result00013\2025-08-25track_linear_solver_merged_and_pruned.csv")).to_numpy()
688
+ centroid = np.load(abs_path(r"D:\Data\Haishan\experiment\ipsc\result00013\2025-08-25-centroid.npy"), allow_pickle=True)
689
+ track = merging_and_pruning(track.T, centroid, min_length=10, prune_leaf=True, thr_merge=30)
690
+ df = pd.DataFrame(data=track, columns=None)
691
+ df.to_csv(join(r'D:\Data\Haishan\track\result', now.strftime("%Y-%m-%d") +
692
+ 'track_linear_solver_merged_and_pruned.csv'), index=False)