本文共 7493 字,大约阅读时间需要 24 分钟。
PyTorch 是一个强大的深度学习框架,广泛应用于图像处理和计算机视觉任务中。对于图像分割任务,PyTorch 提供了丰富的工具和库,使得开发和训练模型变得更加高效和易于使用。以下是基于 PyTorch 的图像分割任务的详细介绍,结合实际开发经验总结出的最佳实践。
图像分割任务的核心目标是对图像中的每个像素进行分类,确定其所属的类别。与图像分类任务不同,图像分割任务需要在像素级别上进行分类,输出的结果是一个大小为图像尺寸的分辨率一致的类别矩阵。
在 PyTorch 中,可以选择预训练模型或自定义模型进行图像分割任务。以下是几种常用的模型架构:
根据任务需求和数据集大小,选择合适的模型架构。对于本文,选择 U-Net 作为基础模型。
U-Net 由编码器和解码器组成,具体结构如下:
import torchimport torch.nn as nnimport torch.nn.functional as Fclass UNet(nn.Module): def __init__(self, in_channels=3, num_classes=2, init_features=32): super(UNet, self).__init__() features = init_features self.encoder1 = UNet._block(in_channels, features, name="enc1") self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2) self.encoder2 = UNet._block(features, features * 2, name="enc2") self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2) self.encoder3 = UNet._block(features * 2, features * 4, name="enc3") self.pool3 = nn.MaxPool2d(kernel_size=2, stride=2) self.bottleneck = UNet._block(features * 4, features * 8, name="bottleneck") self.upconv4 = nn.ConvTranspose2d(features * 8, features * 4, kernel_size=2, stride=2) self.decoder4 = UNet._block((features * 4) * 2, features * 4, name="dec4") self.upconv3 = nn.ConvTranspose2d(features * 4, features * 2, kernel_size=2, stride=2) self.decoder3 = UNet._block((features * 2) * 2, features * 2, name="dec3") self.upconv2 = nn.ConvTranspose2d(features * 2, features, kernel_size=2, stride=2) self.decoder2 = UNet._block(features * 2, features, name="dec2") self.upconv1 = nn.ConvTranspose2d(features, features, kernel_size=2, stride=2) self.decoder1 = UNet._block(features * 2, features, name="dec1") self.conv = nn.Conv2d(in_channels=features, out_channels=num_classes, kernel_size=1) def forward(self, x): enc1 = self.encoder1(x) enc2 = self.encoder2(self.pool1(enc1)) enc3 = self.encoder3(self.pool2(enc2)) enc4 = self.encoder4(self.pool3(enc3)) bottleneck = self.bottleneck(self.pool4(enc4)) dec4 = self.upconv4(bottleneck) dec4 = torch.cat((dec4, enc4), dim=1) dec4 = self.decoder4(dec4) dec3 = self.upconv3(dec4) dec3 = torch.cat((dec3, enc3), dim=1) dec3 = self.decoder3(dec3) dec2 = self.upconv2(dec3) dec2 = torch.cat((dec2, enc2), dim=1) dec2 = self.decoder2(dec2) dec1 = self.upconv1(dec2) dec1 = torch.cat((dec1, enc1), dim=1) dec1 = self.decoder1(dec1) return torch.sigmoid(self.conv(dec1)) @staticmethod def _block(in_channels, features, name): return nn.Sequential( OrderedDict( [ (name + "conv1", nn.Conv2d(in_channels=in_channels, out_channels=features, kernel_size=3, padding=1, bias=False)), (name + "norm1", nn.BatchNorm2d(num_features=features)), (name + "relu1", nn.ReLU(inplace=True)), (name + "conv2", nn.Conv2d(in_channels=features, out_channels=features, kernel_size=3, padding=1, bias=False)), (name + "norm2", nn.BatchNorm2d(num_features=features)), (name + "relu2", nn.ReLU(inplace=True)), ] ) )
import osimport cv2from PIL import Imagefrom torch.utils.data import Dataset, DataLoaderclass ImageSegmentationDataset(Dataset): def __init__(self, img_dir, label_dir, transform=None): self.img_dir = img_dir self.label_dir = label_dir self.transform = transform self.images = [] self.labels = [] for img_name in os.listdir(img_dir): img_path = os.path.join(img_dir, img_name) label_path = os.path.join(label_dir, img_name.replace('.png', '.png')) image = Image.open(img_path).convert('RGB') label = cv2.imread(label_path) self.images.append(image) self.labels.append(label) def __getitem__(self, index): img = self.images[index] label = self.labels[index] if self.transform is not None: img = self.transform(img) return img, label def __len__(self): return len(self.images) import torchimport torch.optim as optimimport torch.nn.functional as Fdef train_model(model, train_loader, val_loader, num_epochs=10, device='cuda'): model.to(device) optimizer = optim.Adam(model.parameters(), lr=1e-4) criterion = nn.CrossEntropyLoss() best_val_loss = float('inf') for epoch in range(num_epochs): model.train() epoch_loss = 0.0 epoch_acc = 0.0 for inputs, labels in train_loader: inputs = inputs.to(device) labels = labels.to(device) outputs = model(inputs) loss = criterion(outputs, labels) preds = torch.argmax(outputs, dim=1) total = labels.size(0) correct = (preds == labels).sum().item() epoch_loss += loss.item() epoch_acc += correct / total avg_loss = epoch_loss / len(train_loader) avg_acc = epoch_acc / len(train_loader.dataset) print(f"Epoch {epoch + 1}/{num_epochs}") print(f"Train Loss: {avg_loss:.4f}") print(f"Train Acc: {avg_acc:.3f}") model.eval() val_loss = 0.0 val_acc = 0.0 with torch.no_grad(): for inputs, labels in val_loader: inputs = inputs.to(device) labels = labels.to(device) outputs = model(inputs) loss = criterion(outputs, labels) preds = torch.argmax(outputs, dim=1) total = labels.size(0) correct = (preds == labels).sum().item() val_loss += loss.item() val_acc += correct / total avg_val_loss = val_loss / len(val_loader) avg_val_acc = val_acc / len(val_loader.dataset) print(f"Val Loss: {avg_val_loss:.4f}") print(f"Val Acc: {avg_val_acc:.3f}") if avg_val_loss < best_val_loss: best_val_loss = avg_val_loss torch.save(model.state_dict(), 'best_model.pth') return model, best_val_loss def evaluate_model(model, test_loader, device='cuda'): model.eval() acc = 0.0 total = 0 with torch.no_grad(): for inputs, labels in test_loader: inputs = inputs.to(device) labels = labels.to(device) outputs = model(inputs) preds = torch.argmax(outputs, dim=1) total += labels.size(0) acc += (preds == labels).sum().item() avg_acc = acc / total if total > 0 else 0.0 return avg_acc
通过以上步骤,可以在 PyTorch 中成功实现图像分割任务。以下是一些优化建议:
通过以上方法,可以在 PyTorch 中高效地完成图像分割任务,并根据需要对模型进行不断优化和调整。
转载地址:http://weofk.baihongyu.com/