"""Auditable image-quality screening before segmentation, using raw intensities."""
from __future__ import annotations
import csv
import io
import json
import math
import os
import tempfile
from datetime import datetime, timezone
from pathlib import Path
from .logging_util import _spacr_home
DEFAULTS = {
'image_qc_mode': 'off',
'image_qc_channels': [],
'image_qc_min_focus': {},
'image_qc_max_saturation': {},
'image_qc_saturation_level': {},
'image_qc_max_nonfinite': 0.0,
'image_qc_classifier': False,
'image_qc_classifier_model': None,
'image_qc_classifier_labels': None,
'image_qc_classifier_threshold': 0.5,
}
REPORT = 'qc/image_quality.json'
[docs]
def quality_policy(settings):
"""Validate a saved per-channel policy without inspecting any image.
:param settings: Mask settings containing image_qc_* options.
:returns: JSON-serializable policy; empty channel list means every channel.
:raises ValueError: unsupported mode, invalid channel or invalid threshold.
"""
policy = {key: settings.get(key, value) for key, value in DEFAULTS.items()}
if policy['image_qc_mode'] not in ('off', 'report', 'exclude'):
raise ValueError('image_qc_mode must be off, report or exclude')
channels = policy['image_qc_channels']
if not isinstance(channels, (list, tuple)) or any(
isinstance(value, bool) or not isinstance(value, int) or value < 0 for value in channels):
raise ValueError('image_qc_channels must contain nonnegative channel indices')
policy['image_qc_channels'] = list(dict.fromkeys(channels))
for key in ('image_qc_min_focus', 'image_qc_max_saturation', 'image_qc_saturation_level'):
values = policy[key]
if not isinstance(values, dict):
raise ValueError(f'{key} must map channel indices to thresholds')
converted = {}
for channel, value in values.items():
if isinstance(channel, bool) or not str(channel).isdigit():
raise ValueError(f'{key}: channel indices must be nonnegative integers')
value = float(value)
if not math.isfinite(value) or value < 0 or (
key == 'image_qc_max_saturation' and value > 1) or (
key == 'image_qc_saturation_level' and value == 0):
raise ValueError(f'{key}: invalid threshold for channel {channel}')
converted[str(int(channel))] = value
policy[key] = converted
limit = float(policy['image_qc_max_nonfinite'])
if not math.isfinite(limit) or not 0 <= limit <= 1:
raise ValueError('image_qc_max_nonfinite must be between 0 and 1')
policy['image_qc_max_nonfinite'] = limit
if not isinstance(policy['image_qc_classifier'], bool):
raise ValueError('image_qc_classifier must be True or False')
for key in ('image_qc_classifier_model', 'image_qc_classifier_labels'):
value = policy[key]
if value is None or (isinstance(value, str) and not value.strip()):
policy[key] = None
elif isinstance(value, (str, os.PathLike)):
policy[key] = os.fspath(value)
else:
raise ValueError(f'{key} must be a file path or blank')
threshold = float(policy['image_qc_classifier_threshold'])
if not math.isfinite(threshold) or not 0 < threshold < 1:
raise ValueError('image_qc_classifier_threshold must be between 0 and 1')
policy['image_qc_classifier_threshold'] = threshold
return policy
[docs]
def assess_image(image, settings, channel_ids=None):
"""Measure focus, saturation and nonfinite pixels on unnormalized channels.
:param image: YX, YXC or leading-dimensions plus YXC array.
:param settings: image_qc_* policy, validated before use.
:param channel_ids: optional acquisition-channel labels for stored C planes.
:returns: one metric/reason record per selected channel. Focus is the best
plane's Laplacian variance in raw intensity units squared, avoiding
rejection solely because a z stack includes out-of-focus planes.
Saturation uses an explicit acquisition level or integer dtype ceiling;
it never uses the brightest observed pixel. Object counts are not read.
:raises ValueError: channels, shape or saturation calibration are unavailable.
"""
import numpy as np
from scipy.ndimage import laplace
policy = quality_policy(settings)
image = np.asarray(image)
if image.ndim == 2:
image = image[..., None]
if image.ndim < 3 or not image.size:
raise ValueError('Image quality requires a nonempty YX or (..., Y, X, C) image')
channel_ids = list(range(image.shape[-1])) if channel_ids is None else list(channel_ids)
if len(channel_ids) != image.shape[-1]:
raise ValueError('Image quality channel mapping must match the raw channel planes')
requested = policy['image_qc_channels'] or list(dict.fromkeys(channel_ids))
if not set(requested) <= set(channel_ids):
raise ValueError('Selected image-quality channels are absent from this field')
for key in ('image_qc_min_focus', 'image_qc_max_saturation', 'image_qc_saturation_level'):
if not set(map(int, policy[key])) <= set(requested):
raise ValueError(f'{key} contains a channel that is not being screened')
records = []
for channel in requested:
key = str(channel)
plane = image[..., channel_ids.index(channel)]
finite = np.isfinite(plane)
invalid = float(1 - finite.mean())
numeric = np.where(finite, plane, 0).astype(np.float64)
planes = numeric.reshape((-1, *numeric.shape[-2:]))
focus = max(float(np.var(laplace(member))) for member in planes)
ceiling = policy['image_qc_saturation_level'].get(key)
if ceiling is None and np.issubdtype(image.dtype, np.integer):
ceiling = float(np.iinfo(image.dtype).max)
fraction = float(np.mean(finite & (plane >= ceiling))) if ceiling is not None else None
if key in policy['image_qc_max_saturation'] and ceiling is None:
raise ValueError(f'Channel {channel}: floating images need image_qc_saturation_level')
reasons = []
if key in policy['image_qc_min_focus'] and focus < policy['image_qc_min_focus'][key]:
reasons.append('focus_below_threshold')
if key in policy['image_qc_max_saturation'] and fraction > policy['image_qc_max_saturation'][key]:
reasons.append('saturation_above_threshold')
if invalid > policy['image_qc_max_nonfinite']:
reasons.append('nonfinite_pixels')
records.append(dict(channel=channel, focus_variance=focus,
saturation_level=ceiling, saturation_fraction=fraction,
nonfinite_fraction=invalid, reasons=reasons))
return records
def _atomic_text(path, text):
"""Publish one complete UTF-8 report without exposing partial writes."""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
descriptor, temporary = tempfile.mkstemp(dir=path.parent, prefix='.image_quality_')
try:
with os.fdopen(descriptor, 'w', encoding='utf-8') as output:
output.write(text)
output.flush()
os.fsync(output.fileno())
os.replace(temporary, path)
finally:
if os.path.exists(temporary):
os.unlink(temporary)
def _write_gallery(destination, fields, paths, channel_ids):
"""Write a read-only review gallery with at most 64 field thumbnails."""
import base64
from html import escape
import numpy as np
from PIL import Image
by_name = {Path(path).name: Path(path) for path in paths}
cards = []
ordered = sorted(fields, key=lambda field: (not bool(field['reasons']), field['field']))
for field in ordered[:64]:
path = by_name[field['field']]
image = np.load(path, mmap_mode='r', allow_pickle=False)
metric = next((record for record in field['channels'] if record['reasons']), field['channels'][0])
channel = metric['channel']
if image.ndim > 2:
position = list(channel_ids).index(channel) if channel_ids is not None else channel
image = image[..., position]
if image.ndim > 2:
image = np.max(image.reshape((-1, *image.shape[-2:])), axis=0)
step = max(1, int(math.ceil(max(image.shape) / 512)))
image = np.asarray(image[::step, ::step], dtype=np.float32)
values = image[np.isfinite(image)]
low, high = np.percentile(values, (1, 99)) if values.size else (0., 1.)
scaled = np.nan_to_num((image - low) / max(float(high - low), 1e-12), nan=0., posinf=1., neginf=0.)
thumbnail = Image.fromarray((scaled.clip(0, 1) * 255).astype(np.uint8))
thumbnail.thumbnail((256, 256))
buffer = io.BytesIO()
thumbnail.save(buffer, format='PNG')
encoded = base64.b64encode(buffer.getvalue()).decode('ascii')
details = '; '.join(field['reasons']) or 'No configured criterion failed.'
cards.append('<article><img src="data:image/png;base64,' + encoded + '"><div><b>' +
escape(field['field']) + '</b> — ' + escape(field['status']) + '. ' +
escape(details) + f' Preview channel {channel}. Focus variance {metric["focus_variance"]:.5g}; ' +
'saturated fraction ' + escape(str(metric['saturation_fraction'])) + '.</div></article>')
html = ('<!doctype html><html><meta charset="utf-8"><title>Image Quality</title><style>'
'body{background:#14171b;color:#edf0f4;font:16px system-ui;margin:24px}'
'main{display:grid;grid-template-columns:repeat(auto-fit,minmax(320px,1fr));gap:12px}'
'article{border:1px solid #43505e;border-radius:12px;padding:12px;background:#20252bcc}'
'img{max-width:100%;display:block;margin:0 auto 8px;border-radius:8px}'
'header{margin-bottom:18px;line-height:1.5}</style><header><b>Image Quality.</b> '
'Flagged fields appear first. Up to 64 previews are shown; the CSV contains every field and channel. '
'Preview contrast is stretched for viewing only; metrics use raw intensities. '
'Volumes use a maximum projection for this preview. Excluded fields are not zero-object detections.'
'</header><main>' + ''.join(cards) + '</main></html>')
_atomic_text(destination.with_suffix('.html'), html)
[docs]
def screen_fields(root, settings, paths=None, channel_ids=None):
"""Save a field/channel report and return fields explicitly excluded by policy.
:param root: project folder; reports go to qc/image_quality.json and .csv.
:param settings: Mask settings; mode off clears a previously active policy.
:param paths: optional iterable of raw NPY paths; defaults to root/stack.
:param channel_ids: optional acquisition-channel labels for supplied arrays.
:returns: excluded basenames, empty in report-only or off mode. Inputs are
never deleted or rewritten. Report and exact policy precede exclusion.
:raises ValueError: enabled screening has no raw fields or invalid calibration.
"""
import numpy as np
from .cancellation import checkpoint
root = Path(root)
policy = quality_policy(settings)
destination = root / REPORT
if policy['image_qc_mode'] == 'off' and not destination.exists():
return []
fields, excluded, rows = [], [], []
paths = [] if policy['image_qc_mode'] == 'off' else paths
classifier = None
if policy['image_qc_mode'] != 'off':
paths = sorted(root.joinpath('stack').glob('*.npy')) if paths is None else list(paths)
if not paths:
raise ValueError('Image quality needs the raw stack fields; restore stack/ or rerun preprocessing')
if policy['image_qc_classifier']:
classifier = _prepare_qc_classifier(root, policy, paths, channel_ids)
for path in paths:
checkpoint()
path = Path(path)
image = np.load(path, mmap_mode='r', allow_pickle=False)
metrics = assess_image(image, policy, channel_ids)
if classifier is not None:
_classify_records(classifier, image, metrics, policy, channel_ids)
reasons = [f"channel {record['channel']}: {reason}"
for record in metrics for reason in record['reasons']]
status = ('excluded_image_quality' if policy['image_qc_mode'] == 'exclude' else 'flagged') if reasons else 'accepted'
fields.append(dict(field=path.name, status=status, reasons=reasons, channels=metrics))
if status == 'excluded_image_quality':
excluded.append(path.name)
rows.extend(dict(field=path.name, status=status, **dict(record, reasons='; '.join(record['reasons'])))
for record in metrics)
ensure_no_retained_measurements(root, excluded)
report = dict(version=1, created_at=datetime.now(timezone.utc).isoformat(),
policy=policy, excluded_fields=excluded, fields=fields)
stream = io.StringIO()
columns = ['field', 'status', 'channel', 'focus_variance', 'saturation_level',
'saturation_fraction', 'nonfinite_fraction', 'reasons']
if classifier is not None:
columns[-1:-1] = [f'p_{name}' for name in _QC_CLASSES]
writer = csv.DictWriter(stream, fieldnames=columns)
writer.writeheader()
writer.writerows(rows)
_write_gallery(destination, fields, paths, channel_ids)
_atomic_text(destination.with_suffix('.csv'), stream.getvalue())
_atomic_text(destination, json.dumps(report, indent=2, allow_nan=False) + '\n')
print(f'Image quality: {len(fields)} fields screened; {len(excluded)} excluded. Report: {destination}')
return excluded
[docs]
def ensure_no_retained_measurements(root, rejected):
"""Refuse exclusions that would leave previously measured fields in reports.
Existing results are never deleted. Re-screening an analyzed project needs
a fresh project if excluded fields already have measurement rows.
:param root: project folder containing measurements/measurements.db.
:param rejected: rejected field filenames, also matched without extensions.
:returns: None if no existing measurement row matches a rejected field.
:raises ValueError: a table's file_name column contains a rejected identity.
"""
from contextlib import closing
from .database_concurrency import connect
database = Path(root) / 'measurements' / 'measurements.db'
if not rejected or not database.is_file():
return
identities = sorted({value for name in rejected for value in (name, Path(name).stem)})
with closing(connect(database, readonly=True)) as connection:
tables = [row[0] for row in connection.execute("SELECT name FROM sqlite_master WHERE type='table'")]
for table in tables:
quoted = '"' + table.replace('"', '""') + '"'
columns = {row[1] for row in connection.execute(f'PRAGMA table_info({quoted})')}
if 'file_name' not in columns:
continue
for offset in range(0, len(identities), 400):
selected = identities[offset:offset + 400]
placeholders = ','.join('?' for _ in selected)
if connection.execute(f'SELECT 1 FROM {quoted} WHERE file_name IN ({placeholders}) LIMIT 1', selected).fetchone():
raise ValueError(
'Image-quality exclusions overlap existing measurements. Use a fresh project '
'for this exclusion policy, or select report mode to review without exclusion. '
'Existing images, measurements and the previous quality policy were retained.')
[docs]
def excluded_fields(root):
"""Read the active saved exclusion policy for downstream field consumers.
:param root: project folder whose qc/image_quality.json is authoritative.
:returns: excluded NPY basenames, or an empty set when no policy is active.
:raises ValueError: a saved report has an unsupported or invalid schema.
"""
path = Path(root) / REPORT
if not path.exists():
return set()
report = json.loads(path.read_text(encoding='utf-8'))
if report.get('version') != 1:
raise ValueError(f'Unsupported image-quality report: {path}')
if report.get('policy', {}).get('image_qc_mode') != 'exclude':
return set()
names = report.get('excluded_fields', [])
if not isinstance(names, list) or any(not isinstance(name, str) or Path(name).name != name for name in names):
raise ValueError(f'Invalid image-quality field identities: {path}')
return set(names)
[docs]
def filter_batch(batch, filenames, settings):
"""Remove policy-excluded fields without substituting empty label images.
:param batch: normalized image batch with its original field axis.
:param filenames: matching basenames in batch order.
:param settings: current run settings carrying image_qc_excluded_fields.
:returns: filtered batch and matching filenames in their original order.
"""
excluded = set(settings.get('image_qc_excluded_fields', ()))
if not excluded:
return batch, filenames
keep = [index for index, name in enumerate(filenames) if str(name) not in excluded]
return batch[keep], [filenames[index] for index in keep]
_QC_CLASSES = ('out_of_focus', 'saturated', 'debris', 'bubble', 'empty')
_QC_TILE = 64
_QC_TILES = 16
_QC_BUILTIN_FIELDS = 2400
_QC_BUILTIN_EPOCHS = 12
_QC_FINE_TUNE_EPOCHS = 15
_QC_MODEL_VERSION = 1
_QC_BUILTIN_CUTOFFS = {'out_of_focus': 0.9, 'saturated': 0.5, 'debris': 0.95,
'bubble': 0.75, 'empty': 0.5}
_QC_LABEL_ALIASES = {'good': (), 'ok': (), 'pass': (), 'blur': ('out_of_focus',),
'blurry': ('out_of_focus',), 'defocus': ('out_of_focus',),
'saturation': ('saturated',), 'bubbles': ('bubble',),
'blank': ('empty',)}
def _builtin_qc_model_path():
"""Where the built-in classifier is cached after it is first trained."""
return _spacr_home() / 'models' / f'image_qc_classifier_v{_QC_MODEL_VERSION}.pt'
def _best_focus_plane(plane):
"""Return the 2-D plane of a volume with the largest Laplacian variance."""
import numpy as np
from scipy.ndimage import laplace
plane = np.nan_to_num(np.asarray(plane, dtype=np.float32))
if plane.ndim == 2:
return plane
planes = plane.reshape((-1, *plane.shape[-2:]))
return planes[int(np.argmax([np.var(laplace(member)) for member in planes]))]
def _qc_inputs(plane, ceiling=None, tiles=_QC_TILES):
"""Turn one raw channel into the classifier's scale-free inputs.
Three maps are derived from the raw plane: intensity above the median
background divided by the brighter of the signal range and 20 noise
standard deviations, so a field of noise alone stays near zero; the
Laplacian magnitude relative to pixel noise, so blur shows as absent
structure rather than a dim image; and the pixels at the saturation
ceiling. ``tiles`` native-resolution 64-pixel tiles spread over the
field show cell-scale detail and a 64 by 64 average of the whole field
shows field-scale structure such as bubbles.
:returns: ``(tiles, 3, 64, 64)`` and ``(3, 64, 64)`` float32 arrays.
"""
import numpy as np
import torch
from scipy.ndimage import laplace
plane = _best_focus_plane(plane)
background = float(np.median(plane))
step = np.diff(plane[:, ::2], axis=1).ravel()
noise = 1.4826 * float(np.median(np.abs(step - np.median(step)))) / math.sqrt(2)
scale = max(float(np.percentile(plane, 99.9)) - background, 20 * noise, 1e-6)
maps = np.stack([np.clip((plane - background) / scale, -1, 3),
np.log1p(np.abs(laplace(plane)) / max(noise, 1e-3 * scale)),
(plane >= ceiling) if ceiling is not None else np.zeros_like(plane)]
).astype(np.float32)
size = _QC_TILE
pad = [(0, 0)] + [(0, max(0, size - extent)) for extent in maps.shape[1:]]
padded = np.pad(maps, pad, mode='reflect' if min(maps.shape[1:]) > 1 else 'edge')
side = int(math.ceil(math.sqrt(tiles)))
rows = np.linspace(0, padded.shape[1] - size, side).astype(int)
columns = np.linspace(0, padded.shape[2] - size, side).astype(int)
stack = np.stack([padded[:, y:y + size, x:x + size] for y in rows for x in columns][:tiles])
tensor = torch.from_numpy(np.ascontiguousarray(padded))[None]
thumbnail = torch.nn.functional.adaptive_avg_pool2d(tensor, (size, size))[0].numpy()
return stack, thumbnail
def _qc_network():
"""Build the two-branch convolutional classifier with random weights.
One branch reads each native-resolution tile and keeps the strongest
response per tile, the other reads the whole-field average; tile
responses are summarised by maximum and mean so a single bad region
still counts. One logit per class in ``_QC_CLASSES``.
"""
import torch
from torch import nn
def trunk():
"""Four 3x3 convolution stages shared in shape by both branches."""
layers, width = [], 3
for index, out in enumerate((16, 32, 48, 48)):
layers += [nn.Conv2d(width, out, 3, padding=1), nn.ReLU()]
if index < 3:
layers.append(nn.MaxPool2d(2))
width = out
return nn.Sequential(*layers)
class QCNet(nn.Module):
"""Tile and whole-field branches joined by one linear layer."""
def __init__(self):
"""Build the two convolutional trunks and the shared head."""
super().__init__()
self.tile = trunk()
self.field = trunk()
self.head = nn.Linear(48 * 3, len(_QC_CLASSES))
def forward(self, tiles, thumbnail):
"""Return class logits for a batch of fields."""
batch, count = tiles.shape[:2]
local = self.tile(tiles.flatten(0, 1)).amax((2, 3)).view(batch, count, -1)
whole = self.field(thumbnail).amax((2, 3))
return self.head(torch.cat([local.amax(1), local.mean(1), whole], 1))
return QCNet()
def _train_qc_network(model, tiles, thumbnails, targets, epochs, seed=0, rate=2e-3, balance=10.0):
"""Fit ``model`` in place on the CPU with flips, rotations and tile dropout.
At most eight CPU threads are used while fitting.
:param tiles: ``(fields, tiles, 3, 64, 64)`` array.
:param thumbnails: ``(fields, 3, 64, 64)`` array.
:param targets: ``(fields, classes)`` array of 0/1 labels; a field with
no defect has an all-zero row.
:param balance: the most a rare class is weighted up relative to its
absent cases; high when training from scratch, low when fine-tuning
so a few labels do not inflate false flags.
:returns: the model, left in evaluation mode.
"""
import numpy as np
import torch
generator = torch.Generator().manual_seed(seed)
tiles = torch.as_tensor(np.asarray(tiles, np.float32))
thumbnails = torch.as_tensor(np.asarray(thumbnails, np.float32))
targets = torch.as_tensor(np.asarray(targets, np.float32))
positive = targets.mean(0).clamp(0.02, 0.98)
loss = torch.nn.BCEWithLogitsLoss(
pos_weight=((1 - positive) / positive).clamp(max=float(balance)))
optimiser = torch.optim.Adam(model.parameters(), lr=rate)
model.train()
keep = max(1, tiles.shape[1] // 2)
threads = torch.get_num_threads()
torch.set_num_threads(min(threads, 8))
try:
for _ in range(int(epochs)):
order = torch.randperm(len(targets), generator=generator)
for start in range(0, len(order), 32):
index = order[start:start + 32]
chosen = torch.randperm(tiles.shape[1], generator=generator)[:keep]
local, whole = tiles[index][:, chosen], thumbnails[index]
turns = int(torch.randint(4, (1,), generator=generator))
local, whole = local.rot90(turns, (-2, -1)), whole.rot90(turns, (-2, -1))
if torch.rand(1, generator=generator) < .5:
local, whole = local.flip(-1), whole.flip(-1)
optimiser.zero_grad()
loss(model(local, whole), targets[index]).backward()
optimiser.step()
finally:
torch.set_num_threads(threads)
return model.eval()
def _predict_qc(model, tiles, thumbnails):
"""Return class probabilities, one row per field."""
import numpy as np
import torch
with torch.no_grad():
outputs = [torch.sigmoid(model(torch.as_tensor(np.asarray(tiles[start:start + 64], np.float32)),
torch.as_tensor(np.asarray(thumbnails[start:start + 64], np.float32))))
for start in range(0, len(tiles), 64)]
return torch.cat(outputs).numpy() if outputs else np.zeros((0, len(_QC_CLASSES)), np.float32)
def _synthetic_qc_field(rng, defects, size=256, ceiling=65535):
"""Render one synthetic fluorescence field carrying the named defects.
Cells are textured ellipses on an uneven background with shot and read
noise. ``out_of_focus`` blurs the optics, ``saturated`` raises the gain
until part of the field clips at ``ceiling``, ``debris`` adds bright
fibres or aggregates, ``bubble`` adds a dimmed disc with a refractive
rim and ``empty`` leaves background and noise only.
"""
import numpy as np
from scipy.ndimage import gaussian_filter
y, x = np.mgrid[:size, :size].astype(np.float32)
background = rng.uniform(80, 0.03 * ceiling)
signal = np.zeros((size, size), np.float32)
if 'empty' not in defects:
radius = rng.uniform(3, 16)
crowd = min(1200, int(rng.uniform(0, 1) * (size / radius) ** 2))
region = np.ones((size, size), bool)
if rng.uniform() < .4:
smooth = gaussian_filter(rng.normal(size=(size, size)), size / rng.uniform(4, 10))
region = smooth > np.percentile(smooth, rng.uniform(20, 70))
spots = np.argwhere(region)
for _ in range(int(rng.integers(2, 12)) + int(rng.integers(0, crowd + 1))):
cy, cx = spots[int(rng.integers(len(spots)))] + rng.uniform(-.5, .5, 2)
a, b = radius * rng.uniform(.6, 1.4, 2)
reach = int(2 * max(a, b)) + 2
top, left = max(0, int(cy) - reach), max(0, int(cx) - reach)
bottom, right = min(size, int(cy) + reach), min(size, int(cx) + reach)
if bottom <= top or right <= left:
continue
angle = rng.uniform(0, np.pi)
dy, dx = y[top:bottom, left:right] - cy, x[top:bottom, left:right] - cx
u = (dx * np.cos(angle) + dy * np.sin(angle)) / a
v = (-dx * np.sin(angle) + dy * np.cos(angle)) / b
signal[top:bottom, left:right] += rng.uniform(.3, 1) * np.clip(1.3 - u * u - v * v, 0, .3) / .3
texture = 1 + rng.uniform(.3, 1.5) * gaussian_filter(rng.normal(size=signal.shape), rng.uniform(.8, 2))
signal = gaussian_filter(signal * np.clip(texture, .3, None), rng.uniform(.5, 1.2))
else:
for _ in range(int(rng.integers(0, 3))):
cy, cx = rng.uniform(0, size, 2)
signal += .3 * np.exp(-((y - cy) ** 2 + (x - cx) ** 2) / 2)
amplitude = rng.uniform(15, 60) * math.sqrt(background) + rng.uniform(0, .5) * (.6 * ceiling - background)
field = signal * amplitude
haze = gaussian_filter(rng.normal(size=(size, size)), size / rng.uniform(2, 8))
haze = haze / max(float(np.abs(haze).max()), 1e-6)
field = field + haze * (rng.uniform(0, 8) * math.sqrt(background) if 'empty' in defects
else rng.uniform(0, .4) * amplitude)
if 'debris' in defects:
for _ in range(int(rng.integers(1, 4))):
debris = np.zeros_like(signal)
if rng.uniform() < .5:
cy, cx = rng.uniform(0, size, 2)
heading = rng.uniform(0, 2 * np.pi)
for _ in range(int(rng.integers(60, 200))):
heading += rng.normal(0, .15)
cy, cx = cy + np.sin(heading), cx + np.cos(heading)
if 0 <= cy < size and 0 <= cx < size:
debris[int(cy), int(cx)] = 1
debris = gaussian_filter(debris, rng.uniform(1, 2.5))
debris /= max(debris.max(), 1e-6)
elif rng.uniform() < .5:
cy, cx = rng.uniform(0, size, 2)
corners = rng.integers(3, 8)
angles = np.sort(rng.uniform(0, 2 * np.pi, corners))
reach = rng.uniform(6, 30) * rng.uniform(.5, 1, corners)
from matplotlib.path import Path as Outline
outline = Outline(np.c_[cx + reach * np.cos(angles), cy + reach * np.sin(angles)])
debris = outline.contains_points(np.c_[x.ravel(), y.ravel()]).reshape(x.shape).astype(np.float32)
else:
cy, cx = rng.uniform(0, size, 2)
for _ in range(int(rng.integers(3, 9))):
oy, ox = rng.normal(0, rng.uniform(4, 12), 2)
spread = rng.uniform(3, 9)
debris += np.exp(-((y - cy - oy) ** 2 + (x - cx - ox) ** 2) / (2 * spread ** 2))
debris = np.clip(debris, 0, 1)
field += debris * max(amplitude, 30 * math.sqrt(background)) * rng.uniform(1.5, 6)
optics = rng.uniform(2, 8) if 'out_of_focus' in defects else rng.uniform(0, 1)
if optics:
field = gaussian_filter(field, optics)
field = field + background * (1 + rng.uniform(-.3, .3) * (x / size) + rng.uniform(-.3, .3) * (y / size))
if 'bubble' in defects:
cy, cx = rng.uniform(-.1, 1.1, 2) * size
radius = rng.uniform(.18, .45) * size
distance = np.hypot(y - cy, x - cx)
inside = gaussian_filter((distance < radius).astype(np.float32), 2)
rim = np.exp(-((distance - radius) ** 2) / (2 * rng.uniform(1, 3) ** 2))
field = field * (1 - inside * rng.uniform(.4, .8)) + rim * rng.choice([-.5, 1.5]) * float(np.median(field))
if 'saturated' in defects:
fraction = rng.uniform(.01, .12)
level = np.percentile(field, 100 * (1 - fraction))
field = background + (field - background) * (ceiling * 1.05 - background) / max(level - background, 1e-6)
else:
peak = float(field.max())
if peak > .85 * ceiling:
field = background + (field - background) * (.85 * ceiling - background) / (peak - background)
field = np.clip(field, 0, None)
gain = rng.uniform(1, 4)
field = rng.poisson(field / gain) * gain + rng.normal(0, rng.uniform(2, 10), field.shape)
return np.clip(field, 0, ceiling).astype(np.uint16 if ceiling <= 65535 else np.float32)
def _synthetic_qc_labels(rng):
"""Draw the defects for one training field: most single, some paired."""
draw = rng.uniform()
if draw < .35:
return ()
first = _QC_CLASSES[int(rng.integers(len(_QC_CLASSES)))]
if first != 'empty' and rng.uniform() < .2:
second = _QC_CLASSES[int(rng.integers(len(_QC_CLASSES) - 1))]
return tuple(sorted({first, second}))
return (first,)
def _train_builtin_qc_model(fields=None, epochs=None, seed=0):
"""Train the built-in classifier on synthetic fields with planted defects.
:returns: a model in evaluation mode.
"""
import numpy as np
import torch
rng = np.random.default_rng(seed)
torch.manual_seed(seed)
tiles, thumbnails, targets = [], [], []
for _ in range(int(fields or _QC_BUILTIN_FIELDS)):
labels = _synthetic_qc_labels(rng)
ceiling = int(rng.choice([4095, 65535]))
image = _synthetic_qc_field(rng, labels, size=int(rng.choice([192, 256, 320])), ceiling=ceiling)
local, whole = _qc_inputs(image, ceiling)
tiles.append(local)
thumbnails.append(whole)
targets.append([name in labels for name in _QC_CLASSES])
return _train_qc_network(_qc_network(), np.stack(tiles), np.stack(thumbnails),
np.asarray(targets, np.float32), epochs or _QC_BUILTIN_EPOCHS, seed)
def _save_qc_model(model, path, source):
"""Save weights with their class order so a later run reads tensors only."""
import torch
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_name('.' + path.name + '.part')
torch.save(dict(version=_QC_MODEL_VERSION, classes=list(_QC_CLASSES), source=str(source),
state=model.state_dict()), temporary)
os.replace(temporary, path)
return path
def _load_qc_model(path):
"""Load a saved classifier as tensors only, refusing another class order."""
import torch
saved = torch.load(Path(path), map_location='cpu', weights_only=True)
if not isinstance(saved, dict) or saved.get('version') != _QC_MODEL_VERSION or \
tuple(saved.get('classes', ())) != _QC_CLASSES:
raise ValueError(f'{path} is not an image-quality classifier saved by this version of spaCR')
model = _qc_network()
model.load_state_dict(saved['state'])
return model.eval()
def _base_qc_model(policy):
"""The saved model named by the policy, or the cached built-in one."""
if policy['image_qc_classifier_model']:
return _load_qc_model(policy['image_qc_classifier_model'])
path = _builtin_qc_model_path()
if path.is_file():
model = _load_qc_model(path)
else:
print('Image quality: training the built-in classifier on synthetic fields '
'(once, on the CPU; about a few minutes)...')
model = _train_builtin_qc_model()
_save_qc_model(model, path, 'built-in synthetic planted defects')
model.qc_cutoffs = dict(_QC_BUILTIN_CUTOFFS)
return model
def _calibrated_qc(model, probabilities):
"""Move each class's operating point of a calibrated model to 0.5.
The built-in classifier, trained on synthetic fields only, scores real
in-focus fields high for debris, bubble and blur. Its ``qc_cutoffs``
were measured on real labelled fields; a probability equal to a class's
cut-off is mapped to 0.5 by a shift in log-odds, so
``image_qc_classifier_threshold`` keeps its meaning and its default
lands on the measured operating point. A model without cut-offs, such
as one fine-tuned on the user's labels, is returned unchanged.
:param model: the classifier, optionally carrying ``qc_cutoffs``.
:param probabilities: ``(fields, classes)`` raw probabilities.
:returns: the calibrated probabilities.
"""
import numpy as np
cutoffs = getattr(model, 'qc_cutoffs', None)
if not cutoffs:
return probabilities
p = np.clip(np.asarray(probabilities, np.float64), 1e-7, 1 - 1e-7)
cut = np.array([cutoffs.get(name, 0.5) for name in _QC_CLASSES], np.float64)
shift = np.log(cut / (1 - cut))
return (1 / (1 + np.exp(-(np.log(p / (1 - p)) - shift)))).astype(np.float32)
def _field_ceiling(image, policy, channel):
"""The saturation level assess_image would use for ``channel``."""
import numpy as np
level = policy['image_qc_saturation_level'].get(str(channel))
if level is None and np.issubdtype(np.asarray(image).dtype, np.integer):
level = float(np.iinfo(np.asarray(image).dtype).max)
return level
def _classify_records(model, image, records, policy, channel_ids=None):
"""Add class probabilities and classifier reasons to assess_image records.
Each screened channel is classified on its own; a class at or above
``image_qc_classifier_threshold`` adds ``classifier_<class>`` to that
channel's reasons, which exclude mode then treats like any other flag.
"""
import numpy as np
image = np.asarray(image)
if image.ndim == 2:
image = image[..., None]
channel_ids = list(range(image.shape[-1])) if channel_ids is None else list(channel_ids)
for record in records:
channel = record['channel']
local, whole = _qc_inputs(image[..., channel_ids.index(channel)],
_field_ceiling(image, policy, channel))
probabilities = _calibrated_qc(model, _predict_qc(model, local[None], whole[None]))[0]
for name, value in zip(_QC_CLASSES, probabilities):
record[f'p_{name}'] = round(float(value), 4)
if value >= policy['image_qc_classifier_threshold']:
record['reasons'].append(f'classifier_{name}')
return records
def _parse_qc_labels(value):
"""Split one annotation cell into classifier classes; good means none."""
names = set()
for token in str(value).replace(',', ';').split(';'):
token = token.strip().lower().replace(' ', '_').replace('-', '_')
if not token:
continue
if token in _QC_CLASSES:
names.add(token)
elif token in _QC_LABEL_ALIASES:
names.update(_QC_LABEL_ALIASES[token])
else:
raise ValueError(f'Unknown image-quality label {value!r}; use good or '
+ ', '.join(_QC_CLASSES))
return names
def _labelled_qc_fields(policy, paths, channel_ids=None):
"""Read the annotation table and pair each labelled field with its inputs.
The table needs a field column (the raw file name, with or without .npy)
and a label column: good, or one or more of the classes separated by
semicolons. An optional channel column picks the channel; otherwise the
first screened channel is used.
:returns: list of dicts with field, channel, labels, tiles, thumbnail and
the rule-based focus and saturation metrics of that channel.
"""
import numpy as np
from .tabular import read_table
table = read_table(policy['image_qc_classifier_labels'], report=None)
table.columns = [str(column).strip().lower() for column in table.columns]
table = table.rename(columns={'fieldid': 'field', 'chanid': 'channel'})
if not {'field', 'label'} <= set(table.columns):
raise ValueError('image_qc_classifier_labels needs field and label columns')
by_name = {}
for path in paths:
by_name[Path(path).name] = Path(path)
by_name[Path(path).stem] = Path(path)
samples = []
for row in table.to_dict('records'):
name = str(row['field']).strip()
path = by_name.get(name) or by_name.get(Path(name).name)
if path is None:
continue
image = np.load(path, mmap_mode='r', allow_pickle=False)
records = assess_image(image, policy, channel_ids)
channel = row.get('channel')
if channel is None or (isinstance(channel, float) and math.isnan(channel)):
record = records[0]
else:
record = next((item for item in records if str(item['channel']) == str(int(channel))), None)
if record is None:
raise ValueError(f'Labelled channel {channel} of {name} is not being screened')
stored = image[..., None] if image.ndim == 2 else image
ids = list(range(stored.shape[-1])) if channel_ids is None else list(channel_ids)
local, whole = _qc_inputs(stored[..., ids.index(record['channel'])],
_field_ceiling(stored, policy, record['channel']))
samples.append(dict(field=path.name, channel=record['channel'], labels=_parse_qc_labels(row['label']),
tiles=local, thumbnail=whole, focus=record['focus_variance'],
saturation=record['saturation_fraction'] or 0.0,
rule_flag=bool(record['reasons'])))
return samples
def _precision_recall(truth, flagged):
"""Precision and recall of boolean calls; undefined values are NaN."""
import numpy as np
truth, flagged = np.asarray(truth, bool), np.asarray(flagged, bool)
hits = float(np.sum(truth & flagged))
precision = hits / flagged.sum() if flagged.sum() else float('nan')
recall = hits / truth.sum() if truth.sum() else float('nan')
return precision, recall
def _tuned_rule(focus, saturation, truth):
"""Best focus and saturation cut-offs for any-defect F1 on training fields.
This is the most the current rule-based metrics can do when their
thresholds are chosen from labelled fields: a field is flagged when its
focus variance is below the focus cut-off or its saturated fraction is
above the saturation cut-off.
"""
import numpy as np
focus, saturation, truth = map(np.asarray, (focus, saturation, truth))
focus_cuts = np.concatenate([[-np.inf], np.unique(focus)])
saturation_cuts = np.concatenate([[np.inf], np.unique(saturation)])
best, choice = -1.0, (-np.inf, np.inf)
for low in focus_cuts:
for high in saturation_cuts:
flagged = (focus < low) | (saturation > high)
hits = np.sum(flagged & truth)
score = 2 * hits / max(flagged.sum() + truth.sum(), 1)
if score > best:
best, choice = score, (low, high)
return choice
def _benchmark_qc(samples, base, policy, folds=5, epochs=None, seed=0):
"""Cross-validated precision and recall: classifier against rule metrics.
Each fold fine-tunes a copy of ``base`` on the other folds and scores the
held-out fields; the rule baseline is scored twice, with the thresholds
of the saved policy and with cut-offs tuned on the same training folds.
:returns: rows with method, defect, precision, recall, positives, fields.
"""
import copy
import numpy as np
truth = np.array([[name in sample['labels'] for name in _QC_CLASSES] for sample in samples])
folds = max(2, min(int(folds), len(samples)))
order = np.random.default_rng(seed).permutation(len(samples))
probabilities = np.zeros(truth.shape, np.float32)
tuned = np.zeros((len(samples), 2), bool)
focus = np.array([sample['focus'] for sample in samples])
saturation = np.array([sample['saturation'] for sample in samples])
tiles = np.stack([sample['tiles'] for sample in samples])
thumbnails = np.stack([sample['thumbnail'] for sample in samples])
for fold in range(folds):
test = order[fold::folds]
train = np.setdiff1d(order, test)
model = _train_qc_network(copy.deepcopy(base), tiles[train], thumbnails[train],
truth[train], epochs or _QC_FINE_TUNE_EPOCHS, seed + fold,
rate=5e-4, balance=3.0)
probabilities[test] = _predict_qc(model, tiles[test], thumbnails[test])
low, high = _tuned_rule(focus[train], saturation[train], truth[train].any(1))
tuned[test] = np.c_[focus[test] < low, saturation[test] > high]
called = probabilities >= policy['image_qc_classifier_threshold']
configured = np.array([sample['rule_flag'] for sample in samples])
rule = np.zeros(truth.shape, bool)
rule[:, _QC_CLASSES.index('out_of_focus')] = tuned[:, 0]
rule[:, _QC_CLASSES.index('saturated')] = tuned[:, 1]
rows = []
for method, calls in (('classifier', called), ('rule_saved_policy', None), ('rule_tuned', rule)):
for index, name in enumerate(('any_defect',) + _QC_CLASSES):
actual = truth.any(1) if index == 0 else truth[:, index - 1]
if calls is None:
flagged = configured if index == 0 else np.zeros(len(samples), bool)
else:
flagged = calls.any(1) if index == 0 else calls[:, index - 1]
precision, recall = _precision_recall(actual, flagged)
rows.append(dict(method=method, defect=name, precision=precision, recall=recall,
positives=int(actual.sum()), fields=len(samples)))
return rows
def _prepare_qc_classifier(root, policy, paths, channel_ids=None):
"""Load or fine-tune the classifier this screening run will apply.
With ``image_qc_classifier_labels`` the labelled fields are first used to
score the classifier against the rule metrics by cross-validation
(qc/image_qc_benchmark.csv), then to fine-tune on all of them; the
result is saved as qc/image_qc_model.pt and applied to every field.
"""
import numpy as np
import pandas as pd
from .tabular import write_table
model = _base_qc_model(policy)
if not policy['image_qc_classifier_labels']:
return model
samples = _labelled_qc_fields(policy, paths, channel_ids)
if len(samples) < 4:
raise ValueError('image_qc_classifier_labels matched fewer than 4 raw fields; '
'label at least 4 fields by their file names')
qc = Path(root) / 'qc'
if len(samples) >= 10:
rows = _benchmark_qc(samples, model, policy)
write_table(pd.DataFrame(rows), qc / 'image_qc_benchmark.csv')
for row in rows:
if row['defect'] == 'any_defect':
print(f"Image quality {row['method']}: precision {row['precision']:.3f}, "
f"recall {row['recall']:.3f} over {row['fields']} labelled fields")
else:
print('Image quality: fewer than 10 labelled fields, so no held-out benchmark was run.')
truth = np.array([[name in sample['labels'] for name in _QC_CLASSES] for sample in samples], np.float32)
model = _train_qc_network(model, np.stack([sample['tiles'] for sample in samples]),
np.stack([sample['thumbnail'] for sample in samples]), truth, _QC_FINE_TUNE_EPOCHS,
rate=5e-4, balance=3.0)
model.qc_cutoffs = None
_save_qc_model(model, qc / 'image_qc_model.pt', policy['image_qc_classifier_labels'])
return model