Aprenda a ajustar transformadores y segmentar cualquier cosa |  de Stefan Todoran |  junio de 2024

Entrene el modelo Segment Anything (SAM) de Meta para segmentar máscaras de alta fidelidad para cualquier dominio

11 minutos de lectura

hace 11 horas

El lanzamiento de varios modelos básicos potentes y de código abierto, junto con los avances en el ajuste fino, han dado lugar a un nuevo paradigma en el aprendizaje automático y la inteligencia artificial. En el centro de esta revolución se encuentra el modelo de transformador.

Si bien los modelos específicos de dominio de alta precisión alguna vez estuvieron fuera del alcance de todos, excepto de las corporaciones mejor financiadas, hoy el paradigma del modelo fundamental permite que incluso los modestos recursos disponibles para estudiantes o investigadores independientes logren resultados que rivalizan con los modelos propietarios de última generación.

El ajuste fino puede mejorar enormemente el rendimiento en tareas fuera de distribución (fuente de la imagen: por el autor).

Este artículo explora la aplicación del Segment Anything Model (SAM) de Meta a la tarea de teledetección de segmentación de píxeles de ríos. Si desea acceder directamente al código, el archivo fuente de este proyecto está disponible en GitHub y los datos están en Cara abrazadaaunque se recomienda leer primero el artículo completo.

El primer paso es encontrar o crear un conjunto de datos adecuado. Según la literatura existente, un buen conjunto de datos de ajuste fino para SAM tendrá al menos entre 200 y 800 imágenes. Una lección clave de la última década de avances en el aprendizaje profundo es que más datos siempre es mejor, por lo que no puedes equivocarte con un conjunto de datos de ajuste más grande. Sin embargo, el objetivo detrás de los modelos fundamentales es permitir que incluso conjuntos de datos relativamente pequeños sean suficientes para un rendimiento sólido.

También será necesario tener una cuenta HuggingFace, que puede ser creado aquíUsando HuggingFace podemos almacenar y recuperar fácilmente nuestro conjunto de datos en cualquier momento y desde cualquier dispositivo, lo que facilita la colaboración y la reproducibilidad.

El último requisito es un dispositivo con GPU en el que podamos ejecutar el flujo de trabajo de entrenamiento. Una GPU Nvidia T4, que está disponible de forma gratuita a través de colaboración de googlees lo suficientemente potente como para entrenar el punto de control del modelo SAM más grande (sam-vit-huge) en 1000 imágenes durante 50 épocas en menos de 12 horas.

Para evitar perder el progreso de los límites de uso en tiempos de ejecución alojados, puede montar Google Drive y guardar allí cada punto de control del modelo. Alternativamente, implemente y conéctese a un Máquina virtual de GCP para eludir los límites por completo. Si nunca has usado GCP antes, eres elegible para un crédito gratuito de $300 dólares, que es suficiente para entrenar el modelo al menos una docena de veces.

Antes de comenzar a entrenar, debemos comprender la arquitectura de SAM. El modelo contiene tres componentes: un codificador de imágenes de un modelo mínimamente modificado codificador automático enmascarado, un codificador de mensajes flexible capaz de procesar diversos tipos de mensajes y un decodificador de máscaras rápido y liviano. Una motivación detrás del diseño es permitir una segmentación rápida y en tiempo real en dispositivos periféricos (por ejemplo, en el navegador), ya que la incrustación de la imagen solo necesita calcularse una vez y el decodificador de máscara puede ejecutarse en aproximadamente 50 ms en la CPU.

La arquitectura del modelo de SAM nos muestra qué entradas acepta el modelo y qué partes del modelo deben entrenarse (fuente de la imagen: SAMGitHub).

En teoría, el codificador de imágenes ya ha aprendido la forma óptima de incrustar una imagen, identificando formas, bordes y otras características visuales generales. De manera similar, en teoría, el codificador de mensajes ya puede codificar mensajes de manera óptima. El decodificador de máscara es la parte de la arquitectura del modelo que toma estas incrustaciones de imágenes y mensajes y realmente crea la máscara operando sobre la imagen y las incrustaciones de mensajes.

Como tal, un enfoque es congelar los parámetros del modelo asociados con la imagen y solicitar a los codificadores durante el entrenamiento y actualizar solo los pesos del decodificador de máscara. Este enfoque tiene la ventaja de permitir tareas posteriores supervisadas y no supervisadas, ya que los puntos de control y los cuadros delimitadores son automatizables y utilizables por humanos.

Diagrama que muestra el codificador de imagen SAM congelado y el decodificador de máscara, junto con el codificador de aviso sobrecargado, utilizado en la arquitectura AutoSAM (fuente: papel AutoSAM).

