satellite / inference /stitching.py
prateeksharmacoder's picture
Deploy ZeroGPU codebase
db89723
Raw
History Blame Contribute Delete
4.33 kB
import numpy as np
from PIL import Image
def create_blend_weight(
height,
width,
overlap
):
"""
Create a smooth 2D blending weight.
Pixels near the center receive higher weight.
Pixels near overlapping boundaries receive lower weight.
"""
if overlap <= 0:
return np.ones(
(height, width),
dtype=np.float32
)
# Horizontal weights
wx = np.ones(width, dtype=np.float32)
transition = min(overlap, width // 2)
if transition > 0:
ramp = np.linspace(
0.01,
1.0,
transition,
dtype=np.float32
)
wx[:transition] = ramp
wx[-transition:] = ramp[::-1]
# Vertical weights
wy = np.ones(height, dtype=np.float32)
transition = min(overlap, height // 2)
if transition > 0:
ramp = np.linspace(
0.01,
1.0,
transition,
dtype=np.float32
)
wy[:transition] = ramp
wy[-transition:] = ramp[::-1]
return wy[:, None] * wx[None, :]
def stitch_tiles(
sr_tiles,
scale=4,
original_size=None,
overlap=32
):
"""
Stitch overlapping super-resolution tiles.
Parameters
----------
sr_tiles : list of dictionaries
Each dictionary must contain:
image
x
y
where x/y are coordinates in the ORIGINAL image.
scale : int
Super-resolution scale factor.
original_size : tuple
(width, height) of original image.
overlap : int
Overlap in ORIGINAL-image pixels.
Returns
-------
PIL.Image
Final stitched SR image.
"""
if not sr_tiles:
raise ValueError("No SR tiles supplied.")
if original_size is None:
raise ValueError(
"original_size must be provided."
)
original_width, original_height = original_size
output_width = original_width * scale
output_height = original_height * scale
# Accumulate weighted RGB values
canvas = np.zeros(
(
output_height,
output_width,
3
),
dtype=np.float32
)
# Accumulate weights
weights = np.zeros(
(
output_height,
output_width
),
dtype=np.float32
)
sr_overlap = overlap * scale
for tile_info in sr_tiles:
image = tile_info["image"]
if not isinstance(image, Image.Image):
image = Image.fromarray(image)
image = image.convert("RGB")
tile = np.asarray(
image,
dtype=np.float32
)
tile_height, tile_width = tile.shape[:2]
# Original-image coordinates → SR coordinates
x = int(tile_info["x"] * scale)
y = int(tile_info["y"] * scale)
# Do not allow the tile to exceed final canvas
valid_width = min(
tile_width,
output_width - x
)
valid_height = min(
tile_height,
output_height - y
)
if valid_width <= 0 or valid_height <= 0:
continue
tile = tile[
:valid_height,
:valid_width
]
# --------------------------------------------------
# Build blending weight
# --------------------------------------------------
weight = create_blend_weight(
valid_height,
valid_width,
sr_overlap
)
# --------------------------------------------------
# Accumulate
# --------------------------------------------------
canvas[
y:y + valid_height,
x:x + valid_width
] += tile * weight[..., None]
weights[
y:y + valid_height,
x:x + valid_width
] += weight
# ------------------------------------------------------
# Normalize overlapping pixels
# ------------------------------------------------------
weights = np.maximum(
weights,
1e-8
)
canvas /= weights[..., None]
canvas = np.clip(
canvas,
0,
255
).round().astype(np.uint8)
return Image.fromarray(
canvas,
mode="RGB"
)