typical-marine-ecological-feature-recognition-code / seaweed_segmentation_dataset.py
cuibinge's picture
Sync YOLO training and evaluation utilities (part 2)
c2b1b26 verified
Raw
History Blame Contribute Delete
5.57 kB
"""Dataset utilities for seaweed binary segmentation."""
from __future__ import annotations
from pathlib import Path
import albumentations as A
import numpy as np
import rasterio
import torch
from PIL import Image
from torch.utils.data import Dataset
from torchvision import transforms
IMAGE_EXTS = {".tif", ".tiff", ".png", ".jpg", ".jpeg"}
MASK_EXTS = (".png", ".tif", ".tiff", ".jpg", ".jpeg")
class SeaweedSegmentationDataset(Dataset):
def __init__(self, image_dir, mask_dir, transform=None, target_size=256, use_4channel=True):
self.image_dir = Path(image_dir)
self.mask_dir = Path(mask_dir)
self.transform = transform
self.target_size = int(target_size)
self.use_4channel = bool(use_4channel)
if not self.image_dir.exists():
raise FileNotFoundError(f"Image directory not found: {self.image_dir}")
if not self.mask_dir.exists():
raise FileNotFoundError(f"Mask directory not found: {self.mask_dir}")
self.images = sorted(p.name for p in self.image_dir.iterdir() if p.is_file() and p.suffix.lower() in IMAGE_EXTS)
self.normalize_3ch = transforms.Normalize(mean=(0.430, 0.411, 0.296), std=(0.213, 0.156, 0.143))
self.normalize_4ch = transforms.Normalize(mean=(0.430, 0.411, 0.296, 0.350), std=(0.213, 0.156, 0.143, 0.180))
def __len__(self):
return len(self.images)
def find_mask_path(self, image_name: str) -> Path:
stem = Path(image_name).stem
for ext in MASK_EXTS:
for suffix in (ext, ext.upper()):
candidate = self.mask_dir / f"{stem}{suffix}"
if candidate.exists():
return candidate
return self.mask_dir / f"{stem}.png"
@staticmethod
def read_raster(path: Path) -> np.ndarray:
if path.suffix.lower() in {".tif", ".tiff"}:
with rasterio.open(path) as src:
return np.transpose(src.read(), (1, 2, 0))
image = Image.open(path)
return np.asarray(image)
@staticmethod
def extract_432_bands(image: np.ndarray) -> np.ndarray:
if image.ndim == 2:
image = image[:, :, None]
if image.shape[2] >= 4:
return image[:, :, [3, 2, 1]]
output = image[:, :, : min(3, image.shape[2])]
while output.shape[2] < 3:
output = np.concatenate([output, output[:, :, -1:]], axis=2)
return output
@staticmethod
def read_mask(path: Path, fallback_shape: tuple[int, int]) -> np.ndarray:
if not path.exists():
return np.zeros(fallback_shape, dtype=np.uint8)
if path.suffix.lower() in {".tif", ".tiff"}:
with rasterio.open(path) as src:
mask = src.read(1)
else:
mask = np.asarray(Image.open(path).convert("L"))
return (mask > 127).astype(np.uint8)
def __getitem__(self, idx):
img_name = self.images[idx]
img_path = self.image_dir / img_name
mask_path = self.find_mask_path(img_name)
try:
image = self.read_raster(img_path)
if image.ndim == 2:
image = image[:, :, None]
mask = self.read_mask(mask_path, image.shape[:2])
if self.use_4channel and image.shape[2] >= 4:
processed = image[:, :, :4]
processed = torch.from_numpy(processed.astype(np.float32))
if processed.max() > 1.0:
processed = processed / 65535.0
processed = self.normalize_4ch(processed.permute(2, 0, 1))
else:
processed = self.extract_432_bands(image)
processed = torch.from_numpy(processed.astype(np.float32))
if processed.max() > 1.0:
processed = processed / 65535.0
processed = self.normalize_3ch(processed.permute(2, 0, 1))
mask_tensor = torch.from_numpy(mask).long()
if self.target_size != processed.shape[1] or self.target_size != processed.shape[2]:
processed = transforms.Resize((self.target_size, self.target_size), antialias=True)(processed)
mask_tensor = transforms.Resize(
(self.target_size, self.target_size),
interpolation=transforms.InterpolationMode.NEAREST,
)(mask_tensor.unsqueeze(0)).squeeze(0)
if self.transform:
augmented = self.transform(image=processed.permute(1, 2, 0).numpy(), mask=mask_tensor.numpy())
processed = torch.from_numpy(augmented["image"]).permute(2, 0, 1).float()
mask_tensor = torch.from_numpy(augmented["mask"]).long()
return {"image": processed, "mask": mask_tensor, "filename": img_name}
except Exception as exc:
print(f"Error loading {img_name}: {exc}")
return None
def get_train_transforms(target_size=256, use_4channel=True):
return A.Compose(
[
A.Resize(target_size, target_size),
A.HorizontalFlip(p=0.5),
A.VerticalFlip(p=0.3),
A.RandomRotate90(p=0.3),
A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5),
A.RandomBrightnessContrast(p=0.3),
A.GaussNoise(p=0.2),
]
)
def get_val_transforms(target_size=256, use_4channel=True):
return None