Un enfoque alternativo es sobrecargar el codificador de mensajes, congelar el codificador de imágenes y el decodificador de máscara y simplemente no utilizar el codificador de máscara SAM original. Por ejemplo, la arquitectura AutoSAM utiliza una red basada en Harmonic Dense Net para producir incrustaciones rápidas basadas en la propia imagen. En este tutorial cubriremos el primer enfoque, congelar la imagen y solicitar codificadores y entrenar solo el decodificador de máscara, pero el código para este enfoque alternativo se puede encontrar en AutoSAM. GitHub y papel.

El siguiente paso es determinar qué tipo de indicaciones recibirá el modelo durante el tiempo de inferencia, de modo que podamos proporcionar ese tipo de indicaciones en el momento del entrenamiento. Personalmente, no recomendaría el uso de indicaciones de texto para ningún proceso serio de visión por computadora, dada la naturaleza impredecible/inconsistente del procesamiento del lenguaje natural. Esto deja puntos y cuadros delimitadores, y la elección depende en última instancia de la naturaleza particular de su conjunto de datos específico, aunque la literatura ha encontrado que los cuadros delimitadores superan a los puntos de control de manera bastante consistente.

Las razones de esto no están del todo claras, pero podría deberse a cualquiera de los siguientes factores o a una combinación de ellos:

  • Los buenos puntos de control son más difíciles de seleccionar en el momento de la inferencia (cuando se desconoce la máscara de verdad fundamental) que los cuadros delimitadores.
  • El espacio de posibles indicaciones puntuales es órdenes de magnitud mayor que el espacio de posibles indicaciones del cuadro delimitador, por lo que no ha sido entrenado tan exhaustivamente.
  • Los autores originales de SAM se centraron en las capacidades de disparo cero y de pocos disparos (contadas en términos de interacciones humanas inmediatas) del modelo, por lo que el entrenamiento previo puede haberse centrado más en los cuadros delimitadores.

De todos modos, la segmentación de ríos es en realidad un caso raro en el que las indicaciones puntuales superan a los cuadros delimitadores (aunque sólo ligeramente, incluso con un dominio extremadamente favorable). Dado que en cualquier imagen de un río la masa de agua se extenderá desde un extremo de la imagen hasta el otro, cualquier cuadro delimitador casi siempre cubrirá la mayor parte de la imagen. Por lo tanto, las indicaciones del cuadro delimitador para porciones muy diferentes del río pueden parecer extremadamente similares, lo que en teoría significa que los cuadros delimitadores proporcionan al modelo significativamente menos información que los puntos de control y, por lo tanto, conducen a un peor rendimiento.

Puntos de control, indicaciones del cuadro delimitador y segmentación de la verdad fundamental superpuestos en dos imágenes de entrenamiento de muestra (fuente de la imagen: por el autor).

Observe cómo en la ilustración anterior, aunque las verdaderas máscaras de segmentación para las dos porciones del río son completamente diferentes, sus respectivos cuadros delimitadores son casi idénticos, mientras que sus indicaciones de puntos difieren (comparativamente) más.

El otro factor importante a considerar es la facilidad con la que se pueden generar indicaciones de entrada en el momento de la inferencia. Si espera tener un humano en el bucle, entonces tanto los cuadros delimitadores como los puntos de control son bastante triviales de adquirir en el momento de la inferencia. Sin embargo, en el caso de que desee tener un proceso completamente automatizado, responder estas preguntas se vuelve más complicado.

Ya sea que se utilicen puntos de control o cuadros delimitadores, la generación del mensaje generalmente implica primero estimar una máscara aproximada para el objeto de interés. Los cuadros delimitadores pueden ser simplemente el cuadro mínimo que envuelve la máscara aproximada, mientras que los puntos de control deben tomarse como muestra de la máscara aproximada. Esto significa que los cuadros delimitadores son más fáciles de obtener cuando se desconoce la máscara de verdad básica, ya que la máscara estimada para el objeto de interés solo necesita coincidir aproximadamente con el mismo tamaño y posición del objeto verdadero, mientras que para los puntos de control la máscara estimada necesitaría coincidir más estrechamente con los contornos del objeto.

Al utilizar una máscara estimada en lugar de la verdad fundamental, la ubicación de los puntos de control puede incluir puntos mal etiquetados, mientras que los cuadros delimitadores generalmente están en el lugar correcto (fuente de la imagen: por el autor).

Para la segmentación de ríos, si tenemos acceso tanto a RGB como a NIR, podemos utilizar métodos de umbralización de índices espectrales para obtener nuestra máscara aproximada. Si solo tenemos acceso a RGB, podemos convertir la imagen a HSV y aplicar un umbral a todos los píxeles dentro de un cierto rango de tono, saturación y valor. Luego, podemos eliminar los componentes conectados por debajo de un cierto umbral de tamaño y utilizar erosion de skimage.morphology para asegurarnos de que los únicos píxels en nuestra máscara sean aquellos que estaban hacia el centro de las grandes manchas azules.

