基于PyTorch的多波段卫星图像多分类UNet训练问题排查
8波段卫星图像UNet分割掩码无意义、训练停滞问题排查
我使用PyTorch中的UNet模型从8波段卫星图像中提取预测掩码,但生成的掩码不符合预期、缺乏连贯性。不确定问题出在训练数据格式、训练代码还是预测代码上,怀疑是训练数据输入模型的方式有问题。
我的数据情况:
- 8波段卫星图像形状为
(8, 512, 512) - 掩码有三种形式:
- 单通道掩码:形状
(512, 512),值范围0到n(0代表背景,1到n为目标类别) - 独热编码(OHE)掩码:形状
(512, 512, 8) - 堆叠掩码:形状
(512, 512, 3)
- 单通道掩码:形状
- 部分掩码包含所有类别,部分仅含少数类别或仅背景,我尝试过使用这三种掩码进行训练。
编辑补充:将softmax的dim=2修改后,输出有所改善,但模型在最初几个热身epoch后完全停止学习:训练损失初期下降后立即停滞或上升,预测掩码变得无意义(全黑或随机斑块)。怀疑问题出在训练流程(代码如下)或类别不平衡(背景类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
内容的提问来源于stack exchange,提问作者andrewr
相关产品推荐
相关产品推荐

