File size: 3,653 Bytes
8d928e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
80
81
82
83
84
85
86
87
88
89
90
91
92
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()