File size: 34,088 Bytes
03e863f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
"""Flow probe system: single probe, multi-probe, and branching support.

The probe is advected by the same velocity field as the particles. Its trajectory
is recorded and can be analyzed for brain region transitions.

Features:
  - Single probe: click to place, follows mean flow
  - Multi-probe: initialize N probes in a local neighborhood
  - Branching: when MDN components are highly uncertain (50/50 split), spawn
    a ghost probe that follows the dominant alternative component
  - Live region highlighting in RED (ghost highlights more transparent)
"""

import numpy as np
import vtk


class FlowProbe:
    """A single probe that follows the flow field."""

    def __init__(self, ren: vtk.vtkRenderer, amin: np.ndarray, amax: np.ndarray,
                 color=(0.0, 1.0, 0.3), opacity=0.9, ghost=False, label=""):
        self.ren = ren
        self.amin = amin.astype(np.float32)
        self.amax = amax.astype(np.float32)
        self.active = False
        self.position = None
        self.path: list[np.ndarray] = []
        self.speeds: list[float] = []
        self.raw_field_mags: list[float] = []  # field magnitude (independent of speed scale)
        self._max_path_len = 50000
        self._mesh_overlay = None
        self._highlighted_regions: set[str] = set()
        self._highlight_actors: dict[str, vtk.vtkActor] = {}
        self._check_interval = 30
        self._step_counter = 0
        self.current_regions: set[str] = set()
        self._boundary_check = None
        self._on_region_change = None
        self.ghost = ghost
        self.label = label
        self._color = color
        self.steps_alive = 0
        self._stuck_counter = 0
        self._stuck_threshold = 150  # steps with near-zero movement before warning
        self._stuck_warned = False
        self._stuck_eps = 1e-6  # minimum displacement per step
        # Debounce: require N consecutive detections before entering, N misses before leaving
        self._debounce_enter = 3   # consecutive checks to confirm entry
        self._debounce_leave = 5   # consecutive misses to confirm exit
        self._region_hit_count: dict[str, int] = {}   # key -> consecutive hit count
        self._region_miss_count: dict[str, int] = {}  # key -> consecutive miss count
        self._nearest_max_dist_mm = 5.0  # hard limit for nearest-region snapping (mm)

        diag = float(np.linalg.norm(amax - amin))
        marker_alpha = 0.4 if ghost else 1.0
        trail_alpha = 0.35 if ghost else 0.9

        # Probe marker (sphere)
        self._sphere = vtk.vtkSphereSource()
        self._sphere.SetRadius(0.012 * diag if not ghost else 0.008 * diag)
        self._sphere.SetThetaResolution(16)
        self._sphere.SetPhiResolution(16)
        mapper = vtk.vtkPolyDataMapper()
        mapper.SetInputConnection(self._sphere.GetOutputPort())
        self.marker_actor = vtk.vtkActor()
        self.marker_actor.SetMapper(mapper)
        self.marker_actor.GetProperty().SetColor(*color)
        self.marker_actor.GetProperty().SetOpacity(marker_alpha)
        self.marker_actor.GetProperty().LightingOff()
        self.marker_actor.VisibilityOff()
        ren.AddActor(self.marker_actor)

        # Trail line
        self._trail_points = vtk.vtkPoints()
        self._trail_cells = vtk.vtkCellArray()
        self._trail_pd = vtk.vtkPolyData()
        self._trail_pd.SetPoints(self._trail_points)
        self._trail_pd.SetLines(self._trail_cells)
        trail_mapper = vtk.vtkPolyDataMapper()
        trail_mapper.SetInputData(self._trail_pd)
        self.trail_actor = vtk.vtkActor()
        self.trail_actor.SetMapper(trail_mapper)
        self.trail_actor.GetProperty().SetColor(*color)
        self.trail_actor.GetProperty().SetLineWidth(3.0 if not ghost else 2.0)
        self.trail_actor.GetProperty().SetOpacity(trail_alpha)
        self.trail_actor.GetProperty().LightingOff()
        self.trail_actor.VisibilityOff()
        ren.AddActor(self.trail_actor)

    def set_mesh_overlay(self, mesh_overlay):
        self._mesh_overlay = mesh_overlay

    def set_boundary_check(self, fn):
        self._boundary_check = fn

    def set_on_region_change(self, fn):
        self._on_region_change = fn

    def place(self, position: np.ndarray):
        """Place probe at position and start recording."""
        self.position = np.array(position, dtype=np.float32).ravel()[:3]

        # Boundary validation
        if self._boundary_check is not None and not self._boundary_check(self.position):
            found = False
            diag = float(np.linalg.norm(self.amax - self.amin))
            for scale in [0.01, 0.02, 0.05, 0.1, 0.2]:
                for _ in range(50):
                    offset = np.random.randn(3).astype(np.float32) * scale * diag
                    candidate = np.clip(self.position + offset, self.amin, self.amax)
                    if self._boundary_check(candidate):
                        self.position = candidate
                        found = True
                        break
                if found:
                    break
            if not found:
                self.position = (self.amin + self.amax) / 2.0
                print(f"[probe{self.label}] Could not find valid position, using domain center")

        self.path = [self.position.copy()]
        self.speeds = [0.0]
        self.raw_field_mags = [0.0]
        self.active = True
        self._step_counter = 0
        self.steps_alive = 0
        self._stuck_counter = 0
        self._stuck_warned = False

        # Reset trail
        self._trail_points = vtk.vtkPoints()
        self._trail_cells = vtk.vtkCellArray()
        self._trail_points.InsertNextPoint(*self.position.tolist())
        self._trail_pd.SetPoints(self._trail_points)
        self._trail_pd.SetLines(self._trail_cells)
        self._trail_pd.Modified()

        self.marker_actor.VisibilityOn()
        self.marker_actor.SetPosition(*self.position.tolist())
        self.trail_actor.VisibilityOn()

        self._update_region_highlights()
        tag = " (ghost)" if self.ghost else ""
        print(f"[probe{self.label}{tag}] placed at "
              f"({self.position[0]:.2f}, {self.position[1]:.2f}, {self.position[2]:.2f})")

    def step(self, sampler, dt_step: float):
        """Advect one step using sampler."""
        if not self.active or self.position is None:
            return
        V = sampler.sample_vec(self.position[None, :])
        velocity = V[0]
        raw_mag = float(np.linalg.norm(velocity))

        new_pos = self.position + velocity * dt_step
        new_pos = np.clip(new_pos, self.amin, self.amax)

        # Boundary constraint
        if self._boundary_check is not None and not self._boundary_check(new_pos):
            half_pos = self.position + velocity * dt_step * 0.5
            half_pos = np.clip(half_pos, self.amin, self.amax)
            if self._boundary_check(half_pos):
                new_pos = half_pos
            else:
                return  # hit boundary

        self.position = new_pos.astype(np.float32)
        speed = float(np.linalg.norm(self.position - self.path[-1])) if self.path else 0.0

        # Stuck / weak flow detection
        if speed < self._stuck_eps and raw_mag < self._stuck_eps:
            self._stuck_counter += 1
            if self._stuck_counter >= self._stuck_threshold and not self._stuck_warned:
                self._stuck_warned = True
                tag = f" (ghost)" if self.ghost else ""
                print(f"\n[probe{self.label}{tag}] Flow is very weak or has ended here. "
                      f"The probe is stuck at ({self.position[0]:.1f}, {self.position[1]:.1f}, {self.position[2]:.1f}).")
                print(f"[probe{self.label}{tag}] Try placing the probe deeper in the brain "
                      f"where flow is stronger (press 'c' to clear, then 'g' + click).\n")
        else:
            self._stuck_counter = 0

        if len(self.path) < self._max_path_len:
            self.path.append(self.position.copy())
            self.speeds.append(speed)
            self.raw_field_mags.append(raw_mag)

        self.marker_actor.SetPosition(*self.position.tolist())

        # Append to trail
        idx = self._trail_points.InsertNextPoint(*self.position.tolist())
        if idx > 0:
            self._trail_cells.InsertNextCell(2)
            self._trail_cells.InsertCellPoint(idx - 1)
            self._trail_cells.InsertCellPoint(idx)
        self._trail_points.Modified()
        self._trail_cells.Modified()
        self._trail_pd.Modified()

        # Periodic region check
        self._step_counter += 1
        self.steps_alive += 1
        if self._step_counter % self._check_interval == 0:
            self._update_region_highlights()

    def _point_in_bbox(self, point, bounds):
        """Fast bounding-box containment check. bounds is (xmin,xmax,ymin,ymax,zmin,zmax)."""
        return (bounds[0] <= point[0] <= bounds[1] and
                bounds[2] <= point[1] <= bounds[3] and
                bounds[4] <= point[2] <= bounds[5])

    def _update_region_highlights(self):
        if self._mesh_overlay is None or self.position is None:
            return

        # --- Raw detection (what region key is at the probe right now?) ---
        raw_key = None
        if hasattr(self._mesh_overlay, 'get_region_at_point'):
            raw_key = self._mesh_overlay.get_region_at_point(self.position)
            if raw_key is None and hasattr(self._mesh_overlay, 'find_nearest_region'):
                raw_key = self._mesh_overlay.find_nearest_region(
                    self.position, search_radius=2,
                    max_distance_mm=self._nearest_max_dist_mm)
        else:
            for key in self._mesh_overlay.get_all_region_keys():
                if hasattr(self._mesh_overlay, 'fast_point_in_mesh'):
                    if self._mesh_overlay.fast_point_in_mesh(self.position, key):
                        raw_key = key
                        break
                elif self._mesh_overlay.point_in_mesh(self.position, key):
                    raw_key = key
                    break

        raw_detected = {raw_key} if raw_key else set()

        # --- Debounce: require consecutive detections to enter, consecutive misses to leave ---
        # Update hit/miss counters for detected key
        for key in raw_detected:
            self._region_hit_count[key] = self._region_hit_count.get(key, 0) + 1
            self._region_miss_count.pop(key, None)

        # Update miss counters for keys that were NOT detected this tick
        for key in list(self._region_hit_count.keys()):
            if key not in raw_detected:
                self._region_miss_count[key] = self._region_miss_count.get(key, 0) + 1
                self._region_hit_count[key] = 0

        # Determine stable set: regions that passed the entry threshold
        # and have not yet exceeded the leave threshold
        new_regions = set()
        for key in set(list(self._region_hit_count.keys()) +
                       list(self._highlighted_regions)):
            hits = self._region_hit_count.get(key, 0)
            misses = self._region_miss_count.get(key, 0)
            if key in self._highlighted_regions:
                # Already highlighted — keep it unless misses exceed threshold
                if misses < self._debounce_leave:
                    new_regions.add(key)
            else:
                # Not yet highlighted — add if hits exceed entry threshold
                if hits >= self._debounce_enter:
                    new_regions.add(key)

        # Clean up stale counters
        for key in list(self._region_miss_count.keys()):
            if self._region_miss_count[key] > self._debounce_leave + 2:
                self._region_miss_count.pop(key, None)
                self._region_hit_count.pop(key, None)

        # --- Apply changes (enter/leave) ---
        left = self._highlighted_regions - new_regions
        left_names = []
        for key in left:
            if key in self._highlight_actors:
                try:
                    self.ren.RemoveActor(self._highlight_actors[key])
                except Exception:
                    pass
                del self._highlight_actors[key]
            name = self._mesh_overlay.get_region_name(key)
            if hasattr(self._mesh_overlay, 'get_hemisphere_label'):
                hemi = self._mesh_overlay.get_hemisphere_label(key, self.position)
                if hemi:
                    name = f"{name} ({hemi})"
            left_names.append(name)
            tag = " [branch]" if self.ghost else ""
            print(f"[probe{self.label}] LEFT: {name}{tag}")

        entered = new_regions - self._highlighted_regions
        entered_names = []
        highlight_opacity = 0.12 if self.ghost else 0.25
        for key in entered:
            poly = None
            if hasattr(self._mesh_overlay, 'get_hemisphere_polydata'):
                poly = self._mesh_overlay.get_hemisphere_polydata(key, self.position)
            if poly is None:
                poly = self._mesh_overlay.get_polydata(key)
            if poly is not None:
                mapper = vtk.vtkPolyDataMapper()
                mapper.SetInputData(poly)
                actor = vtk.vtkActor()
                actor.SetMapper(mapper)
                actor.GetProperty().SetColor(1.0, 0.6, 0.0)  # orange (all regions)
                actor.GetProperty().SetOpacity(highlight_opacity)
                actor.GetProperty().LightingOff()
                self.ren.AddActor(actor)
                self._highlight_actors[key] = actor
            name = self._mesh_overlay.get_region_name(key)
            if hasattr(self._mesh_overlay, 'get_hemisphere_label'):
                hemi = self._mesh_overlay.get_hemisphere_label(key, self.position)
                if hemi:
                    name = f"{name} ({hemi})"
            entered_names.append(name)
            tag = " [branch]" if self.ghost else ""
            # Detailed entry log with position context
            pos_detail = ""
            if self.position is not None and self._mesh_overlay is not None:
                try:
                    center = self._mesh_overlay.get_mesh_center(key)
                    bounds = self._mesh_overlay.get_mesh_bounds(key)
                    if center is not None and bounds is not None:
                        extent = [bounds[1]-bounds[0], bounds[3]-bounds[2],
                                  bounds[5]-bounds[4]]
                        char_size = sum(extent) / 3.0
                        dist = float(np.linalg.norm(self.position - center))
                        depth = max(0.0, 1.0 - min(dist / (char_size * 0.5), 1.0))
                        pos_parts = []
                        diff = self.position - center
                        if abs(diff[2]) > extent[2] * 0.15:
                            pos_parts.append("dorsal" if diff[2] > 0 else "ventral")
                        if abs(diff[0]) > extent[0] * 0.15:
                            pos_parts.append("lateral-R" if diff[0] > 0 else "lateral-L")
                        if abs(diff[1]) > extent[1] * 0.15:
                            pos_parts.append("anterior" if diff[1] > 0 else "posterior")
                        pos_str = "-".join(pos_parts) if pos_parts else "central"
                        pos_detail = f"  [{pos_str}, depth={depth:.0%}]"
                except Exception:
                    pass
            print(f"[probe{self.label}{tag}] ENTERED: {name}{pos_detail}")

        self._highlighted_regions = new_regions
        self.current_regions = new_regions

        # --- Extra parcellation subregion highlighting ---
        # Extra parcellation subregions also get orange, same as main regions.
        if (self._mesh_overlay is not None and
                hasattr(self._mesh_overlay, '_extra') and
                self._mesh_overlay._extra is not None and
                self.position is not None and new_regions):
            try:
                hier = self._mesh_overlay.get_hierarchical_regions_at_point(self.position)
                sub = hier.get("subregion")
                sub_mesh = hier.get("subregion_mesh")
                sub_key = f"_extra_{sub['label_id']}" if sub else None

                # Remove old subregion highlight if changed
                old_sub_key = getattr(self, '_current_subregion_key', None)
                if old_sub_key and old_sub_key != sub_key:
                    if old_sub_key in self._highlight_actors:
                        try:
                            self.ren.RemoveActor(self._highlight_actors[old_sub_key])
                        except Exception:
                            pass
                        del self._highlight_actors[old_sub_key]

                # Add new subregion highlight (orange, same as all regions)
                if sub_key and sub_mesh and sub_key not in self._highlight_actors:
                    display_mesh = sub_mesh
                    try:
                        bounds = [0.0] * 6
                        sub_mesh.GetBounds(bounds)
                        x_extent = bounds[1] - bounds[0]
                        if x_extent > 10.0:
                            x_mid = (bounds[0] + bounds[1]) / 2.0
                            plane = vtk.vtkPlane()
                            plane.SetOrigin(x_mid, 0, 0)
                            if self.position[0] >= x_mid:
                                plane.SetNormal(1, 0, 0)
                            else:
                                plane.SetNormal(-1, 0, 0)
                            clipper = vtk.vtkClipPolyData()
                            clipper.SetInputData(sub_mesh)
                            clipper.SetClipFunction(plane)
                            clipper.SetInsideOut(False)
                            clipper.Update()
                            clipped = clipper.GetOutput()
                            if clipped and clipped.GetNumberOfPoints() > 0:
                                display_mesh = clipped
                    except Exception:
                        pass

                    mapper = vtk.vtkPolyDataMapper()
                    mapper.SetInputData(display_mesh)
                    actor = vtk.vtkActor()
                    actor.SetMapper(mapper)
                    actor.GetProperty().SetColor(1.0, 0.6, 0.0)  # orange
                    actor.GetProperty().SetOpacity(0.3)
                    actor.GetProperty().LightingOff()
                    self.ren.AddActor(actor)
                    self._highlight_actors[sub_key] = actor
                    tag = " [branch]" if self.ghost else ""
                    hemi_str = ""
                    if self.position is not None and self.position[0] >= 0:
                        hemi_str = " (right hemisphere)"
                    elif self.position is not None:
                        hemi_str = " (left hemisphere)"
                    print(f"[probe{self.label}{tag}] SUBREGION: {sub['name']}{hemi_str}")

                # Hide coarser main-atlas regions when we have a finer subregion
                if sub_key and sub_key in self._highlight_actors:
                    hidden = set()
                    for rkey in new_regions:
                        if rkey in self._highlight_actors and not rkey.startswith("_extra_"):
                            self._highlight_actors[rkey].VisibilityOff()
                            hidden.add(rkey)
                    self._hidden_orange_keys = hidden
                elif not sub_key:
                    for rkey in getattr(self, '_hidden_orange_keys', set()):
                        if rkey in self._highlight_actors:
                            self._highlight_actors[rkey].VisibilityOn()
                    self._hidden_orange_keys = set()

                self._current_subregion_key = sub_key
            except Exception:
                pass

        # --- Red hotspot: clip mesh near probe position ---
        # Instead of coloring an entire region red, clip the mesh surface
        # within a sphere around the probe and show that patch in red.
        self._update_hotspot()

        if self._on_region_change and (entered_names or left_names):
            self._on_region_change(entered_names, left_names,
                                   is_ghost=self.ghost, label=self.label)

    def _update_hotspot(self):
        """Color highlighted region meshes with a distance-based heatmap.

        Vertices near the probe are red, fading smoothly to orange further away.
        Uses per-vertex RGBA scalars — no clipping, no extra actors.
        """
        from vtkmodules.util.numpy_support import vtk_to_numpy, numpy_to_vtk

        if self.position is None or not self._highlight_actors:
            return

        probe_pos = self.position.astype(np.float64)
        fade_radius = 15.0  # mm — distance over which red fades to orange

        for key, actor in self._highlight_actors.items():
            mapper = actor.GetMapper()
            if mapper is None:
                continue
            poly = mapper.GetInput()
            if poly is None or poly.GetNumberOfPoints() < 3:
                continue

            pts_vtk = poly.GetPoints()
            if pts_vtk is None:
                continue
            verts = vtk_to_numpy(pts_vtk.GetData()).astype(np.float64)

            # Distance from each vertex to probe
            dists = np.linalg.norm(verts - probe_pos, axis=1)
            t = np.clip(dists / fade_radius, 0.0, 1.0)  # 0=at probe, 1=far

            # Color gradient: red (1,0,0) at probe → orange (1,0.6,0) far
            r = np.full(len(t), 255, np.uint8)
            g = (t * 0.6 * 255).astype(np.uint8)
            b = np.zeros(len(t), np.uint8)

            # Alpha: brighter near probe, dimmer far away
            base_opacity = 0.45 if not self.ghost else 0.2
            far_opacity = 0.2 if not self.ghost else 0.08
            alpha_f = far_opacity + (base_opacity - far_opacity) * (1.0 - t)
            a = (np.clip(alpha_f, 0.0, 1.0) * 255).astype(np.uint8)

            rgba = np.column_stack([r, g, b, a])
            scalars = numpy_to_vtk(rgba, deep=True)
            scalars.SetNumberOfComponents(4)
            scalars.SetName("HeatmapColors")
            poly.GetPointData().SetScalars(scalars)

            mapper.SetColorModeToDirectScalars()
            mapper.SetScalarModeToUsePointData()
            mapper.ScalarVisibilityOn()
            actor.GetProperty().LightingOff()
            # Override the flat color — let scalars drive everything
            actor.GetProperty().SetOpacity(1.0)
            poly.Modified()

    def clear(self):
        self.active = False
        self.position = None
        self.path = []
        self.speeds = []
        self.raw_field_mags = []
        self._step_counter = 0
        self.steps_alive = 0
        self.marker_actor.VisibilityOff()
        self.trail_actor.VisibilityOff()
        self._trail_points = vtk.vtkPoints()
        self._trail_cells = vtk.vtkCellArray()
        self._trail_pd.SetPoints(self._trail_points)
        self._trail_pd.SetLines(self._trail_cells)
        self._trail_pd.Modified()
        for key, actor in self._highlight_actors.items():
            try:
                self.ren.RemoveActor(actor)
            except Exception:
                pass
        self._highlight_actors.clear()
        self._highlighted_regions.clear()
        self.current_regions.clear()

    def destroy(self):
        """Remove all actors from renderer."""
        self.clear()
        try:
            self.ren.RemoveActor(self.marker_actor)
        except Exception:
            pass
        try:
            self.ren.RemoveActor(self.trail_actor)
        except Exception:
            pass

    def get_path_array(self) -> np.ndarray:
        if not self.path:
            return np.zeros((0, 3), np.float32)
        return np.array(self.path, dtype=np.float32)

    def get_speeds_array(self) -> np.ndarray:
        if not self.speeds:
            return np.zeros(0, np.float32)
        return np.array(self.speeds, dtype=np.float32)

    def get_field_mags_array(self) -> np.ndarray:
        if not self.raw_field_mags:
            return np.zeros(0, np.float32)
        return np.array(self.raw_field_mags, dtype=np.float32)


