Reuse shared features for paired 2D and 3D output

#1
Files changed (3) hide show
  1. hecate.py +45 -2
  2. hecate_shared_output.py +203 -0
  3. test_hecate_shared_output.py +208 -0
hecate.py CHANGED
@@ -423,7 +423,50 @@ def main():
423
  if abs(args.spacing_um/model.sampling_um-1) > .02:
424
  parser.error(f'Resample the render to {model.sampling_um} um first (no automatic resampling)')
425
  volume = open_volume(args.input,args.array)
426
- if output is not None:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
427
  output.parent.mkdir(parents=True,exist_ok=True)
428
  with tempfile.TemporaryDirectory(prefix='hecate-') as tmp:
429
  result = np.memmap(Path(tmp)/'prediction',mode='w+',dtype='uint8',shape=volume.shape[-2:])
@@ -442,7 +485,7 @@ def main():
442
  if output is not None:
443
  output.with_suffix('.json').write_text(json.dumps(metadata,indent=2)+'\n')
444
  print(output)
445
- if output3d is not None:
446
  import zarr
447
  from numcodecs import Blosc
448
  output3d.parent.mkdir(parents=True,exist_ok=True)
 
423
  if abs(args.spacing_um/model.sampling_um-1) > .02:
424
  parser.error(f'Resample the render to {model.sampling_um} um first (no automatic resampling)')
425
  volume = open_volume(args.input,args.array)
426
+ if output is not None and output3d is not None:
427
+ from hecate_shared_output import predict_both
428
+ import zarr
429
+ from numcodecs import Blosc
430
+ output.parent.mkdir(parents=True,exist_ok=True)
431
+ output3d.parent.mkdir(parents=True,exist_ok=True)
432
+ partial = output3d.with_name(output3d.name+'.partial')
433
+ if partial.exists():
434
+ parser.error(f'Incomplete output exists: {partial}; remove it or choose another output')
435
+ options = {'zarr_format':2} if int(zarr.__version__.split('.')[0])>=3 else {}
436
+ result3 = zarr.open_array(str(partial),mode='w',shape=volume.shape,
437
+ chunks=(1,256,256),dtype='uint8',fill_value=0,
438
+ compressor=Blosc(cname='zstd',clevel=3,shuffle=Blosc.BITSHUFFLE),**options)
439
+ metadata = {
440
+ 'checkpoint':Path(args.checkpoint).name,'sampling_um':args.spacing_um,
441
+ 'shape_yx':list(volume.shape[-2:]),'reverse_z':args.reverse,
442
+ 'central_depth':model.patch_size[0],'precision':args.precision,
443
+ 'stride':args.stride or model.patch_size[1]//2,
444
+ 'probabilities':'uint8 = round(255 * probability); no min-max scaling'
445
+ }
446
+ with tempfile.TemporaryDirectory(prefix='hecate-') as tmp:
447
+ result2 = np.memmap(Path(tmp)/'prediction',mode='w+',dtype='uint8',
448
+ shape=volume.shape[-2:])
449
+ predict_both(model,volume,result2,result3,reverse=args.reverse,
450
+ stride=args.stride,batch_size=args.batch_size,
451
+ precision=args.precision)
452
+ result2.flush()
453
+ partial_png = output.with_suffix('.partial')
454
+ write_png(result2,partial_png)
455
+ partial_png.replace(output)
456
+ del result2
457
+ depth = volume.shape[0]
458
+ d = model.patch_size[0]
459
+ start = depth//2-d//2
460
+ if args.reverse:
461
+ start = depth-start-d
462
+ result3.attrs.update({**metadata,'axes':['z','y','x'],
463
+ 'spacing_um':[args.spacing_um]*3,'z_order':'same as input',
464
+ 'evaluated_z_interval':[start+model.margin,start+d-model.margin],
465
+ 'outside_evaluated_z':'zero; not evaluated',
466
+ 'background':'empty CT columns zeroed; no intensity threshold'})
467
+ partial.rename(output3d)
468
+ print(output3d)
469
+ elif output is not None:
470
  output.parent.mkdir(parents=True,exist_ok=True)
471
  with tempfile.TemporaryDirectory(prefix='hecate-') as tmp:
472
  result = np.memmap(Path(tmp)/'prediction',mode='w+',dtype='uint8',shape=volume.shape[-2:])
 
485
  if output is not None:
486
  output.with_suffix('.json').write_text(json.dumps(metadata,indent=2)+'\n')
487
  print(output)
