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
)