Para entrenar nuestro modelo, necesitamos un cargador de datos que contenga todos nuestros datos de entrenamiento y que podamos iterar para cada época de entrenamiento. Cuando cargamos nuestro conjunto de datos desde HuggingFace, toma la forma de un datasets.Dataset clase. Si el conjunto de datos es privado, asegúrese de instalar primero la CLI de HuggingFace e iniciar sesión con !huggingface-cli login.

from datasets import load_dataset, load_from_disk, Dataset

hf_dataset_name = "stodoran/elwha-segmentation-v1"
training_data = load_dataset(hf_dataset_name, split="train")
validation_data = load_dataset(hf_dataset_name, split="validation")

Luego necesitamos codificar nuestra propia clase de conjunto de datos personalizado que devuelva no solo una imagen y una etiqueta para cualquier índice, sino también el mensaje. A continuación se muestra una implementación que puede manejar indicaciones tanto de puntos de control como de cuadros delimitadores. Para ser inicializado, se necesita un HuggingFace datasets.Dataset instancia y una instancia de procesador SAM.

from torch.utils.data import Dataset

class PromptType:
CONTROL_POINTS = "pts"
BOUNDING_BOX = "bbox"

class SAMDataset(Dataset):
def __init__(
self,
dataset,
processor,
prompt_type = PromptType.CONTROL_POINTS,
num_positive = 3,
num_negative = 0,
erode = True,
multi_mask = "mean",
perturbation = 10,
image_size = (1024, 1024),
mask_size = (256, 256),
):
# Asign all values to self
...

def __len__(self):
return len(self.dataset)

def __getitem__(self, idx):
datapoint = self.dataset[idx]
input_image = cv2.resize(np.array(datapoint["image"]), self.image_size)
ground_truth_mask = cv2.resize(np.array(datapoint["label"]), self.mask_size)

if self.prompt_type == PromptType.CONTROL_POINTS:
inputs = self._getitem_ctrlpts(input_image, ground_truth_mask)
elif self.prompt_type == PromptType.BOUNDING_BOX:
inputs = self._getitem_bbox(input_image, ground_truth_mask)

inputs["ground_truth_mask"] = ground_truth_mask
return inputs

También tenemos que definir el SAMDataset._getitem_ctrlpts y SAMDataset._getitem_bbox funciones, aunque si solo planea usar un tipo de mensaje, puede refactorizar el código para manejar directamente ese tipo en SAMDataset.__getitem__ y eliminar la función auxiliar.

class SAMDataset(Dataset):
...

def _getitem_ctrlpts(self, input_image, ground_truth_mask):
# Get control points prompt. See the GitHub for the source
# of this function, or replace with your own point selection algorithm.
input_points, input_labels = generate_input_points(
num_positive=self.num_positive,
num_negative=self.num_negative,
mask=ground_truth_mask,
dynamic_distance=True,
erode=self.erode,
)
input_points = input_points.astype(float).tolist()
input_labels = input_labels.tolist()
input_labels = [[x] for x in input_labels]

# Prepare the image and prompt for the model.
inputs = self.processor(
input_image,
input_points=input_points,
input_labels=input_labels,
return_tensors="pt"
)

# Remove batch dimension which the processor adds by default.
inputs = {k: v.squeeze(0) for k, v in inputs.items()}
inputs["input_labels"] = inputs["input_labels"].squeeze(1)

return inputs

def _getitem_bbox(self, input_image, ground_truth_mask):
# Get bounding box prompt.
bbox = get_input_bbox(ground_truth_mask, perturbation=self.perturbation)

# Prepare the image and prompt for the model.
inputs = self.processor(input_image, input_boxes=[[bbox]], return_tensors="pt")
inputs = {k: v.squeeze(0) for k, v in inputs.items()} # Remove batch dimension which the processor adds by default.

return inputs

Juntando todo esto, podemos crear una función que crea y devuelve un cargador de datos de PyTorch dada una de las divisiones del conjunto de datos de HuggingFace. Escribir funciones que devuelvan cargadores de datos en lugar de simplemente ejecutar celdas con el mismo código no solo es una buena práctica para escribir código flexible y fácil de mantener, sino que también es necesario si planea usar HuggingFace Acelerar para ejecutar entrenamiento distribuido.

from transformers import SamProcessor
from torch.utils.data import DataLoader

def get_dataloader(
hf_dataset,
model_size = "base", # One of "base", "large", or "huge"
batch_size = 8,
prompt_type = PromptType.CONTROL_POINTS,
num_positive = 3,
num_negative = 0,
erode = True,
multi_mask = "mean",
perturbation = 10,
image_size = (256, 256),
mask_size = (256, 256),
):
processor = SamProcessor.from_pretrained(f"facebook/sam-vit-{model_size}")