class ProbeSystem:
    """Manages single/multi probes and branching behavior.

    Modes:
      - single: one probe following mean flow
      - multi: N probes in local neighborhood, all following mean flow
      - branching: when PI uncertainty is high, spawn ghost probes

    For state-propagation mode, defaults to multi(4) + branching.
    """

    def __init__(self, ren: vtk.vtkRenderer, win: vtk.vtkRenderWindow,
                 amin: np.ndarray, amax: np.ndarray):
        self.ren = ren
        self.win = win
        self.amin = amin.astype(np.float32)
        self.amax = amax.astype(np.float32)
        self.probes: list[FlowProbe] = []
        self.ghost_probes: list[FlowProbe] = []
        self._mesh_overlay = None
        self._boundary_check = None
        self._on_region_change = None
        self._branching_enabled = False
        self._multi_count = 1
        self._branch_threshold = 0.35  # max ratio between top 2 PI components
        self._branch_min_pi = 0.25     # min weight of 2nd component
        self._branch_check_interval = 50
        self._branch_step_counter = 0
        self._ghost_color = (0.4, 0.7, 1.0)  # pale blue for ghosts
        self._pi_field = None
        self._mus_samplers = None

    def set_mesh_overlay(self, overlay):
        self._mesh_overlay = overlay

    def set_boundary_check(self, fn):
        self._boundary_check = fn

    def set_on_region_change(self, fn):
        self._on_region_change = fn

    def set_branching(self, enabled: bool, pi_field=None, mus_samplers=None):
        """Enable/disable branching.

        Args:
            enabled: toggle branching
            pi_field: the PI weight grid (G,G,G,K) numpy array
            mus_samplers: list of TriLinearSampler for each component
        """
        self._branching_enabled = enabled
        self._pi_field = pi_field
        self._mus_samplers = mus_samplers
        print(f"[probe-system] branching {'ON' if enabled else 'OFF'}")

    def set_multi_count(self, n: int):
        self._multi_count = max(1, n)
        print(f"[probe-system] multi-probe count: {self._multi_count}")

    @property
    def active(self) -> bool:
        return any(p.active for p in self.probes)

    @property
    def path(self):
        """Return path of first probe (for backward compat)."""
        return self.probes[0].path if self.probes else []

    def place(self, position: np.ndarray):
        """Place probe(s) at position."""
        self.clear()
        diag = float(np.linalg.norm(self.amax - self.amin))
        jitter_scale = 0.04 * diag

        for i in range(self._multi_count):
            if i == 0:
                pos = position.copy()
                label = "" if self._multi_count == 1 else f"#{i+1}"
            else:
                offset = np.random.randn(3).astype(np.float32) * jitter_scale
                pos = np.clip(position + offset, self.amin, self.amax)
                label = f"#{i+1}"

            probe = FlowProbe(self.ren, self.amin, self.amax,
                              color=(0.0, 1.0, 0.3), ghost=False, label=label)
            if self._mesh_overlay:
                probe.set_mesh_overlay(self._mesh_overlay)
            if self._boundary_check:
                probe.set_boundary_check(self._boundary_check)
            if self._on_region_change:
                probe.set_on_region_change(self._on_region_change)
            probe.place(pos)
            self.probes.append(probe)

    def step(self, sampler, dt_step: float, pi_sampler=None):
        """Step all probes."""
        for p in self.probes:
            if p.active:
                p.step(sampler, dt_step)

        for g in self.ghost_probes:
            if g.active:
                # Ghost probes follow their specific component sampler
                comp_sampler = getattr(g, '_comp_sampler', sampler)
                g.step(comp_sampler, dt_step)

        # Prune ghost probes that have been alive > 200 steps but did not diverge
        diag = float(np.linalg.norm(self.amax - self.amin))
        prune_dist = 0.03 * diag
        to_remove = []
        for g in self.ghost_probes:
            if not g.active or g.position is None:
                continue
            if g.steps_alive > 200:
                for p in self.probes:
                    if not p.active or p.position is None:
                        continue
                    dist = float(np.linalg.norm(g.position - p.position))
                    if dist < prune_dist:
                        print("[branch] pruned ghost - did not diverge")
                        to_remove.append(g)
                        break
        for g in to_remove:
            g.destroy()
            self.ghost_probes.remove(g)

        # Check for branching
        if self._branching_enabled:
            self._branch_step_counter += 1
            if self._branch_step_counter % self._branch_check_interval == 0:
                self._check_branching(sampler, dt_step)

    def _check_branching(self, mean_sampler, dt_step):
        """Check if any active probe is at a high-uncertainty location and branch."""
        if self._pi_field is None or self._mus_samplers is None:
            return
        if not self.probes:
            return

        from .field_loader import TriLinearSampler

        for p in self.probes:
            if not p.active or p.position is None:
                continue
            # Sample PI weights at probe position
            pos = p.position[None, :]
            pi_vals = mean_sampler.sample_vec(pos)  # dummy, we need PI
            # Actually sample PI field directly
            try:
                K = self._pi_field.shape[-1]
                weights = np.zeros(K, np.float32)
                for k in range(K):
                    pi_k = self._pi_field[..., k:k+1]
                    # Use sampler to interpolate
                    w = mean_sampler.sample_scalar(
                        self._pi_field[..., k],
                        p.position[None, :]
                    )
                    weights[k] = float(w[0])
            except Exception:
                continue

            # Normalize
            ws = weights.sum()
            if ws <= 0:
                continue
            weights /= ws

            # Sort to find top 2
            sorted_idx = np.argsort(weights)[::-1]
            w1 = weights[sorted_idx[0]]
            w2 = weights[sorted_idx[1]] if len(sorted_idx) > 1 else 0.0

            # Check if uncertain enough
            if w2 < self._branch_min_pi:
                continue
            ratio = w2 / max(w1, 1e-9)
            if ratio < self._branch_threshold:
                continue

            # Check we haven't already branched from near this position
            already_branched = False
            for g in self.ghost_probes:
                if g.active and g.path:
                    d = np.linalg.norm(g.path[0] - p.position)
                    if d < 0.02 * float(np.linalg.norm(self.amax - self.amin)):
                        already_branched = True
                        break
            if already_branched:
                continue

            # Branch! Create ghost probe following 2nd component
            comp_idx = int(sorted_idx[1])
            if comp_idx >= len(self._mus_samplers):
                continue

            ghost = FlowProbe(self.ren, self.amin, self.amax,
                              color=self._ghost_color, ghost=True,
                              label=f"~branch(comp{comp_idx+1})")
            if self._mesh_overlay:
                ghost.set_mesh_overlay(self._mesh_overlay)
            if self._boundary_check:
                ghost.set_boundary_check(self._boundary_check)
            if self._on_region_change:
                ghost.set_on_region_change(self._on_region_change)
            ghost._comp_sampler = self._mus_samplers[comp_idx]
            ghost.place(p.position.copy())
            self.ghost_probes.append(ghost)

            print(f"[branch] SPLIT at ({p.position[0]:.2f}, {p.position[1]:.2f}, "
                  f"{p.position[2]:.2f}): comp{sorted_idx[0]+1}={w1:.2f} vs "
                  f"comp{comp_idx+1}={w2:.2f}")
            if self._on_region_change:
                self._on_region_change(
                    [f"BRANCH: comp{comp_idx+1} (weight={w2:.2f})"], [],
                    is_ghost=True, label="branch"
                )

    def clear(self):
        """Clear all probes."""
        for p in self.probes:
            p.destroy()
        for g in self.ghost_probes:
            g.destroy()
        self.probes.clear()
        self.ghost_probes.clear()
        self._branch_step_counter = 0

    def freeze(self):
        """Stop all probes from moving (keep path and highlights intact)."""
        for p in self.probes:
            p.active = False
        for g in self.ghost_probes:
            g.active = False
        print("[probe-system] probes frozen")

    def get_all_probes(self) -> list[FlowProbe]:
        """Return all probes (main + ghost) for analysis."""
        return self.probes + self.ghost_probes

    def get_primary_probe(self) -> FlowProbe | None:
        """Return the first active probe."""
        for p in self.probes:
            if p.active:
                return p
        return None

    def get_all_paths(self) -> list[tuple[np.ndarray, bool, str]]:
        """Return [(path_array, is_ghost, label), ...] for all probes."""
        result = []
        for p in self.probes:
            if p.path:
                result.append((p.get_path_array(), False, p.label))
        for g in self.ghost_probes:
            if g.path:
                result.append((g.get_path_array(), True, g.label))
        return result