Skip to content
Closed
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
10 changes: 10 additions & 0 deletions CHANGELOG.MD
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,16 @@
[//]: # (this file, may be copied, modified, propagated, or distributed except according to)
[//]: # (the terms contained in the file 'LICENCE.txt'.)

## [Unreleased]

### Changed
- **Relaxed the fitsbolt pin to `fitsbolt>=0.3.1,<0.4` and require `cutana>=0.4.0`** in `pyproject.toml`, `environment.yml`, and `environment_CI.yml`. fitsbolt 0.3.1 fixes the `fits_extension=None` validation crash that motivated the `==0.2.0` pin, and cutana 0.4.0 requires it, so the old pin forced cutana back to 0.3.1
- Image resizing in fitsbolt>=0.3 uses OpenCV instead of scikit-image: `interpolation_order` now accepts 0–4 (0=nearest, 1=linear, 2=cubic, 3=lanczos4, 4=area) and downscaling always uses area interpolation. The default (1, linear) is unaffected

### Fixed
- **Checkpoints saved with fitsbolt<0.3 failed at prediction under fitsbolt 0.3** (`TypeError` for every normalisation method except ZSCALE) because their stored `fitsbolt_cfg` lacks the new `*_n_samples` keys and stores midtones parameters as scalars. `load_checkpoint` now upgrades such configs in place, preserving the exact pre-0.3 behaviour
- `test_cutana_vs_training_normalisation` matched cutana cutouts to training images by position, but cutana's streaming orchestrator yields batches in completion order; cutouts are now matched by source ID

## [v1.3.2] – 2026-07-08

### Fixed
Expand Down
28 changes: 28 additions & 0 deletions anomaly_match/data_io/checkpoint_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -232,6 +232,33 @@ def save_checkpoint(save_state: dict[str, Any], path: str | Path) -> Path:
return path


def _upgrade_legacy_fitsbolt_cfg(fb_data: dict) -> dict:
"""Bring a fitsbolt config saved by fitsbolt<0.3 up to the >=0.3 schema.

fitsbolt 0.3 reads ``normalisation.{minmax,percentile,asinh}_n_samples``
directly and expects the midtones parameters as per-channel lists. Older
checkpoints lack those keys / store scalars, which crashes prediction.

Args:
fb_data: Decoded fitsbolt config dictionary from checkpoint metadata.

Returns:
The same dictionary, upgraded in place.
"""
if "normalisation" not in fb_data:
# Partial configs (not produced by fitsbolt's create_config) have nothing to upgrade.
return fb_data
norm = fb_data["normalisation"]
for key in ("minmax_n_samples", "percentile_n_samples", "asinh_n_samples"):
# None means "use all pixels", i.e. the exact pre-0.3 behaviour.
norm.setdefault(key, None)
midtones = norm["midtones"]
for key in ("percentile", "desired_mean"):
if isinstance(midtones[key], (int, float)):
midtones[key] = [midtones[key]]
return fb_data


def load_checkpoint(path: str | Path, device: str = "cpu") -> dict[str, Any]:
"""Load a model checkpoint from a ``.safetensors`` file.

Expand Down Expand Up @@ -315,6 +342,7 @@ def load_checkpoint(path: str | Path, device: str = "cpu") -> dict[str, Any]:
# _dynamic=False prevents DotMap from auto-creating empty child maps
# on missing-key access, which would break fitsbolt's validate_config
# (e.g. channel_combination should stay absent, not become DotMap()).
fb_data = _upgrade_legacy_fitsbolt_cfg(fb_data)
checkpoint["fitsbolt_cfg"] = DotMap(fb_data, _dynamic=False)
else:
checkpoint["fitsbolt_cfg"] = None
Expand Down
6 changes: 3 additions & 3 deletions anomaly_match/utils/get_default_cfg.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,9 +65,9 @@ def get_default_cfg():
cfg.normalisation.channel_combination = None

# further interpolation and normalisation settings
cfg.normalisation.interpolation_order = (
1 # order of interpolation for resizing with skimage, 0-5
)
# interpolation for resizing (fitsbolt>=0.3 uses OpenCV):
# 0=nearest, 1=linear, 2=cubic, 3=lanczos4, 4=area (always used when downscaling)
cfg.normalisation.interpolation_order = 1
cfg.normalisation.normalisation_method = NormalisationMethod.CONVERSION_ONLY
# settings for normalisation:
cfg.normalisation.norm_maximum_value = None # None or float
Expand Down
4 changes: 2 additions & 2 deletions environment.yml
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,8 @@ dependencies:
- pip
- pip:
- albumentations
- cutana>=0.3.0
- fitsbolt==0.2.0
- cutana>=0.4.0
- fitsbolt>=0.3.1,<0.4
- opencv-python-headless
- safetensors
- timm
4 changes: 2 additions & 2 deletions environment_CI.yml
Original file line number Diff line number Diff line change
Expand Up @@ -37,5 +37,5 @@ dependencies:
- albumentations
- safetensors
- timm
- fitsbolt==0.2.0
- cutana>=0.3.0
- fitsbolt>=0.3.1,<0.4
- cutana>=0.4.0
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -32,9 +32,9 @@ classifiers = [
dependencies = [
"albumentations",
"astropy",
"cutana>=0.3.0",
"cutana>=0.4.0",
"dotmap",
"fitsbolt==0.2.0",
"fitsbolt>=0.3.1,<0.4",
"h5py",
"ipykernel",
"ipywidgets",
Expand Down
39 changes: 26 additions & 13 deletions tests/e2e/test_normalisation_consistency.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,26 +81,30 @@ def _make_cutana_config(csv_path, n_output_channels=3):


def _run_cutana_normalised(csv_path, fitsbolt_cfg, n_output_channels=3):
"""Run cutana with external_fitsbolt_cfg and return normalised cutout arrays."""
"""Run cutana with external_fitsbolt_cfg and return normalised cutouts keyed by source ID.

Cutana's streaming orchestrator returns batches in completion order, so cutouts are
matched to sources through the batch metadata rather than by position.
"""
config = _make_cutana_config(csv_path, n_output_channels)
config.external_fitsbolt_cfg = fitsbolt_cfg

orchestrator = cutana.StreamingOrchestrator(config)
orchestrator.init_streaming(batch_size=10, write_to_disk=False)

all_cutouts = []
cutouts_by_source = {}
for _ in range(orchestrator.get_batch_count()):
batch = orchestrator.next_batch()
cutouts = batch["cutouts"]
if isinstance(cutouts, list) and len(cutouts) == 0:
continue
if isinstance(cutouts, list):
cutouts = np.array(cutouts)
for i in range(cutouts.shape[0]):
all_cutouts.append(np.array(cutouts[i]))
for i, source in enumerate(batch["metadata"]):
cutouts_by_source[source["source_id"]] = np.array(cutouts[i])

orchestrator.cleanup()
return all_cutouts
return cutouts_by_source


def _extract_raw_cutouts(csv_path, output_dir):
Expand Down Expand Up @@ -132,7 +136,7 @@ def cutana_test_data(tmp_path_factory):

Returns:
tuple: (clean_fits_paths, rewritten_csv_path)
- clean_fits_paths: list of FITS file paths with raw cutout data in HDU[0]
- clean_fits_paths: dict of source ID -> FITS file path with raw cutout data in HDU[0]
- rewritten_csv_path: path to CSV with absolute FITS tile paths
"""
if not os.path.exists(_CSV_CATALOGUE) or not os.path.exists(_FITS_TILE):
Expand All @@ -153,13 +157,17 @@ def cutana_test_data(tmp_path_factory):
pytest.skip("Cutana did not produce any cutouts from the test tile")

# Save as clean FITS files with data in HDU[0] for the training path
clean_fits_paths = []
with open(rewritten_csv) as f:
source_ids = [row["SourceID"] for row in csv.DictReader(f)]
clean_fits_paths = {}
for i, raw_path in enumerate(raw_fits_paths):
# Cutana embeds the source ID in the cutout filename
(source_id,) = [sid for sid in source_ids if sid in os.path.basename(raw_path)]
with fits.open(raw_path) as hdul:
raw_data = hdul[1].data
clean_path = str(tmp_path / f"cutout_{i}.fits")
fits.PrimaryHDU(raw_data.astype(np.float32)).writeto(clean_path, overwrite=True)
clean_fits_paths.append(clean_path)
clean_fits_paths[source_id] = clean_path

return clean_fits_paths, rewritten_csv

Expand Down Expand Up @@ -210,18 +218,23 @@ def test_cutana_vs_training_normalisation(cutana_test_data):
cutana_normalised = _run_cutana_normalised(
rewritten_csv, cfg.fitsbolt_cfg, n_output_channels=3
)
if len(cutana_normalised) != len(clean_fits_paths):
if cutana_normalised.keys() != clean_fits_paths.keys():
failures.append(
f"{method.name}: cutana returned {len(cutana_normalised)} cutouts "
f"but {len(clean_fits_paths)} raw cutouts were extracted"
f"{method.name}: cutana returned sources {sorted(cutana_normalised)} "
f"but raw cutouts were extracted for {sorted(clean_fits_paths)}"
)
continue
source_ids = sorted(clean_fits_paths)

format_cfg = create_cutana_format_cfg(cfg)
prediction_images = [convert_cutana_cutout(c, format_cfg) for c in cutana_normalised]
prediction_images = [
convert_cutana_cutout(cutana_normalised[sid], format_cfg) for sid in source_ids
]

# --- Training path: load raw FITS via fitsbolt ---
training_pairs = load_and_process_wrapper(clean_fits_paths, cfg, show_progress=False)
training_pairs = load_and_process_wrapper(
[clean_fits_paths[sid] for sid in source_ids], cfg, show_progress=False
)

# --- Compare ---
for i, (pred_img, (_, train_img)) in enumerate(zip(prediction_images, training_pairs)):
Expand Down
25 changes: 25 additions & 0 deletions tests/unit/test_checkpoint_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,31 @@ def test_fitsbolt_cfg_roundtrip(self, tmp_path):
assert loaded_fb.output_dtype == np.uint8
assert np.array_equal(loaded_fb.channel_combination, fb_cfg.channel_combination)

def test_legacy_fitsbolt_cfg_upgraded_for_prediction(self, tmp_path):
"""Checkpoints saved with fitsbolt<0.3 must still preprocess images under fitsbolt>=0.3."""
from fitsbolt.cfg.create_config import create_config

from anomaly_match.data_io.load_images import process_single_wrapper

fb_dict = create_config(size=[32, 32]).toDict()
# Recreate the fitsbolt 0.2 schema: no subsampling keys, scalar midtones parameters.
for key in ("minmax_n_samples", "percentile_n_samples", "asinh_n_samples"):
del fb_dict["normalisation"][key]
fb_dict["normalisation"]["midtones"]["percentile"] = 99.8
fb_dict["normalisation"]["midtones"]["desired_mean"] = 0.2

path = save_checkpoint(_make_full_checkpoint(fitsbolt_cfg=DotMap(fb_dict)), tmp_path / "m")
loaded_fb = load_checkpoint(path)["fitsbolt_cfg"]
assert loaded_fb.normalisation.midtones.percentile == [99.8]
assert loaded_fb.normalisation.minmax_n_samples is None

image = np.random.default_rng(0).random((48, 48, 3)).astype(np.float32)
for method in NormalisationMethod:
cfg = DotMap({"fitsbolt_cfg": DotMap(loaded_fb.toDict(), _dynamic=False)})
cfg.fitsbolt_cfg.normalisation_method = method
out = process_single_wrapper(image, cfg)
assert out.shape == (32, 32, 3), method

def test_labeled_data_csv_roundtrip(self, tmp_path):
"""Verify labeled_data_csv string survives round-trip."""
csv = "filename,label\nimg1.jpg,anomaly\nimg2.jpg,normal\n"
Expand Down
17 changes: 9 additions & 8 deletions tests/unit/test_file_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -839,8 +839,8 @@ def test_image_interpolation_orders(self, test_config):
"""Test different interpolation orders when resizing images.

This test creates a 40x40 image and a 200x200 image, then resizes both to 100x100
using different interpolation orders (0-5), which correspond to different polynomial
interpolation methods in scikit-image.
using different interpolation orders (0-4), which map to OpenCV interpolation flags in
fitsbolt>=0.3: 0=nearest, 1=linear, 2=cubic, 3=lanczos4, 4=area.
"""
# Create a small (40x40) test image with a clear pattern
small_img = np.zeros((40, 40, 3), dtype=np.uint8)
Expand Down Expand Up @@ -882,8 +882,8 @@ def test_image_interpolation_orders(self, test_config):
# Store results from different interpolation orders to compare them
upscaled_results = []

# Check each interpolation order (0-5)
for order in range(6):
# Check each interpolation order (0-4)
for order in range(5):
_update_config(test_config, interpolation_order=order)

# Resize small image (40x40 → 100x100) - upsampling
Expand Down Expand Up @@ -977,7 +977,8 @@ def test_image_interpolation_orders(self, test_config):

# Higher order interpolation (order > 1) should lead to smoother transitions
# This is difficult to quantify precisely, but we can check for values between the extremes for upscaling
if order >= 3:
# Order 4 (INTER_AREA) degenerates to nearest-like output when upscaling.
if order in (2, 3):
# For boundary regions, check that there are intermediate values
# between the pure colors in neighboring quadrants
# Sample near the boundary but not exactly on it
Expand All @@ -993,9 +994,9 @@ def test_image_interpolation_orders(self, test_config):
)

# Compare results between different interpolation orders to verify they're not identical
# We'll compare order 0 (nearest neighbor) with orders 1, 3, and 5
# These should produce visibly different results
for i, upscaled_im in enumerate(upscaled_results):
# These should produce visibly different results. Order 4 (INTER_AREA) is excluded:
# when upscaling OpenCV makes it identical to nearest neighbour (order 0).
for i, upscaled_im in enumerate(upscaled_results[:4]):
if i != 0:
assert not np.array_equal(upscaled_results[0], upscaled_results[i]), (
"Order 0 and order {i} interpolation should produce different results"
Expand Down
Loading