Spaces:
Paused
Paused
Download tests/test_raster_preprocessing.py from Bireswar26/GeoChat: direct link, hf CLI and curl.
- Browser
- Download file 2.51 kB
-
https://huggingface.co/spaces/Bireswar26/GeoChat/resolve/main/tests/test_raster_preprocessing.py
- Command line
-
hf download hf://spaces/Bireswar26/GeoChat/tests/test_raster_preprocessing.py
-
curl -L -o test_raster_preprocessing.py https://huggingface.co/spaces/Bireswar26/GeoChat/resolve/main/tests/test_raster_preprocessing.py
2.51 kB
| 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") | |