Spaces:
Paused
Paused
File size: 2,512 Bytes
a325975 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 | from pathlib import Path
import numpy as np
import pytest
import rasterio
from PIL import Image
from rasterio.transform import from_origin
from geochat.inference.sar import load_image_input, raster_to_rgb, sar1_to_rgb
def write_raster(path: Path, bands: np.ndarray) -> Path:
bands = np.asarray(bands)
if bands.ndim == 2:
bands = bands[None, ...]
with rasterio.open(
path,
"w",
driver="GTiff",
width=bands.shape[2],
height=bands.shape[1],
count=bands.shape[0],
dtype="float32",
transform=from_origin(0, 1, 1, 1),
) as destination:
destination.write(bands.astype(np.float32))
return path
def test_png_and_jpeg_inputs_become_rgb(tmp_path):
image = Image.new("RGBA", (4, 3), (10, 20, 30, 255))
png_path = tmp_path / "image.png"
jpeg_path = tmp_path / "image.jpg"
image.save(png_path)
image.convert("RGB").save(jpeg_path)
assert load_image_input(png_path).mode == "RGB"
assert load_image_input(jpeg_path).mode == "RGB"
def test_rgb_tiff_is_loaded_with_rasterio(tmp_path):
image = raster_to_rgb(write_raster(tmp_path / "rgb.tif", np.arange(48).reshape(3, 4, 4)))
assert image.mode == "RGB"
assert image.size == (4, 4)
def test_invalid_tiff_has_clean_error(tmp_path):
path = tmp_path / "invalid.tif"
path.write_bytes(b"not-a-raster")
with pytest.raises(ValueError, match="Unable to decode"):
raster_to_rgb(path)
def test_unsupported_multispectral_configuration_is_rejected(tmp_path):
path = write_raster(tmp_path / "multispectral.tif", np.ones((4, 2, 2)))
with pytest.raises(ValueError, match="no explicit RGB band mapping"):
raster_to_rgb(path)
def test_sar_vv_vh_pseudo_rgb(tmp_path):
vv = write_raster(tmp_path / "VV.tiff", np.arange(16).reshape(4, 4))
vh = write_raster(tmp_path / "VH.tiff", np.arange(16, 32).reshape(4, 4))
image = sar1_to_rgb(vv, vh)
assert image.mode == "RGB"
assert np.asarray(image).shape == (4, 4, 3)
def test_sar_dimensions_must_match(tmp_path):
vv = write_raster(tmp_path / "VV.tif", np.ones((4, 4)))
vh = write_raster(tmp_path / "VH.tif", np.ones((3, 4)))
with pytest.raises(ValueError, match="dimensions do not match"):
sar1_to_rgb(vv, vh)
def test_missing_vh_is_rejected(tmp_path):
vv = write_raster(tmp_path / "VV.tif", np.ones((4, 4)))
with pytest.raises((FileNotFoundError, ValueError)):
sar1_to_rgb(vv, tmp_path / "missing-vh.tif")
|