From 3a137cbb15fe89d5f248078c58ea249ffe6355c6 Mon Sep 17 00:00:00 2001 From: Tom Hehir <148493038+tom-hehir@users.noreply.github.com> Date: Mon, 31 Aug 2026 16:38:33 +0100 Subject: [PATCH 1/2] fix(image): invert HSC zeropoint rescaling --- aion/codecs/preprocessing/image.py | 2 +- tests/codecs/test_image_codec.py | 32 ++++++++++++++++++++++++++++++ 2 files changed, 33 insertions(+), 1 deletion(-) diff --git a/aion/codecs/preprocessing/image.py b/aion/codecs/preprocessing/image.py index f8b1f17..0ae1d80 100644 --- a/aion/codecs/preprocessing/image.py +++ b/aion/codecs/preprocessing/image.py @@ -102,6 +102,6 @@ def forward(self, image, survey): return image def backward(self, image, survey): - zpscale = self.reverse_zeropoint(27.0) if survey == "HSC" else 1.0 + zpscale = self.convert_zeropoint(27.0) if survey == "HSC" else 1.0 image *= zpscale return image diff --git a/tests/codecs/test_image_codec.py b/tests/codecs/test_image_codec.py index 76fdb36..7234371 100644 --- a/tests/codecs/test_image_codec.py +++ b/tests/codecs/test_image_codec.py @@ -3,9 +3,41 @@ from aion.codecs import ImageCodec from aion.codecs.config import HF_REPO_ID +from aion.codecs.preprocessing.image import RescaleToLegacySurvey from aion.modalities import Image +def test_hsc_rescaling_uses_legacy_survey_zeropoint(): + rescaler = RescaleToLegacySurvey() + image = torch.tensor([[[[-2.5, 0.0], [1.0, 12.5]]]]) + scale = rescaler.convert_zeropoint(27.0) + + rescaled = rescaler.forward(image.clone(), "HSC") + + assert torch.allclose(rescaled, image / scale) + + +def test_hsc_rescaling_round_trip(): + rescaler = RescaleToLegacySurvey() + image = torch.tensor([[[[-2.5, 0.0], [1.0, 12.5]]]]) + + rescaled = rescaler.forward(image.clone(), "HSC") + restored = rescaler.backward(rescaled, "HSC") + + assert torch.allclose(restored, image) + + +def test_non_hsc_rescaling_is_unchanged(): + rescaler = RescaleToLegacySurvey() + image = torch.tensor([[[[-2.5, 0.0], [1.0, 12.5]]]]) + + rescaled = rescaler.forward(image.clone(), "DES") + restored = rescaler.backward(rescaled, "DES") + + assert torch.equal(rescaled, image) + assert torch.equal(restored, image) + + @pytest.mark.parametrize("embedding_dim", [5, 10]) @pytest.mark.parametrize("multisurvey_projection_dims", [12, 24]) @pytest.mark.parametrize("hidden_dims", [8, 16]) From 0904e070a89aa1ff398aa3dbd768530084f458b5 Mon Sep 17 00:00:00 2001 From: Tom Hehir <148493038+tom-hehir@users.noreply.github.com> Date: Mon, 31 Aug 2026 16:51:32 +0100 Subject: [PATCH 2/2] test(image): cover reversible HSC preprocessing --- tests/codecs/test_image_codec.py | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/tests/codecs/test_image_codec.py b/tests/codecs/test_image_codec.py index 7234371..dfb6e63 100644 --- a/tests/codecs/test_image_codec.py +++ b/tests/codecs/test_image_codec.py @@ -38,6 +38,33 @@ def test_non_hsc_rescaling_is_unchanged(): assert torch.equal(restored, image) +def test_hsc_reversible_preprocessing_round_trip(): + codec = ImageCodec( + quantizer_levels=[1] * 5, + hidden_dims=8, + multisurvey_projection_dims=12, + n_compressions=2, + num_consecutive=1, + embedding_dim=5, + ) + bands = ["HSC-G", "HSC-R", "HSC-I", "HSC-Z", "HSC-Y"] + image = torch.linspace(-10.0, 10.0, 5 * 96 * 96).reshape(1, 5, 96, 96) + + # Crop and clamp are lossless for this already-cropped, in-range input. + processed = codec.center_crop(image.clone()) + processed = codec.clamp(processed, bands) + processed = codec.rescaler.forward(processed, codec._get_survey(bands)) + processed = codec._range_compress(processed) + processed, channel_mask = codec.image_padder.forward(processed, bands) + + restored = codec._reverse_range_compress(processed) + restored = codec.image_padder.backward(restored, bands) + restored = codec.rescaler.backward(restored, codec._get_survey(bands)) + + assert channel_mask.sum().item() == len(bands) + assert torch.allclose(restored, image, rtol=1e-5, atol=1e-6) + + @pytest.mark.parametrize("embedding_dim", [5, 10]) @pytest.mark.parametrize("multisurvey_projection_dims", [12, 24]) @pytest.mark.parametrize("hidden_dims", [8, 16])