Skip to content
Merged
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
2 changes: 1 addition & 1 deletion aion/codecs/preprocessing/image.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
59 changes: 59 additions & 0 deletions tests/codecs/test_image_codec.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])
Expand Down
Loading