488
+ if output3d is not None and output is None:
489
  import zarr
490
  from numcodecs import Blosc
491
  output3d.parent.mkdir(parents=True,exist_ok=True)
hecate_shared_output.py ADDED
@@ -0,0 +1,203 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Shared-feature tiled Hecate prediction for paired 2D and 3D outputs.
2
+
3
+ This keeps the pinned Hecate provider and checkpoint untouched. It computes the
4
+ existing 2D attention head and 3D head from one ``features_logits`` call per tile,
5
+ then blends both outputs with the same Hann weights. It is an implementation
6
+ helper, not a provider qualification or authorization to run on target material.
7
+ """
8
+ from __future__ import annotations
9
+
10
+ from collections.abc import Callable
11
+
12
+ import numpy as np
13
+
14
+
15
+ def shared_logits(model, image, valid):
16
+ """Return the pinned Hecate 2D and 3D logits using one shared feature pass.
17
+
18
+ ``image`` is normalized B1ZYX float input; ``valid`` is the corresponding
19
+ boolean support mask. The projections mirror ``Hecate.forward`` and
20
+ ``Hecate.forward_3d`` in the pinned provider, including 2.4 um depth margins.
21
+ """
22
+ import torch
23
+ import torch.nn.functional as F
24
+
25
+ if image.ndim != 5 or image.shape[1] != 1:
26
+ raise ValueError("image must have B1ZYX shape")
27
+ if valid is not None and tuple(valid.shape) != tuple(image.shape):
28
+ raise ValueError("valid support must match image shape")
29
+
30
+ features, logits = model.features_logits(image)
31
+ depth = image.shape[2]
32
+ margin = int(model.margin)
33
+ inner_depth = depth - 2 * margin
34
+ if inner_depth <= 0:
35
+ raise ValueError("model margin leaves no evaluated depth")
36
+
37
+ scores = model.canonical.decoder.depth_collapse.attn_conv(features).float()
38
+ half = scores.shape[2] // 2
39
+ z = (torch.arange(scores.shape[2], device=scores.device, dtype=torch.float32) - half) / half
40
+ scores = scores + model.depth_coordinate_scale * z.view(1, 1, -1, 1, 1)
41
+ if valid is None:
42
+ support = torch.ones_like(logits, dtype=torch.bool)
43
+ else:
44
+ support = F.max_pool3d(
45
+ valid[:, :, margin:depth - margin].float(),
46
+ (1, model.xy_stride, model.xy_stride),
47
+ ).bool()
48
+ weights = scores.masked_fill(~support, -1e4).softmax(2) * support
49
+ weights = weights / weights.sum(2, keepdim=True).clamp_min(1e-8)
50
+ probability_2d = F.interpolate(
51
+ (weights * logits).sum(2).sigmoid(), size=image.shape[-2:],
52
+ mode="bilinear", align_corners=False,
53
+ )
54
+ eps = torch.finfo(torch.float32).eps
55
+ logits_2d = torch.logit(probability_2d.clamp(eps, 1 - eps))
56
+
57
+ logits_3d = F.interpolate(
58
+ logits, size=(inner_depth, *image.shape[-2:]),
59
+ mode="trilinear", align_corners=False,
60
+ )
61
+ if margin:
62
+ logits_3d = F.pad(logits_3d, (0, 0, 0, 0, margin, margin), value=-20.0)
63
+ return logits_2d.float(), logits_3d.float()
64
+
65
+
66
+ def _axis_starts(length: int, patch: int, stride: int) -> list[int]:
67
+ # Kept byte-for-byte equivalent in behavior to the pinned helper.
68
+ return list(range(0, max(0, length - patch), stride)) + [max(0, length - patch)]
69
+
70
+
71
+ def predict_both(
72
+ model,
73
+ volume,
74
+ output_2d: np.ndarray,
75
+ output_3d: np.ndarray,
76
+ *,
77
+ reverse: bool = False,
78
+ stride: int | None = None,
79
+ batch_size: int = 1,
80
+ precision: str = "fp32",
81
+ before_batch: Callable[[], None] | None = None,
82
+ progress: Callable[[int, int], None] | None = None,
83
+ ) -> dict:
84
+ """Stream one input into both uint8 outputs, sharing Hecate's feature pass.
85
+
86
+ Accumulators are disk-backed temporary arrays. Inputs and outputs follow
87
+ the pinned provider's ZYX/HW layouts, central-depth selection, reverse order,
88
+ Hann overlap weighting, and uint8 probability conversion.
89
+ """
90
+ import tempfile
91
+ from contextlib import ExitStack
92
+
93
+ import torch
94
+
95
+ if len(volume.shape) != 3 or np.dtype(volume.dtype) != np.dtype("uint8"):
96
+ raise ValueError("volume must be an unnormalized uint8 Z,Y,X array")
97
+ depth, height, width = map(int, volume.shape)
98
+ model_depth, patch, _ = model.patch_size
99
+ stride = patch // 2 if stride is None else int(stride)
100
+ batch_size = int(batch_size)
101
+ if depth < model_depth or not 0 < stride <= patch or batch_size < 1:
102
+ raise ValueError("insufficient depth, invalid stride, or invalid batch size")
103
+ if tuple(output_2d.shape) != (height, width) or output_2d.dtype != np.uint8:
104
+ raise ValueError("output_2d must be uint8 with shape (Y,X)")
105
+ if tuple(output_3d.shape) != (depth, height, width) or output_3d.dtype != np.uint8:
106
+ raise ValueError("output_3d must be uint8 with shape (Z,Y,X)")
107
+
108
+ device = next(model.parameters()).device
109
+ if precision not in ("fp32", "bf16") or (precision == "bf16" and device.type != "cuda"):
110
+ raise ValueError("bf16 requires CUDA; otherwise use fp32")
111
+
112
+ z0 = depth // 2 - model_depth // 2
113
+ selected_z0 = depth - z0 - model_depth if reverse else z0
114
+ zs = slice(selected_z0, selected_z0 + model_depth)
115
+ window_t = torch.hann_window(patch, periodic=False)
116
+ window_t = window_t[:, None] * window_t[None, :]
117
+ window_t = (window_t / window_t.max().clamp_min(torch.finfo(torch.float32).eps)).clamp_min(0.001)
118
+ window_np = window_t.numpy()
119
+ window_t = window_t.to(device)
120
+ origins = [(y, x) for y in _axis_starts(height, patch, stride)
121
+ for x in _axis_starts(width, patch, stride)]
122
+ output_3d[:] = 0
123
+ blank_tiles = 0
124
+
125
+ with tempfile.TemporaryDirectory(prefix="argus-hecate-shared-") as tmp, ExitStack() as stack:
126
+ def accumulator(name: str, shape: tuple[int, ...]):
127
+ array = np.memmap(f"{tmp}/{name}", mode="w+", dtype="float32", shape=shape)
128
+ stack.callback(array._mmap.close)
129
+ return array
130
+
131
+ numerator_2d = accumulator("sum2d", (height, width))
132
+ numerator_3d = accumulator("sum3d", (model_depth, height, width))
133
+ denominator = accumulator("weight", (height, width))
134
+ numerator_2d[:] = 0
135
+ numerator_3d[:] = 0
136
+ denominator[:] = 0
137
+
138
+ for start in range(0, len(origins), batch_size):
139
+ if before_batch is not None:
140
+ before_batch()
141
+ active, patches = [], []
142
+ for y, x in origins[start:start + batch_size]:
143
+ ph, pw = min(patch, height - y), min(patch, width - x)
144
+ denominator[y:y + ph, x:x + pw] += window_np[:ph, :pw]
145
+ raw = np.asarray(volume[zs, y:y + ph, x:x + pw])
146
+ if reverse:
147
+ raw = raw[::-1]
148
+ if not raw.any():
149
+ blank_tiles += 1
150
+ continue
151
+ padded = np.zeros((model_depth, patch, patch), dtype=np.uint8)
152
+ padded[:, :ph, :pw] = raw
153
+ patches.append(padded)
154
+ active.append((y, x, ph, pw))
155
+
156
+ if patches:
157
+ raw_batch = np.stack(patches)
158
+ support = torch.from_numpy(raw_batch).to(device).any(1)
159
+ normalized = raw_batch.astype(np.float32)
160
+ normalized /= model.divisor
161
+ image = torch.from_numpy(normalized[:, None]).to(device)
162
+ valid = support[:, None, None].expand(-1, 1, model_depth, -1, -1)
163
+ with torch.inference_mode(), torch.autocast(
164
+ device_type=device.type, dtype=torch.bfloat16, enabled=precision == "bf16"
165
+ ):
166
+ logits_2d, logits_3d = shared_logits(model, image, valid)
167
+ probability_2d = logits_2d.float().sigmoid()[:, 0]
168
+ probability_3d = logits_3d.float().sigmoid()[:, 0]
169
+ weighted_2d = (probability_2d * support * window_t).cpu().numpy()
170
+ weighted_3d = (probability_3d * support[:, None] * window_t).cpu().numpy()
171
+ for pred2, pred3, (y, x, ph, pw) in zip(weighted_2d, weighted_3d, active):
172
+ numerator_2d[y:y + ph, x:x + pw] += pred2[:ph, :pw]
173
+ numerator_3d[:, y:y + ph, x:x + pw] += pred3[:, :ph, :pw]
174
+ if progress is not None:
175
+ progress(min(start + batch_size, len(origins)), len(origins))
176
+
177
+ if model.margin:
178
+ numerator_3d[:model.margin] = 0
179
+ numerator_3d[-model.margin:] = 0
180
+ for y in range(0, height, 128):
181
+ for x in range(0, width, 256):
182
+ y1, x1 = min(height, y + 128), min(width, x + 256)
183
+ den = denominator[y:y1, x:x1]
184
+ if not np.isfinite(den).all() or np.any(den <= 0):
185
+ raise FloatingPointError("non-finite or uncovered output pixels")
186
+ prob2 = numerator_2d[y:y1, x:x1] / den
187
+ prob3 = numerator_3d[:, y:y1, x:x1] / den[None]
188
+ if not np.isfinite(prob2).all() or not np.isfinite(prob3).all():
189
+ raise FloatingPointError("non-finite inference output")
190
+ values2 = np.rint(prob2.clip(0, 1) * 255).astype(np.uint8)
191
+ values3 = np.rint(prob3.clip(0, 1) * 255).astype(np.uint8)
192
+ output_2d[y:y1, x:x1] = values2
193
+ output_3d[zs, y:y1, x:x1] = values3[::-1] if reverse else values3
194
+
195
+ numerator_2d.flush()
196
+ numerator_3d.flush()
197
+ denominator.flush()
198
+ del numerator_2d, numerator_3d, denominator
199
+
200
+ return {"tile_count": len(origins), "blank_tiles_skipped": blank_tiles,
201
+ "evaluated_z_interval": [selected_z0 + model.margin,
202
+ selected_z0 + model_depth - model.margin],
203
+ "reverse": bool(reverse), "stride": stride, "precision": precision}
test_hecate_shared_output.py ADDED
@@ -0,0 +1,208 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Parity tests for Hecate's shared 2D/3D feature projection."""
2
+ from __future__ import annotations
3
+
4
+ import torch
5
+ import torch.nn.functional as F
6
+ import numpy as np
7
+ from torch import nn
8
+
9
+ from hecate_shared_output import predict_both, shared_logits
10
+
11
+
12
+ class _Collapse(nn.Module):
13
+ def __init__(self):
14
+ super().__init__()
15
+ self.attn_conv = nn.Conv3d(1, 1, 1, bias=True)
16
+
17
+
18
+ class _Decoder(nn.Module):
19
+ def __init__(self):
20
+ super().__init__()
21
+ self.depth_collapse = _Collapse()
22
+
23
+
24
+ class _Canonical(nn.Module):
25
+ def __init__(self):
26
+ super().__init__()
27
+ self.decoder = _Decoder()
28
+
29
+
30
+ class _FakeHecate(nn.Module):
31
+ def __init__(self, margin=0, xy_stride=1):
32
+ super().__init__()
33
+ self.anchor = nn.Parameter(torch.zeros(()))
34
+ self.canonical = _Canonical()
35
+ self.depth_coordinate_scale = nn.Parameter(torch.tensor(0.17))
36
+ self.margin = margin
37
+ self.xy_stride = xy_stride
38
+ self.divisor = 255.0
39
+ self.patch_size = (6, 8, 8)
40
+ self.feature_calls = 0
41
+
42
+ def features_logits(self, image):
43
+ self.feature_calls += 1
44
+ features = image[:, :, self.margin:image.shape[2] - self.margin]
45
+ if self.xy_stride > 1:
46
+ features = F.avg_pool3d(
47
+ features, (1, self.xy_stride, self.xy_stride),
48
+ stride=(1, self.xy_stride, self.xy_stride),
49
+ )
50
+ logits = features * 0.73 + 0.11
51
+ return features, logits
52
+
53
+ def forward(self, image, valid=None):
54
+ features, logits = self.features_logits(image)
55
+ scores = self.canonical.decoder.depth_collapse.attn_conv(features).float()
56
+ half = scores.shape[2] // 2
57
+ z = (torch.arange(scores.shape[2], device=scores.device, dtype=torch.float32) - half) / half
58
+ scores = scores + self.depth_coordinate_scale * z.view(1, 1, -1, 1, 1)
59
+ if valid is None:
60
+ support = torch.ones_like(logits, dtype=torch.bool)
61
+ else:
62
+ support = F.max_pool3d(
63
+ valid[:, :, self.margin:image.shape[2] - self.margin].float(),
64
+ (1, self.xy_stride, self.xy_stride),
65
+ ).bool()
66
+ weights = scores.masked_fill(~support, -1e4).softmax(2) * support
67
+ weights = weights / weights.sum(2, keepdim=True).clamp_min(1e-8)
68
+ probability = F.interpolate(
69
+ (weights * logits).sum(2).sigmoid(), size=image.shape[-2:],
70
+ mode="bilinear", align_corners=False,
71
+ )
72
+ eps = torch.finfo(torch.float32).eps
73
+ return torch.logit(probability.clamp(eps, 1 - eps))
74
+
75
+ def forward_3d(self, image):
76
+ _, logits = self.features_logits(image)
77
+ full = F.interpolate(
78
+ logits, size=(image.shape[2] - 2 * self.margin, *image.shape[-2:]),
79
+ mode="trilinear", align_corners=False,
80
+ )
81
+ return F.pad(full, (0, 0, 0, 0, self.margin, self.margin), value=-20.0)
82
+
83
+
84
+ def test_shared_heads_match_separate_pinned_head_math_without_duplicate_features():
85
+ torch.manual_seed(7)
86
+ model = _FakeHecate().eval()
87
+ image = torch.rand(2, 1, 6, 12, 10)
88
+ valid = torch.ones_like(image, dtype=torch.bool)
89
+ valid[:, :, :, :2, :3] = False
90
+
91
+ expected_2d = model(image, valid)
92
+ expected_3d = model.forward_3d(image)
93
+ separate_calls = model.feature_calls
94
+ assert separate_calls == 2
95
+
96
+ actual_2d, actual_3d = shared_logits(model, image, valid)
97
+ assert model.feature_calls - separate_calls == 1
98
+ assert torch.equal(actual_2d, expected_2d)
99
+ assert torch.equal(actual_3d, expected_3d)
100
+
101
+
102
+ def test_shared_heads_preserve_two_point_four_micron_margin_padding():
103
+ torch.manual_seed(11)
104
+ model = _FakeHecate(margin=1, xy_stride=2).eval()
105
+ image = torch.rand(1, 1, 6, 12, 10)
106
+ valid = torch.ones_like(image, dtype=torch.bool)
107
+ valid[:, :, :, :2, :3] = False
108
+
109
+ expected_2d = model(image, valid)
110
+ expected_3d = model.forward_3d(image)
111
+ actual_2d, actual_3d = shared_logits(model, image, valid)
112
+ assert torch.equal(actual_2d, expected_2d)
113
+ assert torch.equal(actual_3d, expected_3d)
114
+ assert torch.all(actual_3d[:, :, 0] == -20.0)
115
+ assert torch.all(actual_3d[:, :, -1] == -20.0)
116
+
117
+
118
+ @torch.inference_mode()
119
+ def _stock_predict(model, volume, output, *, reverse, stride, three_d):
120
+ """Small CPU reference matching the pinned provider's separate-head tiler."""
121
+ d, patch, _ = model.patch_size
122
+ depth, height, width = volume.shape
123
+ z0 = depth // 2 - d // 2
124
+ zs = slice(depth - z0 - d, depth - z0) if reverse else slice(z0, z0 + d)
125
+ window = torch.hann_window(patch, periodic=False)
126
+ window = window[:, None] * window[None, :]
127
+ window = (window / window.max().clamp_min(torch.finfo(torch.float32).eps)).clamp_min(0.001)
128
+ numerator = torch.zeros((d, height, width) if three_d else (height, width), dtype=torch.float32)
129
+ denominator = torch.zeros((height, width), dtype=torch.float32)
130
+ starts_y = list(range(0, max(0, height - patch), stride)) + [max(0, height - patch)]
131
+ starts_x = list(range(0, max(0, width - patch), stride)) + [max(0, width - patch)]
132
+ for y in starts_y:
133
+ for x in starts_x:
134
+ ph, pw = min(patch, height - y), min(patch, width - x)
135
+ denominator[y:y + ph, x:x + pw] += window[:ph, :pw]
136
+ raw = np.asarray(volume[zs, y:y + ph, x:x + pw])
137
+ if reverse:
138
+ raw = raw[::-1]
139
+ if not raw.any():
140
+ continue
141
+ padded = np.zeros((d, patch, patch), dtype=np.uint8)
142
+ padded[:, :ph, :pw] = raw
143
+ support = torch.from_numpy(padded[None]).any(1)
144
+ image = torch.from_numpy(padded.astype(np.float32)[None, None] / model.divisor)
145
+ valid = support[:, None, None].expand(-1, 1, d, -1, -1)
146
+ logits = model.forward_3d(image) if three_d else model(image, valid)
147
+ probability = logits.float().sigmoid()[:, 0]
148
+ if three_d:
149
+ weighted = probability * support[:, None] * window
150
+ numerator[:, y:y + ph, x:x + pw] += weighted[0, :, :ph, :pw]
151
+ else:
152
+ weighted = probability * support * window
153
+ numerator[y:y + ph, x:x + pw] += weighted[0, :ph, :pw]
154
+ if three_d and model.margin:
155
+ numerator[:model.margin] = 0
156
+ numerator[-model.margin:] = 0
157
+ probability = numerator / (denominator[None] if three_d else denominator)
158
+ values = torch.round(probability.clamp(0, 1) * 255).to(torch.uint8).numpy()
159
+ if three_d:
160
+ output[zs] = values[::-1] if reverse else values
161
+ else:
162
+ output[:] = values
163
+
164
+
165
+ def test_tiled_outputs_match_two_stock_passes_and_skip_blank_tiles():
166
+ torch.manual_seed(19)
167
+ model = _FakeHecate().eval()
168
+ volume = np.zeros((6, 13, 15), dtype=np.uint8)
169
+ volume[:, :7, :7] = np.random.default_rng(19).integers(1, 255, size=(6, 7, 7), dtype=np.uint8)
170
+ expected_2d = np.zeros((13, 15), dtype=np.uint8)
171
+ expected_3d = np.zeros((6, 13, 15), dtype=np.uint8)
172
+ _stock_predict(model, volume, expected_2d, reverse=False, stride=4, three_d=False)
173
+ _stock_predict(model, volume, expected_3d, reverse=False, stride=4, three_d=True)
174
+
175
+ actual_2d = np.zeros_like(expected_2d)
176
+ actual_3d = np.zeros_like(expected_3d)
177
+ result = predict_both(model, volume, actual_2d, actual_3d, stride=4, batch_size=1)
178
+ assert result["tile_count"] == 9
179
+ assert result["blank_tiles_skipped"] > 0
180
+ assert np.array_equal(actual_2d, expected_2d)
181
+ assert np.array_equal(actual_3d, expected_3d)
182
+
183
+
184
+ def test_center_depth_crop_and_reverse_order_match_two_stock_passes():
185
+ """Longer inputs select the centered model-depth window in either direction."""
186
+ for reverse in (False, True):
187
+ torch.manual_seed(31)
188
+ model = _FakeHecate().eval()
189
+ rng = np.random.default_rng(31)
190
+ volume = rng.integers(0, 256, size=(9, 11, 13), dtype=np.uint8)
191
+ expected_2d = np.zeros((11, 13), dtype=np.uint8)
192
+ expected_3d = np.zeros((9, 11, 13), dtype=np.uint8)
193
+ _stock_predict(model, volume, expected_2d, reverse=reverse, stride=4, three_d=False)
194
+ _stock_predict(model, volume, expected_3d, reverse=reverse, stride=4, three_d=True)
195
+
196
+ actual_2d = np.zeros_like(expected_2d)
197
+ actual_3d = np.zeros_like(expected_3d)
198
+ result = predict_both(
199
+ model, volume, actual_2d, actual_3d, reverse=reverse, stride=4, batch_size=1,
200
+ )
201
+ z0 = volume.shape[0] // 2 - model.patch_size[0] // 2
202
+ selected_z0 = volume.shape[0] - z0 - model.patch_size[0] if reverse else z0
203
+ assert result["tile_count"] == 6
204
+ assert np.array_equal(actual_2d, expected_2d)
205
+ assert np.array_equal(actual_3d, expected_3d)
206
+ assert np.all(actual_3d[:selected_z0] == 0)
207
+ assert np.all(actual_3d[selected_z0 + model.patch_size[0]:] == 0)
208
+ assert np.any(actual_3d[selected_z0:selected_z0 + model.patch_size[0]] > 0)