sam_dataset = SAMDataset(
dataset=hf_dataset,
processor=processor,
prompt_type=prompt_type,
num_positive=num_positive,
num_negative=num_negative,
erode=erode,
multi_mask=multi_mask,
perturbation=perturbation,
image_size=image_size,
mask_size=mask_size,
)
dataloader = DataLoader(sam_dataset, batch_size=batch_size, shuffle=True)

return dataloader

Después de esto, el entrenamiento es simplemente una cuestión de cargar el modelo, congelar la imagen y los codificadores de indicaciones, y entrenar para la cantidad deseada de iteraciones.

model = SamModel.from_pretrained(f"facebook/sam-vit-{model_size}")
optimizer = AdamW(model.mask_decoder.parameters(), lr=learning_rate, weight_decay=weight_decay)

# Train only the decoder.
for name, param in model.named_parameters():
if name.startswith("vision_encoder") or name.startswith("prompt_encoder"):
param.requires_grad_(False)

A continuación se muestra el esquema básico del código del bucle de entrenamiento. Tenga en cuenta que forward_pass, calculate loss, evaluate_modely save_model_checkpoint Las funciones se han omitido por razones de brevedad, pero las implementaciones están disponibles en GitHub. El código de paso hacia adelante diferirá ligeramente según el tipo de solicitud, y el cálculo de la pérdida también necesita un caso especial según el tipo de solicitud; cuando se utilizan solicitudes de puntos, SAM devuelve una máscara predicha para cada punto de entrada, por lo que para obtener una única máscara que se pueda comparar con la verdad fundamental, se deben promediar las máscaras predichas o se debe seleccionar la mejor máscara predicha (identificada según las puntuaciones de IoU predichas de SAM).

train_losses = []
validation_losses = []
epoch_loop = tqdm(total=num_epochs, position=epoch, leave=False)
batch_loop = tqdm(total=len(train_dataloader), position=0, leave=True)

while epoch < num_epochs:
epoch_losses = []

batch_loop.n = 0 # Loop Reset
for idx, batch in enumerate(train_dataloader):
# Forward Pass
batch = {k: v.to(accelerator.device) for k, v in batch.items()}
outputs = forward_pass(model, batch, prompt_type)

# Compute Loss
ground_truth_masks = batch["ground_truth_mask"].float()
train_loss = calculate_loss(outputs, ground_truth_masks, prompt_type, loss_fn, multi_mask="best")
epoch_losses.append(train_loss)

# Backward Pass & Optimizer Step
optimizer.zero_grad()
accelerator.backward(train_loss)
optimizer.step()
lr_scheduler.step()

batch_loop.set_description(f"Train Loss: {train_loss.item():.4f}")
batch_loop.update(1)

validation_loss = evaluate_model(model, validation_dataloader, accelerator.device, loss_fn)
train_losses.append(torch.mean(torch.Tensor(epoch_losses)))
validation_losses.append(validation_loss)

if validation_loss < best_loss:
save_model_checkpoint(
accelerator,
best_checkpoint_path,
model,
optimizer,
lr_scheduler,
epoch,
train_history,
validation_loss,
train_losses,
validation_losses,
loss_config,
model_descriptor=model_descriptor,
)
best_loss = validation_loss

epoch_loop.set_description(f"Best Loss: {best_loss:.4f}")
epoch_loop.update(1)
epoch += 1

Para el proyecto del río Elwha, la mejor configuración lograda entrenó el modelo “sam-vit-base” utilizando un conjunto de datos de más de 1.000 máscaras de segmentación utilizando una instancia de GCP en menos de 12 horas.

En comparación con el SAM de referencia, el ajuste fino mejoró drásticamente el rendimiento, y la máscara mediana pasó de inutilizable a muy precisa.

El ajuste fino de SAM mejora en gran medida el rendimiento de la segmentación en relación con el SAM base con el mensaje predeterminado (fuente de la imagen: por el autor).

Un hecho importante a tener en cuenta es que el conjunto de datos de entrenamiento de imágenes de ríos de 1k era imperfecto, con etiquetas de segmentación que variaban mucho en la cantidad de píxeles clasificados correctamente. Como tal, las métricas que se muestran arriba se calcularon en un conjunto de datos de píxeles perfectos de 225 imágenes de ríos.

Un comportamiento interesante observado fue que el modelo aprendió a generalizar a partir de los datos de entrenamiento imperfectos. Al evaluar los puntos de datos donde el ejemplo de entrenamiento contenía errores de clasificación obvios, podemos observar que la predicción del modelo evita el error. Observe cómo las imágenes en la fila superior que muestran muestras de entrenamiento contienen máscaras que no llenan el río hasta la orilla, mientras que la fila inferior que muestra las predicciones del modelo segmenta los límites del río de manera más precisa.