Spaces:
Running on Zero
Running on Zero
File size: 4,885 Bytes
2407511 | 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 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 | """Official ChangerEx test preprocessing without Open-CD runtime dependencies."""
from __future__ import annotations
from pathlib import Path
from typing import Union
import numpy as np
import torch
from PIL import Image
from torch.nn import functional as functional
from .config import DEFAULT_MAXIMUM_DIMENSION, RGB_MEAN, RGB_STD, SIZE_DIVISOR
from .schemas import PreprocessingDetails
ImageInput = Union[str, Path, Image.Image, np.ndarray]
class ImageValidationError(ValueError):
pass
def load_rgb_image(image: ImageInput) -> Image.Image:
if isinstance(image, (str, Path)):
path = Path(image).expanduser()
if not path.is_file():
raise ImageValidationError(f"Image file does not exist: {path}")
with Image.open(path) as opened:
return opened.convert("RGB").copy()
if isinstance(image, Image.Image):
return image.convert("RGB").copy()
if isinstance(image, np.ndarray):
array = np.asarray(image)
if array.ndim != 3 or array.shape[2] not in (3, 4):
raise ImageValidationError("NumPy images must have shape HxWx3 or HxWx4")
if array.dtype != np.uint8:
if not np.issubdtype(array.dtype, np.number) or not np.isfinite(array).all():
raise ImageValidationError("Image array must contain finite numeric values")
if array.min() < 0 or array.max() > 255:
raise ImageValidationError("Image array values must be within [0, 255]")
array = array.astype(np.uint8)
return Image.fromarray(array[..., :3])
raise ImageValidationError(f"Unsupported image type: {type(image).__name__}")
def _image_to_tensor(image: Image.Image) -> torch.Tensor:
array = np.asarray(image, dtype=np.uint8).copy()
return torch.from_numpy(array).permute(2, 0, 1).to(dtype=torch.float32)
def _official_resize_shape(width: int, height: int, maximum_dimension: int) -> tuple[int, int, float]:
if maximum_dimension <= 0:
raise ImageValidationError("maximum_dimension must be a positive integer")
scale = maximum_dimension / max(width, height)
resized_width = max(1, int(width * scale + 0.5))
resized_height = max(1, int(height * scale + 0.5))
return resized_width, resized_height, scale
def preprocess_pair(
earlier_image: ImageInput,
later_image: ImageInput,
*,
maximum_dimension: int = DEFAULT_MAXIMUM_DIMENSION,
) -> tuple[torch.Tensor, PreprocessingDetails, Image.Image]:
"""Return normalized NCHW pair, details, and the source-sized later image."""
earlier = load_rgb_image(earlier_image)
later = load_rgb_image(later_image)
if earlier.size != later.size:
raise ImageValidationError(
f"Paired images must have identical dimensions; got {earlier.size} and {later.size}"
)
width, height = earlier.size
if width < 1 or height < 1:
raise ImageValidationError("Images must have non-zero dimensions")
resized_width, resized_height, scale = _official_resize_shape(
width, height, maximum_dimension
)
pair = torch.cat((_image_to_tensor(earlier), _image_to_tensor(later)), dim=0).unsqueeze(0)
if (resized_height, resized_width) != (height, width):
pair = functional.interpolate(
pair, size=(resized_height, resized_width), mode="bilinear", align_corners=False
)
mean = pair.new_tensor(RGB_MEAN * 2).view(1, 6, 1, 1)
std = pair.new_tensor(RGB_STD * 2).view(1, 6, 1, 1)
pair = (pair - mean) / std
padded_width = ((resized_width + SIZE_DIVISOR - 1) // SIZE_DIVISOR) * SIZE_DIVISOR
padded_height = ((resized_height + SIZE_DIVISOR - 1) // SIZE_DIVISOR) * SIZE_DIVISOR
pad_right = padded_width - resized_width
pad_bottom = padded_height - resized_height
if pad_right or pad_bottom:
pair = functional.pad(pair, (0, pad_right, 0, pad_bottom), value=0.0)
details = PreprocessingDetails(
source_size=(width, height),
resized_size=(resized_width, resized_height),
padded_size=(padded_width, padded_height),
scale=scale,
channel_order="RGB",
pair_order="earlier RGB, then later RGB",
input_range="float32 source values in [0, 255] before normalization",
mean=RGB_MEAN,
std=RGB_STD,
resize_interpolation="bilinear, align_corners=False",
pad=(0, 0, pad_right, pad_bottom),
size_divisor=SIZE_DIVISOR,
)
return pair.contiguous(), details, later
def preprocessing_tensor_summary(tensor: torch.Tensor) -> dict[str, object]:
data = tensor.detach().cpu().contiguous()
return {
"shape": list(data.shape),
"dtype": str(data.dtype),
"sum": float(data.sum()),
"mean": float(data.mean()),
"std": float(data.std()),
"min": float(data.min()),
"max": float(data.max()),
}
|