Abrir en Colab

Flujos de trabajo para la segmentación de imágenes con CLIPSeg#

La segmentación de imágenes es otra herramienta útil para encontrar patrones en imágenes. Puede utilizarse para calcular qué porcentaje de una imagen corresponde a una característica determinada, como el cielo o el suelo, o para eliminar el fondo de una imagen.

En este cuaderno exploraremos la segmentación de imágenes mediante modelos de código abierto y datos de ciencia participativa provenientes de plataformas como iNaturalist y NASA GLOBE Land Cover.

1. Segmentación general con el modelo Segment Anything (SAM)#

Comencemos a explorar la segmentación de imágenes con el modelo Segment Anything. Este modelo tiene una amplia variedad de aplicaciones y puede segmentar las distintas partes de una imagen para diferenciar objetos y regiones.

# Instalar los paquetes necesarios
!pip install -q segment-anything opencv-python matplotlib pillow requests
!pip install -q git+https://github.com/facebookresearch/segment-anything.git

import torch
import numpy as np
import matplotlib.pyplot as plt
from PIL import Image, ImageDraw
import requests
from io import BytesIO
import cv2
from segment_anything import (
    sam_model_registry,
    SamAutomaticMaskGenerator,
    SamPredictor
)

# Cargar el modelo Segment Anything
print("Cargando el modelo SAM. Este proceso puede tardar un momento...")

# Descargar el punto de control del modelo
!wget -q https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth

# Inicializar SAM
model_type = "vit_h"
checkpoint = "sam_vit_h_4b8939.pth"
device = "cuda" if torch.cuda.is_available() else "cpu"

sam = sam_model_registry[model_type](checkpoint=checkpoint)
sam.to(device=device)
# Crear el generador de máscaras para la segmentación automática
mask_generator = SamAutomaticMaskGenerator(
    model=sam,
    points_per_side=32,
    pred_iou_thresh=0.86,
    stability_score_thresh=0.92,
    crop_n_layers=1,
    crop_n_points_downscale_factor=2,
    min_mask_region_area=100,
)

print("Modelo cargado")

def load_image(url):
    """Carga una imagen a partir de una URL."""
    try:
        response = requests.get(url, stream=True, timeout=30)
        response.raise_for_status()
        image = Image.open(BytesIO(response.content)).convert("RGB")
        return image
    except Exception as e:
        print(f"Error al cargar la imagen: {e}")
        return None

def show_anns(anns, ax):
    """
    Muestra las máscaras de segmentación con diferentes colores.

    Argumentos:
        anns: lista de diccionarios de anotaciones generados por SAM
        ax: eje de Matplotlib sobre el que se dibujarán las máscaras
    """
    if len(anns) == 0:
        return

    # Ordenar las máscaras por área, de mayor a menor
    sorted_anns = sorted(
        anns,
        key=lambda x: x["area"],
        reverse=True
    )

    # Crear una capa de máscaras con colores
    img = np.ones(
        (
            sorted_anns[0]["segmentation"].shape[0],
            sorted_anns[0]["segmentation"].shape[1],
            4
        )
    )
    img[:, :, 3] = 0  # Fondo transparente

    for ann in sorted_anns:
        mask = ann["segmentation"]

        # Generar un color aleatorio para cada segmento
        color_mask = np.concatenate([
            np.random.random(3),
            [0.5]
        ])
        img[mask] = color_mask

    ax.imshow(img)

def create_mask_overlay(image, masks):
    """
    Crea una visualización con las máscaras de segmentación superpuestas
    sobre la imagen original.

    Argumentos:
        image: imagen de PIL
        masks: lista de diccionarios de máscaras generados por SAM

    Devuelve:
        imagen de PIL con las máscaras superpuestas
    """
    # Convertir la imagen de PIL en un arreglo de NumPy
    img_array = np.array(image)

    # Crear la capa superpuesta
    overlay = img_array.copy()

    # Ordenar las máscaras por área
    sorted_masks = sorted(
        masks,
        key=lambda x: x["area"],
        reverse=True
    )

    for mask_dict in sorted_masks:
        mask = mask_dict["segmentation"]

        # Generar un color aleatorio para cada segmento
        color = np.random.randint(0, 255, size=3)
        overlay[mask] = overlay[mask] * 0.5 + color * 0.5

    return Image.fromarray(overlay.astype(np.uint8))

