BiliSakura's picture
Upload CROMA transformers checkpoints with model card metadata
930dfab verified
Raw
History Blame
4.64 kB
# Copyright 2023 The CROMA Authors and The HuggingFace Inc. team.
"""Custom CROMA dual-modality feature extraction pipeline."""
from typing import Any, Optional, Union
from transformers.image_utils import load_image
from transformers.pipelines.base import GenericTensor, Pipeline
from transformers.pipelines.image_feature_extraction import ImageFeatureExtractionPipeline
class CROMAImageFeatureExtractionPipeline(ImageFeatureExtractionPipeline):
"""SAR-optical feature extraction for CROMA."""
def _sanitize_parameters(
self,
image_processor_kwargs=None,
return_tensors=None,
pool=None,
sar_images=None,
optical_images=None,
use_8_bit=None,
output_mode=None,
**kwargs,
):
preprocess_params = {} if image_processor_kwargs is None else dict(image_processor_kwargs)
if sar_images is not None:
preprocess_params["sar_images"] = sar_images
if optical_images is not None:
preprocess_params["optical_images"] = optical_images
if use_8_bit is not None:
preprocess_params["use_8_bit"] = use_8_bit
if "timeout" in kwargs:
preprocess_params["timeout"] = kwargs["timeout"]
postprocess_params = {}
if pool is not None:
postprocess_params["pool"] = pool
if return_tensors is not None:
postprocess_params["return_tensors"] = return_tensors
if output_mode is not None:
postprocess_params["output_mode"] = output_mode
return preprocess_params, {}, postprocess_params
def preprocess(
self,
image=None,
sar_images=None,
optical_images=None,
timeout=None,
**image_processor_kwargs,
) -> dict[str, GenericTensor]:
if sar_images is None:
sar_images = image_processor_kwargs.pop("sar_images", None)
if optical_images is None:
optical_images = image_processor_kwargs.pop("optical_images", None)
if sar_images is not None and not isinstance(sar_images, (list, tuple)) and not hasattr(sar_images, "shape"):
sar_images = load_image(sar_images, timeout=timeout)
if optical_images is not None and not isinstance(optical_images, (list, tuple)) and not hasattr(
optical_images, "shape"
):
optical_images = load_image(optical_images, timeout=timeout)
if image is not None and optical_images is None and sar_images is None:
if not isinstance(image, (list, tuple)) and not hasattr(image, "shape"):
image = load_image(image, timeout=timeout)
optical_images = image
model_inputs = self.image_processor.preprocess(
images=optical_images if optical_images is not None else sar_images,
sar_images=sar_images,
optical_images=optical_images,
return_tensors="pt",
**image_processor_kwargs,
)
model_inputs = model_inputs.to(self.dtype)
return model_inputs
def postprocess(self, model_outputs, pool=None, return_tensors=False, output_mode="primary"):
pool = pool if pool is not None else False
output_mode = output_mode or "primary"
if output_mode == "sar":
outputs = model_outputs["sar_pooler_output"] if pool else model_outputs["sar_hidden_states"]
elif output_mode == "optical":
outputs = model_outputs["optical_pooler_output"] if pool else model_outputs["optical_hidden_states"]
elif output_mode == "joint":
outputs = model_outputs["joint_pooler_output"] if pool else model_outputs["joint_hidden_states"]
elif pool:
outputs = model_outputs["pooler_output"]
else:
outputs = model_outputs[0]
if return_tensors:
return outputs
return outputs.tolist()
def __call__(
self,
images: Optional[Union[str, Any, list[Any]]] = None,
sar_images: Optional[Union[str, Any, list[Any]]] = None,
optical_images: Optional[Union[str, Any, list[Any]]] = None,
**kwargs: Any,
) -> list[Any]:
if images is None and sar_images is None and optical_images is None:
raise ValueError("Provide at least one of `images`, `sar_images`, or `optical_images`.")
if images is None:
images = optical_images if optical_images is not None else sar_images
return Pipeline.__call__(self, images, sar_images=sar_images, optical_images=optical_images, **kwargs)
__all__ = ["CROMAImageFeatureExtractionPipeline"]