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..dfb6e63 100644 --- a/tests/codecs/test_image_codec.py +++ b/tests/codecs/test_image_codec.py @@ -3,9 +3,68 @@ 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) + + +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])