Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion workers/annotations/illumination_correction/Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ COPY ./workers/annotations/illumination_correction/pipeline.py /
LABEL isUPennContrastWorker="" \
isGPUWorker="true" \
isAnnotationWorker="" \
workerVersion="1.0.3" \
workerVersion="1.1.0" \
interfaceName="Stitch Refinement + Illumination Correction" \
interfaceCategory="Image Processing" \
description="Refines composite translations and corrects raw-tile illumination; existing annotations are not shifted" \
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -42,8 +42,11 @@ The implementation follows the supplied raw-tile v7 model:
- Robust Huber IRLS fit of all nonconstant order-5 2-D DCT terms
- Ridge penalty `4 × (1 + order_y² + order_x²)²`
- Six IRLS iterations
- Interior anchoring: a fixed quadratic prior pulling the DCT correction toward zero at camera pixels that no used overlap observes, weighted at 0.3× the overlap sample count (evaluated on at most 4,096 interior pixels) and never Huber-reweighted
- Optional overlap-derived per-position gains with ridge 8 and a 1.10-fold cap

Overlaps only observe the tile margins (about 38% of each tile on a 49-tile, ~10%-overlap grid), so the overlap terms carry no information about the tile interior. Without anchoring, the low-order DCT correction extrapolated into that unobserved interior and made the field too peaked, leaving every corrected channel 2.5–6.4% darker at tile centres than at tile edges. That residual is most visible in low-contrast background channels such as YFP when display contrast is stretched. Anchoring keeps the correction near the log-median base field where there is no overlap evidence; on that dataset it brought the centre/edge residual to within ±0.2% in every channel with unchanged seam agreement, and the result is insensitive to the weight between 0.03 and 0.3. Diagnostics record `overlap_coverage_fraction`, `interior_anchor_weight`, and `interior_anchor_samples`.

The flat-field reference is the ND2 Z-stack home plane at T=0 (or the middle Z plane when the metadata has no valid home index). Training reads those T=0 P×Z camera frames one at a time and never materializes the full time series; correction still streams and preserves every time point. Fields are fitted at 128×128 and bicubically expanded to the raw camera dimensions. Corrected data is clipped only when written back to lossless uint16.

Each scratch TIFF page records explicit `IndexC`, `IndexT`, and `IndexZ` frame metadata so `large_image` groups positions while preserving the source channel/Z/time axes. The image pins the matching bundled `pyvips`/`libvips` wheel used by `large_image` and asserts the native binding version during its build; conversion is serialized to avoid oversubscribing the container. The image is GPU-queue routed (`isGPUWorker=true`) so this long-running compute does not occupy the CPU queue that serves interactive worker-interface requests. This label controls Celery placement; the implementation does not require CUDA.
Expand Down
2 changes: 1 addition & 1 deletion workers/annotations/illumination_correction/entrypoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@


