diff --git a/CHANGELOG.MD b/CHANGELOG.MD index 1d27047..3796460 100644 --- a/CHANGELOG.MD +++ b/CHANGELOG.MD @@ -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 diff --git a/anomaly_match/data_io/checkpoint_io.py b/anomaly_match/data_io/checkpoint_io.py index 191b63c..a1cef57 100644 --- a/anomaly_match/data_io/checkpoint_io.py +++ b/anomaly_match/data_io/checkpoint_io.py @@ -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. @@ -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 diff --git a/anomaly_match/utils/get_default_cfg.py b/anomaly_match/utils/get_default_cfg.py index 969efbc..51dbb3c 100644 --- a/anomaly_match/utils/get_default_cfg.py +++ b/anomaly_match/utils/get_default_cfg.py @@ -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 diff --git a/environment.yml b/environment.yml index f2d07e7..618bf9a 100644 --- a/environment.yml +++ b/environment.yml @@ -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 diff --git a/environment_CI.yml b/environment_CI.yml index 6cc0bb8..dd29357 100644 --- a/environment_CI.yml +++ b/environment_CI.yml @@ -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 diff --git a/pyproject.toml b/pyproject.toml index 11c83a1..435057a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/tests/e2e/test_normalisation_consistency.py b/tests/e2e/test_normalisation_consistency.py index 13e5bec..88a3ccb 100644 --- a/tests/e2e/test_normalisation_consistency.py +++ b/tests/e2e/test_normalisation_consistency.py @@ -81,14 +81,18 @@ 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"] @@ -96,11 +100,11 @@ def _run_cutana_normalised(csv_path, fitsbolt_cfg, n_output_channels=3): 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): @@ -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): @@ -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 @@ -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)): diff --git a/tests/unit/test_checkpoint_io.py b/tests/unit/test_checkpoint_io.py index c8848e9..d983fb3 100644 --- a/tests/unit/test_checkpoint_io.py +++ b/tests/unit/test_checkpoint_io.py @@ -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" diff --git a/tests/unit/test_file_io.py b/tests/unit/test_file_io.py index c46a01a..ade4c6f 100644 --- a/tests/unit/test_file_io.py +++ b/tests/unit/test_file_io.py @@ -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) @@ -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 @@ -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 @@ -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"