相同配置下TensorFlow与PyTorch实现Unet分割模型的性能差异求助
相同配置下TensorFlow与PyTorch实现Unet分割模型的性能差异求助
大家好,这是之前那个“TensorFlow和PyTorch性能差异排查”问题的后续。之前的结论是TF的model.fit默认采用mini-batch(batch size=32)而PyTorch是普通batching导致性能差距,但我遇到了即使batch size完全相同,两者训练出来的模型性能依然有明显差异的情况,想请大家帮忙分析下可能的原因?
我分别用了两个主流的分割模型库来实现Unet:TF端用的是segmentation-models(简称sm),PyTorch端用的是segmentation_models_pytorch(简称smp),我已经尽可能对齐了模型架构和所有训练参数,具体配置如下:
- batch_size: 2
- 训练轮数(Epoch): 50
- 优化器: Adam(学习率lr=0.001)
- 损失函数: dice_loss
- 评估指标: accuracy、auc、iou(交并比)
遗憾的是我没法提供训练数据供大家复现。
TensorFlow 实现
import tensorflow as tf from tensorflow import keras import tensorflow.keras.backend as K import os # iou 指标计算函数 def iou(y_true, y_pred, threshold=0.5): y_pred = tf.cast(y_pred > threshold, tf.float32) y_true = tf.cast(y_true, tf.float32) intersection = tf.reduce_sum(y_true * y_pred, axis=[1, 2, 3]) union = tf.reduce_sum(y_true, axis=[1, 2, 3]) + tf.reduce_sum(y_pred, axis=[1, 2, 3]) - intersection iou = tf.reduce_mean(intersection / (union + tf.keras.backend.epsilon())) return iou # dice_loss 损失函数 def dice_loss(y_true, y_pred): y_true_f = tf.cast(K.flatten(y_true), tf.float32) y_pred_f = tf.cast(K.flatten(y_pred), tf.float32) intersection = K.sum(y_true_f*y_pred_f) val = (2. * intersection + K.epsilon()) / (K.sum(y_true_f * y_true_f) + K.sum(y_pred_f * y_pred_f) + K.epsilon()) return 1. - val # 数据分批 train_dataset = train_dataset.batch(config.BATCH_SIZE) test_dataset = test_dataset.batch(config.BATCH_SIZE) # 构建Unet模型(适配4通道输入) pretrained_base_model = sm.Unet(encoder_weights='imagenet', classes=1) # 4通道转3通道(适配预训练模型输入) inp = keras.Input(shape=(None, None, 4)) l1 = keras.layers.Conv2D(3, (1, 1))(inp) # 将4通道映射为3通道 out = pretrained_base_model(l1) model = keras.Model(inp, out, name=pretrained_base_model.name) model.compile( optimizer='Adam', loss=dice_loss, # 自定义dice_loss metrics=[ iou, # 自定义IoU指标 keras.metrics.AUC(), keras.metrics.BinaryAccuracy() ] ) # 回调函数定义 backup_callback = keras.callbacks.BackupAndRestore( backup_dir="./keras_backups" ) checkpoint_callback = keras.callbacks.ModelCheckpoint( filepath= os.path.join(config.SAVE_MODEL_DIR, config.VERSION, "best_mdl.keras"), monitor='val_iou', save_best_only=True, mode='max', verbose=0 ) # 开始训练 callbacks = [backup_callback, checkpoint_callback] artery_history = model.fit(train_dataset, epochs=config.EPOCHS, callbacks=callbacks, validation_data=test_dataset)
PyTorch 实现
import torch import torch.nn as nn import numpy as np from tqdm import tqdm from sklearn.metrics import roc_auc_score import os from torch.utils.data import DataLoader # 4通道转3通道的卷积层 channel_mapper = nn.Conv2d(in_channels=4, out_channels=3, kernel_size=1) # 包装Unet模型 class WrappedModel(nn.Module): def __init__(self, base_model, channel_mapper): super(WrappedModel, self).__init__() self.channel_mapper = channel_mapper self.base_model = base_model def forward(self, x): x = self.channel_mapper(x) return self.base_model(x) # 初始化模型并移至设备 model = WrappedModel( base_model=smp.Unet(encoder_name="vgg16", encoder_weights="imagenet", in_channels=3, classes=1), channel_mapper=channel_mapper, ).to(device) # 数据加载器 train_loader = DataLoader(train_dataset,batch_size=config.BATCH_SIZE, shuffle=True, num_workers=4, pin_memory=True) test_loader = DataLoader(test_dataset, batch_size=config.BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True) # iou 指标计算函数 def iou(y_true, y_pred, threshold=0.5): if isinstance(y_true, np.ndarray): y_true = torch.tensor(y_true) if isinstance(y_pred, np.ndarray): y_pred = torch.tensor(y_pred) y_pred = (y_pred > threshold).float() y_true = y_true.float() intersection = torch.sum(y_true * y_pred) union = torch.sum(y_true) + torch.sum(y_pred) - intersection # IoU计算,加小epsilon防止除零 iou = intersection / (union + 1e-6) return iou.item() # 返回标量IoU值 # dice_loss 损失函数 def dice_loss(y_true, y_pred): y_true_f = y_true.view(-1).float() y_pred_f = y_pred.view(-1).float() # 计算交集 intersection = torch.sum(y_true_f * y_pred_f) # Dice系数计算 dice_score = (2. * intersection + 1e-6) / (torch.sum(y_true_f * y_true_f) + torch.sum(y_pred_f * y_pred_f) + 1e-6) # Dice损失为1减去Dice系数 return 1. - dice_score # 初始化优化器和损失函数 optimizer = torch.optim.Adam(model.parameters(), lr=0.001) loss_fn = dice_loss # 训练模型 history = train_model( model=model, train_loader=train_loader, val_loader=test_loader, optimizer=optimizer, loss_fn=loss_fn, num_epochs=config.EPOCHS, save_dir= os.path.join(config.SAVE_MODEL_DIR,config.VERSION), checkpoint_path="./checkpoints/checkpoint.pth", model_name= "unet.pth", device=device, ) # 自定义训练循环 def train_model( model, train_loader, val_loader, optimizer, loss_fn, num_epochs, save_dir, device, checkpoint_path, model_name , metrics=[ "iou", "accuracy", "auc"], ): # 初始化历史记录(如果没有加载到checkpoint) def init_history(): return { "train_loss": [], "val_loss": [], "train_metrics": {m: [] for m in metrics}, "val_metrics": {m: [] for m in metrics} } # 简化的checkpoint加载/保存逻辑(示例) def load_checkpoint(model, optimizer, path): if os.path.exists(path): checkpoint = torch.load(path) model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) start_epoch = checkpoint['epoch'] history = checkpoint['history'] return model, optimizer, start_epoch, history return model, optimizer, 0, init_history() def save_checkpoint(model, optimizer, epoch, history, path): torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'history': history }, path) os.makedirs("./checkpoints", exist_ok=True) best_model_path = os.path.join(save_dir, model_name) # 加载 checkpoint(如果存在) model, optimizer, start_epoch, history = load_checkpoint(model, optimizer, checkpoint_path) sigmoid = torch.nn.Sigmoid() best_val_iou = max(history["val_metrics"]["iou"], default=0.0) if "iou" in history["val_metrics"] else 0.0 for epoch in range(start_epoch, num_epochs): print(f"\nEpoch {epoch + 1}/{num_epochs}") print("-" * 40) # 训练阶段 model.train() train_loss = 0 all_train_preds = [] all_train_targets = [] for images, masks in tqdm(train_loader, desc="Training"): images = images.to(device) masks = masks.to(device) outputs = model(images) outputs = sigmoid(outputs) loss = loss_fn(outputs, masks) optimizer.zero_grad() loss.backward() optimizer.step() train_loss += loss.item() preds = (outputs > 0.5).float() all_train_preds.append(preds.cpu().numpy()) all_train_targets.append(masks.cpu().numpy()) train_loss /= len(train_loader) history["train_loss"].append(train_loss) # 计算训练集指标 all_train_preds = np.concatenate(all_train_preds).flatten() all_train_targets = np.concatenate(all_train_targets).flatten() train_metrics = {} for metric in metrics: if metric == "iou": train_metrics["iou"] = iou(all_train_targets, all_train_preds) elif metric == "accuracy": train_metrics["accuracy"] = np.mean(all_train_preds == all_train_targets) elif metric == "auc": if len(np.unique(all_train_targets)) > 1: # 避免只有一类样本导致计算错误 train_metrics["auc"] = roc_auc_score(all_train_targets, all_train_preds) for metric, value in train_metrics.items(): history["train_metrics"][metric].append(value) # 打印训练结果 print(f"Training Loss: {train_loss:.4f}") for metric, value in train_metrics.items(): print(f"Training {metric.capitalize()}: {value:.4f}") # 验证阶段 model.eval() val_loss = 0 all_val_preds = [] all_val_targets = [] with torch.no_grad(): for images, masks in tqdm(val_loader, desc="Validation"): images = images.to(device) masks = masks.to(device) outputs = model(images) outputs = sigmoid(outputs) loss = loss_fn(outputs, masks) val_loss += loss.item() preds = (outputs > 0.5).float() all_val_preds.append(preds.cpu().numpy()) all_val_targets.append(masks.cpu().numpy()) val_loss /= len(val_loader) history["val_loss"].append(val_loss) # 计算验证集指标 all_val_preds = np.concatenate(all_val_preds).flatten() all_val_targets = np.concatenate(all_val_targets).flatten() val_metrics = {} for metric in metrics: if metric == "iou": intersection = np.sum(all_val_preds * all_val_targets) union = np.sum(all_val_preds) + np.sum(all_val_targets) - intersection val_metrics["iou"] = intersection / (union + 1e-6) elif metric == "accuracy": val_metrics["accuracy"] = np.mean(all_val_preds == all_val_targets) elif metric == "auc": val_metrics["auc"] = roc_auc_score(all_val_targets, all_val_preds) for metric, value in val_metrics.items(): history["val_metrics"][metric].append(value) # 打印验证结果 print(f"Validation Loss: {val_loss:.4f}") for metric, value in val_metrics.items(): print(f"Validation {metric.capitalize()}: {value:.4f}") # 保存最优模型 current_val_iou = val_metrics["iou"] if current_val_iou > best_val_iou: best_val_iou = current_val_iou torch.save(model.state_dict(), best_model_path) print(f"Best model updated with val_iou: {best_val_iou:.4f}") # 保存checkpoint save_checkpoint(model, optimizer, epoch + 1, history, checkpoint_path)
备注:内容来源于stack exchange,提问作者Shawn Pan
相关产品推荐
相关产品推荐