def run_segmentation(image_urls):
    """
    Ejecuta la segmentación de SAM en una lista de URL de imágenes.

    Argumentos:
        image_urls: lista de URL de imágenes o una sola URL
    """
    # Convertir una sola URL en una lista
    if isinstance(image_urls, str):
        image_urls = [image_urls]

    for idx, url in enumerate(image_urls, 1):
        print(f"\n{'=' * 60}")
        print(f"Procesando imagen {idx}/{len(image_urls)}")
        print(f"URL: {url}")

        # Cargar la imagen
        image = load_image(url)
        if image is None:
            continue

        # Convertir la imagen en un arreglo de NumPy para SAM
        image_array = np.array(image)

        # Generar las máscaras
        print("Generando máscaras de segmentación...")
        masks = mask_generator.generate(image_array)

        print(f"Se encontraron {len(masks)} segmentos")

        # Crear la visualización
        fig, axes = plt.subplots(1, 3, figsize=(20, 6))

        # Imagen original
        axes[0].imshow(image)
        axes[0].set_title(
            "Imagen original",
            fontsize=14,
            fontweight="bold"
        )
        axes[0].axis("off")

        # Mostrar únicamente las máscaras de segmentación
        axes[1].imshow(image)
        show_anns(masks, axes[1])
        axes[1].set_title(
            f"Máscaras de segmentación ({len(masks)} segmentos)",
            fontsize=14,
            fontweight="bold"
        )
        axes[1].axis("off")

        # Mostrar las máscaras sobre la imagen original
        overlay_img = create_mask_overlay(image, masks)
        axes[2].imshow(overlay_img)
        axes[2].set_title(
            "Máscaras sobre la imagen original",
            fontsize=14,
            fontweight="bold"
        )
        axes[2].axis("off")

        plt.tight_layout()
        plt.show()
# Aplicar la segmentación a fotografías de iNaturalist
example_urls = [
    "https://inaturalist-open-data.s3.amazonaws.com/photos/639462446/large.jpg",
    "https://inaturalist-open-data.s3.amazonaws.com/photos/639441608/large.jpg"
]

run_segmentation(example_urls)

2. Segmentación mediante indicaciones de texto con CLIPSeg y SAM#

En el ejemplo anterior segmentamos todo lo que aparecía en las imágenes. Esto incluyó el animal principal, como la abeja de la primera imagen y el pez de la segunda, además de elementos del fondo como flores, hojas y rocas. Sin embargo, es posible que solo quieras encontrar una característica específica.

En el siguiente ejemplo exploraremos cómo segmentar una imagen para encontrar una característica concreta mediante una indicación de texto. Utilizaremos el modelo CLIPSeg, desarrollado por Timo Lüddecke y Alexander Ecker. Después de que CLIPSeg identifique la característica solicitada, utilizaremos SAM para refinar la segmentación y obtener un contorno más preciso.

Puedes proporcionar la URL de una imagen y una indicación de texto que describa lo que deseas encontrar. El modelo se encargará del resto.

Nota: CLIPSeg suele funcionar mejor con indicaciones escritas en inglés, por lo que los ejemplos conservan términos como bee, flower, leaf, snow, tree y sky.

# Instalar el paquete necesario
!pip install -q transformers

from transformers import (
    CLIPSegProcessor,
    CLIPSegForImageSegmentation
)
import torch
from matplotlib import pyplot as plt
import numpy as np
from PIL import Image
import requests
from io import BytesIO

# Cargar el modelo CLIPSeg
print(
    "Cargando el modelo CLIPSeg para segmentación mediante "
    "indicaciones de texto..."
)

