satellite / geotiff_test.py
prateeksharmacoder's picture
Deploy ZeroGPU compatible code with Sen2SR
8d928e8 verified
Raw History Blame Contribute Delete
12.1 kB
"""
Phase 4 -- GeoTIFF Geospatial Verification Test
satellite_sr_deploy
Tests:
1. save_sr_geotiff() writes a valid GeoTIFF from a PIL image + profile
2. CRS is preserved exactly in the output
3. Geographic bounds are preserved (same extent as input)
4. Output pixel dimensions are exactly 4x the input
5. Affine transform is correctly scaled (pixel size = input/4)
6. Band ordering is R=1, G=2, B=3
7. Pixel values are preserved correctly (uint8 RGB)
8. Round-trip: create synthetic GeoTIFF -> read -> SR -> save -> verify
9. Verify against an existing real GeoTIFF (from Colab outputs)
10. normalize_rgb() clamps properly and returns uint8
Run from satellite_sr_deploy/:
$env:PYTHONIOENCODING="utf-8"; .venv\\Scripts\\python.exe geotiff_test.py
"""
import sys
import math
from pathlib import Path
import numpy as np
from PIL import Image
import rasterio
from rasterio.transform import Affine
from rasterio.crs import CRS
ROOT = Path(__file__).resolve().parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
# Import only geotiff.py (no model loading needed for most tests)
from inference.geotiff import save_sr_geotiff, normalize_rgb
results = []
OUTDIR = ROOT / "outputs" / "geotiff_test"
OUTDIR.mkdir(parents=True, exist_ok=True)
PASS_SYM = "[PASS]"
FAIL_SYM = "[FAIL]"
def record(name, passed, note=""):
tag = PASS_SYM if passed else FAIL_SYM
msg = f"{tag} {name}"
if note:
msg += f" -- {note}"
print(msg)
results.append((name, passed, note))
def make_synthetic_geotiff(path, width=50, height=40, crs_epsg=4326):
"""
Create a small synthetic single-band GeoTIFF with known transform/CRS.
Returns the profile used to write it.
"""
# Mimic a Sentinel-2 scene over India (approximate EPSG:4326 coords)
west, north = 77.0, 13.0 # top-left corner
pixel_size_x = 0.0001 # degrees/pixel (~10 m)
pixel_size_y = -0.0001 # negative (north-to-south)
transform = Affine.translation(west, north) * Affine.scale(pixel_size_x, pixel_size_y)
profile = {
"driver": "GTiff",
"height": height,
"width": width,
"count": 1,
"dtype": "uint16",
"crs": CRS.from_epsg(crs_epsg),
"transform": transform,
}
data = np.random.randint(0, 3000, (height, width), dtype=np.uint16)
with rasterio.open(path, "w", **profile) as dst:
dst.write(data, 1)
return profile
# ================================================================
# GROUP 1 -- save_sr_geotiff: basic output verification
# ================================================================
print("=" * 65)
print("PHASE 4 -- GEOTIFF GEOSPATIAL VERIFICATION")
print("=" * 65)
print()
print("-" * 65)
print("GROUP 1 -- save_sr_geotiff: CRS, transform, dimensions, bands")
print("-" * 65)
# Build a synthetic source profile (50x40, EPSG:4326)
SRC_W, SRC_H = 50, 40
SCALE = 4
EXPECTED_W = SRC_W * SCALE
EXPECTED_H = SRC_H * SCALE
west, north = 77.0, 13.0
px_x, px_y = 0.0001, -0.0001
src_transform = Affine.translation(west, north) * Affine.scale(px_x, px_y)
source_profile = {
"driver": "GTiff",
"height": SRC_H,
"width": SRC_W,
"count": 1,
"dtype": "uint16",
"crs": CRS.from_epsg(4326),
"transform": src_transform,
}
# Synthetic SR image (4x size, gradient RGB)
sr_arr = np.zeros((EXPECTED_H, EXPECTED_W, 3), dtype=np.uint8)
sr_arr[:, :, 0] = np.linspace(0, 255, EXPECTED_W, dtype=np.uint8) # R
sr_arr[:, :, 1] = np.linspace(0, 200, EXPECTED_H, dtype=np.uint8).reshape(-1, 1) # G
sr_arr[:, :, 2] = 128 # B flat
sr_pil = Image.fromarray(sr_arr, mode="RGB")
output_tif = OUTDIR / "synthetic_sr_4x.tif"
try:
save_sr_geotiff(sr_pil, source_profile, output_tif, scale=SCALE)
saved_ok = True
except Exception as e:
saved_ok = False
record("save_sr_geotiff() runs without error", False, str(e))
if saved_ok:
record("save_sr_geotiff() runs without error", True)
with rasterio.open(output_tif) as out:
out_profile = out.profile
out_transform = out.transform
out_crs = out.crs
out_w = out.width
out_h = out.height
out_count = out.count
out_dtype = out.dtypes[0]
band_r = out.read(1)
band_g = out.read(2)
band_b = out.read(3)
# 1a. Dimensions
record(
"Output dimensions are 4x input",
out_w == EXPECTED_W and out_h == EXPECTED_H,
f"expected {EXPECTED_W}x{EXPECTED_H}, got {out_w}x{out_h}"
)
# 1b. CRS preserved
record(
"CRS preserved (EPSG:4326)",
out_crs.to_epsg() == 4326,
f"got CRS: {out_crs}"
)
# 1c. Transform: pixel size scaled by 1/4
expected_px_x = px_x / SCALE
expected_px_y = px_y / SCALE
tol = 1e-12
px_x_ok = abs(out_transform.a - expected_px_x) < tol
px_y_ok = abs(out_transform.e - expected_px_y) < tol
record(
"Pixel size scaled to 1/4 of source (x)",
px_x_ok,
f"expected {expected_px_x:.8f}, got {out_transform.a:.8f}"
)
record(
"Pixel size scaled to 1/4 of source (y)",
px_y_ok,
f"expected {expected_px_y:.8f}, got {out_transform.e:.8f}"
)
# 1d. Origin (top-left corner) unchanged
origin_ok = (
abs(out_transform.c - west) < tol and
abs(out_transform.f - north) < tol
)
record(
"Geographic origin (top-left) unchanged",
origin_ok,
f"expected ({west}, {north}), got ({out_transform.c:.6f}, {out_transform.f:.6f})"
)
# 1e. Geographic extent / bounds preserved
src_right = west + SRC_W * px_x
src_bottom = north + SRC_H * px_y
out_right = out_transform.c + out_w * out_transform.a
out_bottom = out_transform.f + out_h * out_transform.e
bounds_ok = (
abs(out_right - src_right) < tol and
abs(out_bottom - src_bottom) < tol
)
record(
"Geographic bounds (extent) preserved",
bounds_ok,
f"src right={src_right:.6f}, out right={out_right:.6f} | "
f"src bottom={src_bottom:.6f}, out bottom={out_bottom:.6f}"
)
# 1f. Band count
record("Band count is 3 (RGB)", out_count == 3, f"got {out_count}")
# 1g. dtype is uint8
record("Output dtype is uint8", out_dtype == "uint8", f"got {out_dtype}")
# 1h. Band ordering R=1, G=2, B=3
# Check representative pixels in the gradient
r_corner = int(sr_arr[0, EXPECTED_W - 1, 0]) # max-R corner
g_corner = int(sr_arr[EXPECTED_H - 1, 0, 1]) # max-G corner
b_mid = int(sr_arr[EXPECTED_H // 2, EXPECTED_W // 2, 2]) # flat B=128
r_read = int(band_r[0, EXPECTED_W - 1])
g_read = int(band_g[EXPECTED_H - 1, 0])
b_read = int(band_b[EXPECTED_H // 2, EXPECTED_W // 2])
record(
"Band 1 = Red channel",
r_read == r_corner,
f"expected R={r_corner}, got R={r_read}"
)
record(
"Band 2 = Green channel",
g_read == g_corner,
f"expected G={g_corner}, got G={g_read}"
)
record(
"Band 3 = Blue (flat 128)",
b_read == 128,
f"expected B=128, got B={b_read}"
)
print()
# ================================================================
# GROUP 2 -- normalize_rgb: clipping and dtype
# ================================================================
print("-" * 65)
print("GROUP 2 -- normalize_rgb: output range and dtype")
print("-" * 65)
# Typical Sentinel-2 DN range (0-10000+ reflectance)
test_raw = np.random.randint(0, 10000, (40, 50, 3), dtype=np.uint16).astype(np.float32)
try:
norm = normalize_rgb(test_raw.copy())
record(
"normalize_rgb returns uint8",
norm.dtype == np.uint8,
f"got dtype={norm.dtype}"
)
record(
"normalize_rgb output range [0, 255]",
norm.min() >= 0 and norm.max() <= 255,
f"min={norm.min()}, max={norm.max()}"
)
record(
"normalize_rgb shape preserved (H, W, 3)",
norm.shape == test_raw.shape,
f"expected {test_raw.shape}, got {norm.shape}"
)
except Exception as e:
record("normalize_rgb runs without error", False, str(e))
# Edge case: single-value band (all pixels same value)
flat = np.full((10, 10, 3), 5000, dtype=np.float32)
try:
norm_flat = normalize_rgb(flat)
record(
"normalize_rgb handles flat (constant) band",
norm_flat.dtype == np.uint8,
f"min={norm_flat.min()}, max={norm_flat.max()}"
)
except Exception as e:
record("normalize_rgb handles flat band", False, str(e))
print()
# ================================================================
# GROUP 3 -- save_sr_geotiff with numpy array input (not PIL)
# ================================================================
print("-" * 65)
print("GROUP 3 -- save_sr_geotiff accepts numpy array input")
print("-" * 65)
np_input_tif = OUTDIR / "numpy_input_sr.tif"
try:
save_sr_geotiff(sr_arr, source_profile, np_input_tif, scale=SCALE)
with rasterio.open(np_input_tif) as f:
dims_ok = f.width == EXPECTED_W and f.height == EXPECTED_H
record("save_sr_geotiff accepts numpy ndarray", dims_ok,
f"{EXPECTED_W}x{EXPECTED_H} ok={dims_ok}")
except Exception as e:
record("save_sr_geotiff accepts numpy ndarray", False, str(e))
print()
# ================================================================
# GROUP 4 -- Verify existing real GeoTIFF from Colab outputs
# ================================================================
print("-" * 65)
print("GROUP 4 -- Verify existing Colab-produced GeoTIFFs")
print("-" * 65)
REAL_TIFS = [
Path(r"D:\downloads\satellite_resolution_increaser\satellite_resolution_increaser\outputs\hatsat\results\sentinel2_hatsat_4x.tif"),
Path(r"D:\downloads\satellite_resolution_increaser\satellite_resolution_increaser\outputs\hatsat\results\sentinel2_hatsat_4x_georef.tif"),
Path(r"D:\downloads\satellite_resolution_increaser\satellite_resolution_increaser\outputs\esrgan\results\sentinel2_esrgan_4x_georef.tif"),
]
for tif_path in REAL_TIFS:
label = tif_path.name
if not tif_path.exists():
record(f"{label} -- file exists", False, "not found on disk")
continue
try:
with rasterio.open(tif_path) as src:
crs_epsg = src.crs.to_epsg() if src.crs else None
band_count = src.count
dtype = src.dtypes[0]
w, h = src.width, src.height
t = src.transform
bounds = src.bounds
record(
f"{label} -- opens without error",
True,
f"{w}x{h}, {band_count} bands, {dtype}, EPSG:{crs_epsg}"
)
record(
f"{label} -- has CRS",
crs_epsg is not None,
f"EPSG:{crs_epsg}"
)
record(
f"{label} -- has 3 bands (RGB)",
band_count == 3,
f"bands={band_count}"
)
record(
f"{label} -- dimensions multiple of 4 (4x SR)",
w % 4 == 0 and h % 4 == 0,
f"{w}x{h}"
)
# Check pixel size is positive x / negative y (north-up)
north_up = t.a > 0 and t.e < 0
record(
f"{label} -- north-up orientation",
north_up,
f"x_res={t.a:.8f}, y_res={t.e:.8f}"
)
except Exception as e:
record(f"{label} -- opens without error", False, str(e))
print()
# ================================================================
# SUMMARY
# ================================================================
print("=" * 65)
print("PHASE 4 SUMMARY")
print("=" * 65)
total = len(results)
passed = sum(1 for _, p, _ in results if p)
failed = total - passed
for name, p, note in results:
tag = "+" if p else "x"
print(f" [{tag}] {name:<50} {note}")
print()
print(f" Result: {passed}/{total} passed", end="")
if failed:
print(f" -- {failed} FAILED")
else:
print(" -- ALL PASSED")
print("=" * 65)
print(f" Output GeoTIFFs: {OUTDIR}")
print("=" * 65)
sys.exit(0 if failed == 0 else 1)