Reuse shared features for paired 2D and 3D output
#1
by darthceltic85 - opened
- hecate.py +45 -2
- hecate_shared_output.py +203 -0
- 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)
|