crop-segmentation / predictor.py
meabd's picture
Upload 2 files
6ba8390 verified
Raw
History Blame Contribute Delete
9.23 kB
"""
predictor.py
------------
Model loading and inference for the Crop vs No-Crop SegFormer-B2 app.
Architecture is rebuilt from the `nvidia/mit-b2` backbone config with
num_labels=2, then the fine-tuned `best.pt` state dict is loaded on top.
Robust loader: auto-detects common key-naming differences, matches tensors by
name AND shape, and reports detailed diagnostics (including whether the
segmentation HEAD/classifier loaded) so an all-one-class prediction can be
debugged precisely.
"""
from __future__ import annotations
import logging
import os
import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image
from torchvision import transforms
from transformers import SegformerConfig, SegformerForSemanticSegmentation
logger = logging.getLogger(__name__)
# --------------------------------------------------------------------------- #
# Configuration
# --------------------------------------------------------------------------- #
MODEL_NAME = os.environ.get("MODEL_NAME", "nvidia/mit-b2")
NUM_LABELS = 2
INPUT_SIZE = int(os.environ.get("INPUT_SIZE", "512"))
IMAGENET_MEAN = (0.485, 0.456, 0.406)
IMAGENET_STD = (0.229, 0.224, 0.225)
WEIGHTS_FILENAME = os.environ.get("WEIGHTS_FILENAME", "best.pt")
HF_REPO_ID = os.environ.get("HF_REPO_ID", "").strip()
HF_TOKEN = os.environ.get("HF_TOKEN", "").strip() or None
class ModelLoadError(Exception):
"""Raised when the model or its weights cannot be loaded."""
class InferenceError(Exception):
"""Raised when inference fails for a given image."""
def _resolve_weights_path() -> str:
"""Return a local path to best.pt (download from HF Hub if HF_REPO_ID set)."""
if HF_REPO_ID:
try:
from huggingface_hub import hf_hub_download
return hf_hub_download(
repo_id=HF_REPO_ID, filename=WEIGHTS_FILENAME, token=HF_TOKEN
)
except Exception as exc: # noqa: BLE001
raise ModelLoadError(
f"Failed to download '{WEIGHTS_FILENAME}' from '{HF_REPO_ID}': {exc}"
) from exc
here = os.path.join(os.path.dirname(__file__), WEIGHTS_FILENAME)
if os.path.exists(here):
return here
if os.path.exists(WEIGHTS_FILENAME):
return WEIGHTS_FILENAME
raise ModelLoadError(
f"Weights file '{WEIGHTS_FILENAME}' not found. Commit it to the repo "
"or set the HF_REPO_ID environment variable."
)
def _extract_state_dict(checkpoint):
if isinstance(checkpoint, dict):
for key in ("state_dict", "model_state_dict", "model"):
inner = checkpoint.get(key)
if isinstance(inner, dict):
return inner
return checkpoint
def _strip_module(state_dict):
cleaned = {}
for key, value in state_dict.items():
new_key = key
while new_key.startswith("module."):
new_key = new_key[len("module."):]
cleaned[new_key] = value
return cleaned
def _align_state_dict(ckpt_sd, model_sd):
"""
Match checkpoint tensors to model tensors by name (with a few common
transforms) AND shape. Returns (aligned_dict, info) where info reports
matched count, the model layers still MISSING weights, the checkpoint
layers that were UNUSED, and whether the classifier head loaded.
"""
ckpt_sd = _strip_module(ckpt_sd)
transforms_to_try = [
lambda k: k, # as-is
lambda k: "segformer." + k, # add prefix
lambda k: k[len("segformer."):] # strip prefix
if k.startswith("segformer.") else None,
lambda k: k[len("model."):] if k.startswith("model.") else None,
]
best = {}
used_ckpt_keys = set()
for key, value in ckpt_sd.items():
for transform in transforms_to_try:
new_key = transform(key)
if (
new_key
and new_key in model_sd
and new_key not in best
and tuple(model_sd[new_key].shape) == tuple(value.shape)
):
best[new_key] = value
used_ckpt_keys.add(key)
break
missing = sorted(set(model_sd) - set(best))
unexpected = sorted(set(ckpt_sd) - used_ckpt_keys)
# The decode-head classifier is what turns features into crop/background.
classifier_loaded = any("classifier" in k for k in best)
info = {
"model_total": len(model_sd),
"ckpt_total": len(ckpt_sd),
"matched": len(best),
"missing": missing,
"unexpected": unexpected,
"classifier_loaded": classifier_loaded,
}
return best, info
class CropSegmenter:
"""Loads the fine-tuned SegFormer-B2 model and performs inference."""
def __init__(self) -> None:
self.device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu"
)
logger.info("Using device: %s", self.device)
self.last_logit_summary: dict | None = None
# Build architecture from config (weights come from best.pt)
try:
seg_config = SegformerConfig.from_pretrained(
MODEL_NAME, num_labels=NUM_LABELS
)
self.model = SegformerForSemanticSegmentation(seg_config)
except Exception as exc: # noqa: BLE001
try:
self.model = SegformerForSemanticSegmentation.from_pretrained(
MODEL_NAME, num_labels=NUM_LABELS, ignore_mismatched_sizes=True
)
except Exception as exc2: # noqa: BLE001
raise ModelLoadError(
f"Could not build SegFormer from '{MODEL_NAME}': {exc2}"
) from exc
weights_path = _resolve_weights_path()
try:
try:
checkpoint = torch.load(weights_path, map_location=self.device)
except Exception:
checkpoint = torch.load(
weights_path, map_location=self.device, weights_only=False
)
except Exception as exc: # noqa: BLE001
raise ModelLoadError(
f"Failed to read weights file '{weights_path}': {exc}"
) from exc
ckpt_sd = _extract_state_dict(checkpoint)
aligned, self.load_info = _align_state_dict(
ckpt_sd, self.model.state_dict()
)
self.model.load_state_dict(aligned, strict=False)
logger.info(
"Weights matched %d/%d (ckpt had %d). classifier_loaded=%s",
self.load_info["matched"],
self.load_info["model_total"],
self.load_info["ckpt_total"],
self.load_info["classifier_loaded"],
)
if self.load_info["matched"] == 0:
raise ModelLoadError(
"None of the weights in best.pt matched the SegFormer-B2 "
"architecture. The checkpoint is likely from a different model."
)
self.model.to(self.device)
self.model.eval()
def _build_preprocess(self, size: int, normalization: str):
steps = [transforms.Resize((size, size)), transforms.ToTensor()]
if normalization == "imagenet":
steps.append(transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD))
elif normalization == "minus_one_one":
steps.append(transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)))
# "zero_one" -> ToTensor only (pixels in [0, 1])
return transforms.Compose(steps)
@torch.no_grad()
def predict(
self,
image: Image.Image,
input_size: int | None = None,
swap_classes: bool = False,
normalization: str = "imagenet",
) -> np.ndarray:
"""
Run segmentation. Returns a uint8 mask (H, W) in {0, 1} at the original
resolution. Also stores per-call logit stats in self.last_logit_summary.
"""
try:
size = int(input_size or INPUT_SIZE)
if image.mode != "RGB":
image = image.convert("RGB")
orig_w, orig_h = image.size
preprocess = self._build_preprocess(size, normalization)
tensor = preprocess(image).unsqueeze(0).to(self.device)
logits = self.model(pixel_values=tensor).logits # (1, 2, h/4, w/4)
upsampled = F.interpolate(
logits, size=(orig_h, orig_w),
mode="bilinear", align_corners=False,
)
class0 = upsampled[:, 0]
class1 = upsampled[:, 1]
self.last_logit_summary = {
"frac_pixels_pred_crop": float((class1 > class0).float().mean()),
"mean_gap_crop_minus_bg": float((class1 - class0).mean().item()),
"logit_min": float(upsampled.min().item()),
"logit_max": float(upsampled.max().item()),
}
mask = upsampled.argmax(dim=1).squeeze(0).to("cpu").numpy().astype(np.uint8)
if swap_classes:
mask = (1 - mask).astype(np.uint8)
return mask
except Exception as exc: # noqa: BLE001
raise InferenceError(f"Inference failed: {exc}") from exc