Мультиклассовый UNet с n-мерными спутниковыми изображениямиPython

Программы на Python
Anonymous
Мультиклассовый UNet с n-мерными спутниковыми изображениями

Сообщение Anonymous »

Я пытаюсь использовать UNet в Pytorch для извлечения масок прогнозирования из многомерных (8-канальных) спутниковых изображений. У меня возникли проблемы с тем, чтобы маски прогнозирования выглядели несколько ожидаемо/последовательно. Я не уверен, заключается ли проблема в том, как форматируются мои обучающие данные, в моем обучающем коде или в коде, который я использую для прогнозирования. Я подозреваю, что именно так мои обучающие данные передаются в модель. У меня есть 8-канальные спутниковые изображения и одноканальные маски со значениями в диапазоне от 0 до n, количество классов, где 0 — фон, а 1-n — целевые метки, например:
Изображение

С формой изображения (8, 512, 512) и формой маски это (512, 512) в случае одноканального примера, (512, 512, 8) в случае OHE и (512, 512, 3) в многоканальном случае.
Некоторые маски могут содержать все метки классов, некоторые могут иметь только пару или быть только фоновыми метками. Я пробовал использовать эти одноканальные маски, я также преобразовывал их в трехканальные маски, причем первый канал представлял собой все метки для данного изображения, а также я пробовал их горячее кодирование так, чтобы каждая маска была 0- n измерений и для каждого канала своя метка с двоичными значениями 0-1 для фона/цели.
РЕДАКТИРОВАНИЕ
После изменения softmax dim=2, результаты стали выглядеть немного лучше. Тем не менее, похоже, что модель вообще не обучается после первых нескольких эпох прогрева, поскольку потери при обучении сначала уменьшаются, но затем сразу же выходят на плато или увеличиваются, и маски прогнозирования перестают иметь смысл (либо все черные, либо случайные пятна). Я подозреваю, что возникла проблема с моим конвейером обучения (ниже) или, возможно, из-за дисбаланса классов с классом 0 (фон).
import os
import torch
import numpy as np
from skimage import io
from tqdm import tqdm
import torch.nn as nn
import torch.optim as optim
import segmentation_models_pytorch as smp

image_dir = r'test_segmentation\images'
mask_dir = r'test_segmentation\masks'

data_dir=r'unet_training'
os.makedirs(data_dir, exist_ok=True)

model_dir = os.path.join(data_dir, 'models')
os.makedirs(model_dir, exist_ok=True)

pred_dir = os.path.join(data_dir, 'predictions')
os.makedirs(pred_dir, exist_ok=True)

num_bands = 8
num_classes = 9
epochs = 10
learning_rate = 0.001
weight_decay = 0
encoder = 'resnet50'
encoder_weights = 'imagenet'

model = smp.Unet(in_channels=num_bands, encoder_name=encoder, encoder_weights=encoder_weights, classes=num_classes).to(device)
optimizer = optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)
loss_function = nn.CrossEntropyLoss() if num_classes > 1 else nn.BCEWithLogitsLoss()

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

for epoch in range(1, epochs + 1):
train_loss = 0
val_loss = 0

train_loop = tqdm(enumerate(train_loader), total=len(train_loader), desc=f"Epoch {epoch} Training")

model.train()

for batch_idx, (data, targets) in train_loop:
optimizer.zero_grad()

data = data.float().to(device)
targets = targets.long().to(device)
predictions = model(data)
loss = loss_function(predictions, targets)

train_loss += loss.item()

loss.backward()
optimizer.step()

train_loop.set_postfix(loss=train_loss)

val_loop = tqdm(enumerate(val_loader), total=len(val_loader), desc=f"Epoch {epoch} Validation")

model.eval()

for batch_idx, (data, targets) in val_loop:
data, targets = data.to(device).float(), targets.to(device).long()
preds = model(data)

val_loss = loss_function(preds, targets).item()

softmax = torch.nn.Softmax(dim=2)
preds = torch.argmax(softmax(preds), dim=1).cpu().numpy()
preds = np.array(preds[0, :, :], dtype=np.uint8)
labels = np.array(targets.cpu().numpy()[0, :, :], dtype=np.uint8)

#save prediction and label mask
pred_path = os.path.join(pred_dir, f"{epoch}_{batch_idx}_pred.png")
label_path = os.path.join(pred_dir, f"{epoch}_{batch_idx}_label.png")
io.imsave(pred_path, preds)
io.imsave(label_path, labels)

val_loop.set_postfix(loss=val_loss)

avg_train_loss = train_loss / (batch_idx + 1)
avg_val_loss = val_loss/ (batch_idx + 1)

print(f"\nEpoch {epoch} Train Loss: {avg_train_loss}, Val Loss: {avg_val_loss}")

checkpoint_name = os.path.join(model_dir, f"{modeltype}_bands{num_bands}_classes{num_classes}_{encoder}_{learning_rate}_{epoch}.pt")

if epoch == 1:
torch.save(model.state_dict(), checkpoint_name)
elif epoch % 10 == 0:
torch.save(model.state_dict(), checkpoint_name)
elif epoch == epochs:
torch.save(model.state_dict(), checkpoint_name)
else:
pass


Подробнее здесь: https://stackoverflow.com/questions/784 ... ite-images

Вернуться в «Python»