processor = CLIPSegProcessor.from_pretrained(
    "CIDAS/clipseg-rd64-refined"
)
clipseg_model = CLIPSegForImageSegmentation.from_pretrained(
    "CIDAS/clipseg-rd64-refined"
)

# Utilizar la GPU si está disponible
device = "cuda" if torch.cuda.is_available() else "cpu"
clipseg_model.to(device)

print("El modelo CLIPSeg se cargó correctamente")
print(f"Dispositivo utilizado: {device}")

sam_predictor = SamPredictor(sam)

def segment_with_prompt(
    image_url,
    text_prompts,
    threshold=0.4,
    use_sam_refinement=True
):
    """
    Segmenta objetos de una imagen a partir de indicaciones de texto.

    Argumentos:
        image_url: URL de la imagen
        text_prompts: una indicación o una lista de indicaciones
            (por ejemplo, "cat" o ["cat", "dog", "person"])
        threshold: umbral de confianza para la segmentación
            (entre 0 y 1; un valor mayor es más estricto)
        use_sam_refinement: determina si SAM refinará las máscaras
            (mejora la calidad, pero requiere más tiempo)
    """
    # Convertir una sola indicación en una lista
    if isinstance(text_prompts, str):
        text_prompts = [text_prompts]

    print(f"\n{'=' * 60}")
    print(f"Indicaciones de texto: {text_prompts}")
    print(f"URL de la imagen: {image_url}")
    print("=" * 60)

    # Cargar la imagen
    image = load_image(image_url)
    if image is None:
        return

    # Preparar los datos de entrada para CLIPSeg
    inputs = processor(
        text=text_prompts,
        images=[image] * len(text_prompts),
        padding=True,
        return_tensors="pt"
    ).to(device)

    # Generar las máscaras de segmentación
    print(
        "Generando máscaras de segmentación para "
        f"{len(text_prompts)} indicación(es)..."
    )

    with torch.no_grad():
        outputs = clipseg_model(**inputs)

    # Obtener las predicciones
    preds = outputs.logits

    # Procesar la máscara correspondiente a cada indicación
    image_array = np.array(image)
    masks = []
    refined_masks = []

    for idx, prompt in enumerate(text_prompts):
        # Obtener la máscara de esta indicación
        mask = torch.sigmoid(preds[idx]).cpu().numpy()

        # Ajustar la máscara al tamaño de la imagen original
        mask_resized = np.array(
            Image.fromarray(mask).resize(
                image.size,
                Image.BILINEAR
            )
        )

        # Aplicar el umbral
        binary_mask = mask_resized > threshold
        masks.append((binary_mask, mask_resized, prompt))

        # Refinar la máscara con SAM, cuando corresponda
        if use_sam_refinement and binary_mask.sum() > 0:
            # Encontrar el recuadro delimitador de la máscara
            coords = np.argwhere(binary_mask)

            if len(coords) > 0:
                y_min, x_min = coords.min(axis=0)
                y_max, x_max = coords.max(axis=0)

                # Utilizar SAM para refinar la máscara
                sam_predictor.set_image(image_array)
                box = np.array([x_min, y_min, x_max, y_max])

                refined_mask, _, _ = sam_predictor.predict(
                    box=box,
                    multimask_output=False
                )
                refined_masks.append((refined_mask[0], prompt))
            else:
                refined_masks.append((binary_mask, prompt))
        else:
            if binary_mask.sum() > 0:
                print(f"Se encontró '{prompt}'")
            else:
                print(
                    f"No se detectó '{prompt}'. "
                    "Prueba con un umbral menor."
                )

    # Crear la visualización
    num_plots = 3 if use_sam_refinement else 2
    fig, axes = plt.subplots(
        1,
        num_plots,
        figsize=(7 * num_plots, 6)
    )

    if num_plots == 2:
        axes = [axes[0], axes[1]]

    # Imagen original
    axes[0].imshow(image)
    axes[0].set_title(
        "Imagen original",
        fontsize=14,
        fontweight="bold"
    )
    axes[0].axis("off")

    # Máscaras de CLIPSeg con un mapa de intensidad superpuesto
    overlay = image_array.copy().astype(float)
    combined_mask = np.zeros((*image.size[::-1], 3))

    for binary_mask, mask_resized, prompt in masks:
        if binary_mask.sum() > 0:
            # Generar un color diferente
            color = np.random.randint(50, 255, size=3)

            # Crear una máscara con color
            mask_3d = np.stack([binary_mask] * 3, axis=-1)
            combined_mask += mask_3d * color

            # Agregar el mapa de intensidad
            heatmap = plt.cm.jet(mask_resized)[:, :, :3] * 255
            overlay = (
                overlay * 0.6
                + heatmap
                * 0.4
                * np.expand_dims(binary_mask, -1)
            )

    axes[1].imshow(overlay.astype(np.uint8))
    axes[1].set_title(
        "Segmentación de CLIPSeg",
        fontsize=14,
        fontweight="bold"
    )
    axes[1].axis("off")

    # Máscaras refinadas por SAM, cuando se habilite esta opción
    if use_sam_refinement and len(refined_masks) > 0:
        refined_overlay = image_array.copy()

        for refined_mask, prompt in refined_masks:
            if refined_mask.sum() > 0:
                color = np.random.randint(50, 255, size=3)
                refined_overlay[refined_mask] = (
                    refined_overlay[refined_mask] * 0.4
                    + color * 0.6
                )

        axes[2].imshow(refined_overlay.astype(np.uint8))
        axes[2].set_title(
            "Segmentación refinada por SAM",
            fontsize=14,
            fontweight="bold"
        )
        axes[2].axis("off")

    plt.tight_layout()
    plt.show()

    # Mostrar un resumen
    print("\nResumen de la segmentación:")

    for binary_mask, _, prompt in masks:
        coverage = (
            binary_mask.sum() / binary_mask.size
        ) * 100

        if binary_mask.sum() > 0:
            print(
                f"  • '{prompt}': "
                f"{coverage:.2f} % de la imagen"
            )
        else:
            print(f"  • '{prompt}': no se detectó")

