satellite / test_app_logic.py
prateeksharmacoder's picture
Deploy ZeroGPU compatible code with Sen2SR
8d928e8 verified
Raw History Blame Contribute Delete
3.65 kB
import sys
from pathlib import Path
import numpy as np
import rasterio
from rasterio.transform import Affine
from PIL import Image
sys.path.insert(0, str(Path(__file__).resolve().parent))
from app_zerogpu import run_super_resolution
def create_test_images():
print("Creating test images...")
# 1. 16-bit GeoTIFF (Single multiband file)
tiff_path = "test_16bit.tif"
data = np.random.randint(0, 65535, (3, 64, 64), dtype=np.uint16)
transform = Affine.translation(100.0, 50.0) * Affine.scale(10.0, -10.0)
profile = {
'driver': 'GTiff',
'height': 64,
'width': 64,
'count': 3,
'dtype': 'uint16',
'crs': 'EPSG:4326',
'transform': transform,
}
with rasterio.open(tiff_path, 'w', **profile) as dst:
dst.write(data)
# 2. 16-bit GeoTIFFs (3 separate single-band files)
b02_path = "Sentinel-2_L2A_B02_(Raw).tif"
b03_path = "Sentinel-2_L2A_B03_(Raw).tif"
b04_path = "Sentinel-2_L2A_B04_(Raw).tif"
profile_single = profile.copy()
profile_single['count'] = 1
with rasterio.open(b02_path, 'w', **profile_single) as dst: dst.write(data[2:3])
with rasterio.open(b03_path, 'w', **profile_single) as dst: dst.write(data[1:2])
with rasterio.open(b04_path, 'w', **profile_single) as dst: dst.write(data[0:1])
# 3. PNG
png_path = "test.png"
Image.fromarray(np.random.randint(0, 255, (64, 64, 3), dtype=np.uint8)).save(png_path)
# 4. JPG
jpg_path = "test.jpg"
Image.fromarray(np.random.randint(0, 255, (64, 64, 3), dtype=np.uint8)).save(jpg_path)
return tiff_path, [b02_path, b03_path, b04_path], png_path, jpg_path
class DummyProgress:
def __call__(self, value, desc=None):
pass
def main():
tiff_path, multi_tiff_paths, png_path, jpg_path = create_test_images()
models = ["HAT-SAT (Recommended)", "ESRGAN (Baseline)"]
for img_path in [tiff_path, multi_tiff_paths, png_path, jpg_path]:
for model in models:
print(f"\n--- Testing {img_path} with {model} ---")
try:
orig_img, result, out_png, out_tiff = run_super_resolution(img_path, model, progress=DummyProgress())
# Extract value from gr.update() dict
out_png_val = out_png["value"] if isinstance(out_png, dict) else getattr(out_png, "value", out_png)
out_tiff_val = out_tiff["value"] if isinstance(out_tiff, dict) else getattr(out_tiff, "value", out_tiff)
out_file_path = out_tiff_val if out_tiff_val is not None else out_png_val
print(f"Result size: {result.size}")
print(f"Output file: {out_file_path}")
if img_path == tiff_path:
with rasterio.open(out_file_path) as src:
print("TIFF Metadata:")
print(" CRS:", src.crs)
print(" Transform:", src.transform)
print(" Bounds:", src.bounds)
print(" Dimensions:", src.width, "x", src.height)
# Verify scale
expected_transform = Affine.translation(100.0, 50.0) * Affine.scale(2.5, -2.5) # 10 / 4 = 2.5
assert abs(src.transform.a - expected_transform.a) < 1e-6
assert src.width == 64 * 4
print("Scale and bounds verified!")
except Exception as e:
print(f"Error: {e}")
if __name__ == "__main__":
main()