WORKER_NAME = "Stitch Refinement + Illumination Correction"
WORKER_VERSION = "1.0.3"
WORKER_VERSION = "1.1.0"
ALGORITHM_OPTIONS = (
"Overlap DCT + tile gains (recommended)",
"Overlap DCT",
Expand Down
74 changes: 68 additions & 6 deletions workers/annotations/illumination_correction/illumination.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,15 @@
VALIDATION_LOW_PERCENTILE = 10.0
VALIDATION_HIGH_PERCENTILE = 90.0
MAX_DCT_SAMPLES = 250_000
# Overlaps only observe the tile margins, so the DCT correction is
# unconstrained in the interior and extrapolates there (measured: tile
# centres over-corrected by 2.5-6.4% on a 49-tile ND2). A fixed quadratic
# prior pulling the correction toward zero at un-overlapped pixels keeps it
# near the log-median base field where no overlap evidence exists. The prior
# is not Huber-reweighted, and its total weight is this fraction of the
# overlap rows' unit weight.
INTERIOR_ANCHOR_WEIGHT = 0.3
MAX_INTERIOR_ANCHOR_SAMPLES = 4096


@dataclass(frozen=True)
Expand Down Expand Up @@ -114,7 +123,11 @@ def robust_ridge(
response: np.ndarray,
penalty: np.ndarray,
iterations: int = OVERLAP_DCT_IRLS_ITERATIONS,
prior: np.ndarray | None = None,
) -> tuple[np.ndarray, dict[str, float]]:
"""Huber IRLS ridge fit; ``prior`` is a fixed quadratic penalty matrix
that is never reweighted, unlike the ``design`` observations."""
regularizer = np.diag(penalty) if prior is None else np.diag(penalty) + prior
weights = np.ones(response.shape[0], dtype=np.float64)
coefficient = np.zeros(design.shape[1], dtype=np.float64)
residual = response.copy()
Expand All @@ -123,7 +136,7 @@ def robust_ridge(
root_weight = np.sqrt(weights)
weighted_design = design * root_weight[:, None]
weighted_response = response * root_weight
system = weighted_design.T @ weighted_design + np.diag(penalty)
system = weighted_design.T @ weighted_design + regularizer
target = weighted_design.T @ weighted_response
coefficient = np.linalg.solve(system, target)
residual = response - design @ coefficient
Expand Down Expand Up @@ -228,6 +241,44 @@ def _overlap_design(
return design, response.astype(np.float64)


def overlap_coverage(
shape: tuple[int, int],
measurements: Sequence[PairMeasurement],
raw_shape: tuple[int, int],
) -> np.ndarray:
"""Camera pixels observed by at least one overlap, in either tile role."""
covered = np.zeros(shape, dtype=bool)
for measurement in measurements:
shift_x, shift_y = _scaled_shift(measurement, shape, raw_shape)
slices = aligned_slices(shape, shift_x, shift_y)
if slices is not None:
covered[slices[0]] = True
covered[slices[1]] = True
return covered


def _interior_anchor_design(
covered: np.ndarray,
terms: Sequence[tuple[int, int]],
overlap_rows: int,
) -> np.ndarray | None:
height, width = covered.shape
interior_y, interior_x = np.nonzero(~covered)
if interior_y.size == 0:
return None
if interior_y.size > MAX_INTERIOR_ANCHOR_SAMPLES:
selection = np.linspace(
0, interior_y.size - 1, MAX_INTERIOR_ANCHOR_SAMPLES, dtype=np.int64
)
interior_y = interior_y[selection]
interior_x = interior_x[selection]
coordinates = np.column_stack(
(interior_y / max(height - 1, 1), interior_x / max(width - 1, 1))
)
scale = math.sqrt(INTERIOR_ANCHOR_WEIGHT * overlap_rows / interior_y.size)
return dct_basis(coordinates, terms) * scale


def overlap_chunk_mask(
shape: tuple[int, int], axis: str, fit: bool
) -> np.ndarray:
Expand Down Expand Up @@ -354,13 +405,13 @@ def fit_overlap_dct(
]
design_parts = []
response_parts = []
used_pairs = 0
used = []
for measurement in accepted:
part = _overlap_design(base_corrected, measurement, terms, raw_shape)
if part is not None:
design_parts.append(part[0])
response_parts.append(part[1])
used_pairs += 1
used.append(measurement)
if not design_parts:
raise ValueError("confident pairs contained no usable overlap pixels")
design = np.concatenate(design_parts, axis=0)
Expand All @@ -371,14 +422,22 @@ def fit_overlap_dct(
)
design = design[selection]
response = response[selection]
overlap_samples = int(response.size)
covered = overlap_coverage(values.shape[-2:], used, raw_shape)
anchor = _interior_anchor_design(covered, terms, overlap_samples)
penalty = np.asarray(
[
OVERLAP_DCT_RIDGE * (1.0 + order_y**2 + order_x**2) ** 2
for order_y, order_x in terms
],
dtype=np.float64,
)
coefficient, fit_diagnostics = robust_ridge(design, response, penalty)
coefficient, fit_diagnostics = robust_ridge(
design,
response,
penalty,
prior=None if anchor is None else anchor.T @ anchor,
)
grid_y, grid_x = np.meshgrid(
np.linspace(0.0, 1.0, values.shape[-2]),
np.linspace(0.0, 1.0, values.shape[-1]),
Expand Down Expand Up @@ -407,12 +466,15 @@ def fit_overlap_dct(
"base_method": "log_median",
"training_tiles": int(values.shape[0]),
"confident_pairs": len(accepted),
"used_pairs": used_pairs,
"used_pairs": len(used),
"dct_order": OVERLAP_DCT_ORDER,
"dct_terms": len(terms),
"ridge": OVERLAP_DCT_RIDGE,
"irls_iterations": OVERLAP_DCT_IRLS_ITERATIONS,
"dct_samples": int(response.size),
"dct_samples": overlap_samples,
"overlap_coverage_fraction": float(np.mean(covered)),
"interior_anchor_weight": INTERIOR_ANCHOR_WEIGHT,
"interior_anchor_samples": 0 if anchor is None else int(anchor.shape[0]),
"delta_log_range": float(np.max(delta) - np.min(delta)),
"flat_min": float(np.min(flat)),
"flat_max": float(np.max(flat)),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -195,7 +195,7 @@ def test_compute_refines_fits_all_channels_converts_and_uploads() -> None:
assert fit.call_args.kwargs["adaptive_tile_gains"] is True
metadata = upload.call_args.args[-1]
assert metadata["tool"] == "Stitch Refinement + Illumination Correction"
assert metadata["worker_version"] == "1.0.3"
assert metadata["worker_version"] == "1.1.0"
assert metadata["refinement"]["pairs_matched"] == 1
assert metadata["parameters"]["refinement_channel_name"] == "DAPI"
assert metadata["source"]["original_nd2_item_id"] == "source-item"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -127,3 +127,61 @@ def test_overlap_dct_can_disable_adaptive_tile_gains() -> None:

np.testing.assert_array_equal(model.gains, np.ones(2, dtype=np.float32))
assert model.diagnostics["method"] == "overlap_dct"


def test_overlap_dct_does_not_overcorrect_unobserved_tile_interior() -> None:
# Narrow (~10%) overlaps leave the tile centre unobserved by any pair, as
# in real ND2 grids. Without interior anchoring the DCT correction
# extrapolated there and the fitted field came out 7-9% too peaked.
rng = np.random.default_rng(0)
grid, tile_size, pitch = 5, 128, 116
size = pitch * (grid - 1) + tile_size
scene = (
2000.0
+ 300.0 * ndimage.gaussian_filter(rng.normal(size=(size, size)), 1.5)
+ 8000.0 * ndimage.gaussian_filter(rng.normal(size=(size, size)), 25)
).astype(np.float32)
coordinate = np.linspace(-1.0, 1.0, tile_size, dtype=np.float32)
true_flat = np.exp(-0.3 * (coordinate[:, None] ** 2 + coordinate[None, :] ** 2))
true_flat /= np.mean(true_flat)

tiles, stages, origins = [], [], []
for row in range(grid):
order = range(grid) if row % 2 == 0 else range(grid - 1, -1, -1)
for column in order:
x0, y0 = column * pitch, row * pitch
tiles.append(scene[y0 : y0 + tile_size, x0 : x0 + tile_size] * true_flat)
stages.append((column * 100.0, row * 100.0))
origins.append((x0, y0))
origins = np.asarray(origins)
measurements = []
for edge in build_adjacency(np.asarray(stages)):
shift = origins[edge.second] - origins[edge.first]
measurements.append(
PairMeasurement(
first=edge.first,
second=edge.second,
axis=edge.axis,
predicted_shift_x=int(shift[0]),
predicted_shift_y=int(shift[1]),
shift_x=int(shift[0]),
shift_y=int(shift[1]),
ncc=0.99,
accepted=True,
)
)

model = fit_overlap_dct(
np.stack(tiles).astype(np.float32),
measurements,
adaptive_tile_gains=False,
)

ratio = model.flatfield / true_flat
margin = np.concatenate(
[ratio[:8].ravel(), ratio[-8:].ravel(), ratio[:, :8].ravel(), ratio[:, -8:].ravel()]
)
centre_bias = float(np.mean(ratio[40:88, 40:88]) / np.mean(margin) - 1.0)
assert abs(centre_bias) < 0.025
assert 0.0 < model.diagnostics["overlap_coverage_fraction"] < 0.5
assert model.diagnostics["interior_anchor_samples"] > 0
Loading