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

PyTorch如何避免每轮重复计算预训练模型特征 仅训练分类层

问题解答

1. 关于仅训练分类层的实践合理性

仅冻结预训练 backbone 训练顶部分类层是迁移学习领域最常用的标准实践之一,你没有找到对应代码大概率是搜索关键词匹配问题,PyTorch生态有大量相关实现,和TensorFlow-Keras的逻辑完全共通。

2. 避免重复特征计算的方案

和你在Keras里的实现思路一致,你可以提前用冻结的预训练模型一次性提取所有训练数据的特征,序列化保存到本地,后续训练分类头时直接加载特征即可,完全不需要每轮训练都重复执行特征提取步骤,对应PyTorch实现如下:


完整PyTorch对应实现代码

import torch
import torch.nn as nn
from torch.utils.data import TensorDataset, DataLoader
from torchvision.models import mobilenet_v2, MobileNet_V2_Weights
import joblib

# --------------------------
# 步骤1:加载冻结的预训练特征提取器
# --------------------------
# 加载ImageNet预训练的MobileNetV2,对应Keras的weights='imagenet'
weights = MobileNet_V2_Weights.IMAGENET1K_V1
pretrained_model = mobilenet_v2(weights=weights)
# 去掉顶部分类层,仅保留特征提取部分,对应Keras的include_top=False
backbone = nn.Sequential(*list(pretrained_model.children())[:-1])
# 冻结所有特征提取器参数,不需要计算梯度
backbone.eval()
# 自动适配GPU/CPU
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
backbone = backbone.to(device)

# --------------------------
# 步骤2:一次性提取所有训练数据特征并保存
# --------------------------
# 这里的X、Y_train是你原始的训练图片和标签,X格式默认为(样本数, 图片尺寸, 图片尺寸, 3) 和Keras输入对齐
preprocess = weights.transforms()
# 转换为PyTorch要求的NCHW格式,再执行预处理
X_torch = torch.tensor(X).permute(0, 3, 1, 2)
X_preprocessed = preprocess(X_torch)

features_list = []
extract_batch_size = 32
# 关闭梯度计算节省显存和速度
with torch.no_grad():
    for idx in range(0, len(X_preprocessed), extract_batch_size):
        batch_x = X_preprocessed[idx:idx+extract_batch_size].to(device)
        batch_features = backbone(batch_x).cpu()
        features_list.append(batch_features)

# 拼接所有特征并压平,对应Keras的Flatten层
features_x = torch.cat(features_list, dim=0).flatten(start_dim=1)
# 保存特征到本地,对应Keras的joblib.dump逻辑
joblib.dump(features_x.numpy(), "features_x.dat")

# --------------------------
# 步骤3:定义分类头并训练
# --------------------------
# 加载保存的特征和标签
features_x = torch.tensor(joblib.load("features_x.dat"))
Y_train = torch.tensor(Y_train) # Y_train为one-hot编码,格式为(样本数, 类别数)

# 构造数据集加载器
train_dataset = TensorDataset(features_x, Y_train)
train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)

# 定义分类头,结构和你Keras实现完全对齐
class ClassifierHead(nn.Module):
    def __init__(self, input_dim, hidden_dim=100, num_classes=Y_train.shape[1]):
        super().__init__()
        self.fc1 = nn.Linear(input_dim, hidden_dim, bias=True)
        self.relu = nn.ReLU()
        self.fc2 = nn.Linear(hidden_dim, num_classes, bias=False)
        self.softmax = nn.Softmax(dim=1)
    
    def forward(self, x):
        x = self.fc1(x)
        x = self.relu(x)
        x = self.fc2(x)
        return self.softmax(x)

model = ClassifierHead(input_dim=features_x.shape[1]).to(device)
# 定义损失和优化器,和你Keras配置对齐
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

# 训练循环
Epochs = 50 # 替换为你自己的轮次数
for epoch in range(Epochs):
    model.train()
    total_loss = 0
    for batch_x, batch_y in train_loader:
        batch_x, batch_y = batch_x.to(device), batch_y.to(device)
        optimizer.zero_grad()
        outputs = model(batch_x)
        loss = criterion(outputs, batch_y)
        loss.backward()
        optimizer.step()
        total_loss += loss.item()
    print(f"Epoch {epoch+1}/{Epochs}, 平均损失: {total_loss/len(train_loader):.4f}")

补充说明

如果你的数据集太大,内存无法一次性加载所有特征,可以选择训练时实时冻结backbone的方式:不需要提前提取保存特征,只需要给backbone所有参数设置requires_grad=False,训练时每轮都会过backbone但不会更新它的参数,适合大体积数据集场景。

内容的提问来源于stack exchange,提问作者Rami Hachicha

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.23 15:21:00