¡Segmentemos la abeja, las flores y las hojas de la imagen de iNaturalist!

example_url = (
    "https://inaturalist-open-data.s3.amazonaws.com/"
    "photos/639462446/large.jpg"
)

# Las indicaciones se conservan en inglés para mejorar los resultados
print("\nEjemplo: segmentación de 'bee' en la imagen")
segment_with_prompt(
    example_url,
    "bee",
    threshold=0.2
)

print("\nEjemplo: segmentación de varios objetos")
segment_with_prompt(
    example_url,
    ["bee", "flower", "leaf"],
    threshold=0.2
)

Este modelo también es útil para comprender los paisajes. Probémoslo con una imagen enviada a la iniciativa de ciencia participativa Land Cover de NASA GLOBE Observer.

example_url = (
    "https://data.globe.gov/system/photos/"
    "2024/12/31/4325697/original.jpg"
)

# Las indicaciones se conservan en inglés para mejorar los resultados
print("\nEjemplo: segmentación de 'snow' en la imagen")
segment_with_prompt(
    example_url,
    "snow",
    threshold=0.2
)

print("\nEjemplo: segmentación de varios objetos")
segment_with_prompt(
    example_url,
    ["snow", "tree", "sky"],
    threshold=0.2
)

Para utilizar este código con tus propias imágenes, modifica el siguiente bloque:

image_url = "enlace a la imagen"

# Escribe las indicaciones en inglés para obtener mejores resultados
prompts = ["indicación 1 en inglés", "indicación 2 en inglés"]

# Aumenta el umbral para excluir segmentaciones con menor confianza
threshold = 0.2

segment_with_prompt(
    image_url,
    prompts,
    threshold=threshold
)