You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

相同配置下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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.15 03:44:31