BEiT模型beit_base_patch16_224.in22k_ft_in22k迁移学习指南:如何在自己的数据集上微调

BEiT模型beit_base_patch16_224.in22k_ft_in22k迁移学习指南:如何在自己的数据集上微调

【免费下载链接】beit_base_patch16_224.in22k_ft_in22k 【免费下载链接】beit_base_patch16_224.in22k_ft_in22k 项目地址: https://ai.gitcode.com/hf_mirrors/timm/beit_base_patch16_224.in22k_ft_in22k

beit_base_patch16_224.in22k_ft_in22k是一款基于BEiT架构的图像分类模型,它在ImageNet-22k数据集上通过自监督掩码图像建模(MIM)预训练,并在ImageNet-22k上进行了微调。本指南将为新手和普通用户提供简单快速的方法,教你如何在自己的数据集上对该模型进行迁移学习微调。

模型基础认知:为什么选择beit_base_patch16_224.in22k_ft_in22k?

beit_base_patch16_224.in22k_ft_in22k作为一款强大的图像分类/特征骨干模型,具有以下优势:

  • 参数规模:102.6M参数,能够捕捉图像中的丰富特征
  • 计算效率:17.6 GMACs,在性能与效率间取得平衡
  • 输入规格:支持224x224分辨率图像输入
  • 预训练基础:基于ImageNet-22k大规模数据集训练,具备良好的特征提取能力

该模型基于两篇重要论文构建:

  • BEiT: BERT Pre-Training of Image Transformers(https://arxiv.org/abs/2106.08254)
  • An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale(https://arxiv.org/abs/2010.11929v2)

准备工作:环境与数据集准备

1. 安装必要依赖

首先确保你的环境中安装了timm库和PyTorch:

pip install timm torch torchvision

2. 获取模型代码库

克隆模型仓库到本地:

git clone https://gitcode.com/hf_mirrors/timm/beit_base_patch16_224.in22k_ft_in22k
cd beit_base_patch16_224.in22k_ft_in22k

3. 数据集准备

你的数据集应按照以下结构组织:

dataset/
├── train/
│   ├── class1/
│   │   ├── img1.jpg
│   │   └── img2.jpg
│   └── class2/
│       ├── img1.jpg
│       └── img2.jpg
└── val/
    ├── class1/
    │   ├── img1.jpg
    │   └── img2.jpg
    └── class2/
        ├── img1.jpg
        └── img2.jpg

确保图像分辨率接近224x224,或准备好图像预处理步骤。

快速开始:微调模型的完整步骤

加载预训练模型

首先加载预训练模型,并查看其配置信息:

import timm

# 加载模型
model = timm.create_model('beit_base_patch16_224.in22k_ft_in22k', pretrained=True)
print(model)

# 获取模型配置
data_config = timm.data.resolve_model_data_config(model)
print("模型配置:", data_config)

调整模型输出层

根据你的数据集类别数量修改模型的输出层:

num_classes = 10  # 替换为你的数据集类别数
model.reset_classifier(num_classes=num_classes)

数据预处理

使用模型特定的预处理转换:

transforms = timm.data.create_transform(**data_config, is_training=True)

准备数据加载器

创建DataLoader来加载你的数据集:

from torchvision.datasets import ImageFolder
from torch.utils.data import DataLoader

train_dataset = ImageFolder('dataset/train', transform=transforms)
val_dataset = ImageFolder('dataset/val', transform=transforms)

train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4)

设置训练参数

配置优化器、损失函数和学习率调度器:

import torch
import torch.optim as optim
from torch.optim.lr_scheduler import StepLR

criterion = torch.nn.CrossEntropyLoss()
optimizer = optim.AdamW(model.parameters(), lr=1e-4)
scheduler = StepLR(optimizer, step_size=10, gamma=0.1)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)

模型训练循环

实现简单的训练循环:

num_epochs = 20

for epoch in range(num_epochs):
    model.train()
    train_loss = 0.0
    
    for inputs, labels in train_loader:
        inputs, labels = inputs.to(device), labels.to(device)
        
        optimizer.zero_grad()
        outputs = model(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        
        train_loss += loss.item() * inputs.size(0)
    
    train_loss = train_loss / len(train_loader.dataset)
    scheduler.step()
    
    # 验证阶段
    model.eval()
    val_loss = 0.0
    correct = 0
    total = 0
    
    with torch.no_grad():
        for inputs, labels in val_loader:
            inputs, labels = inputs.to(device), labels.to(device)
            outputs = model(inputs)
            loss = criterion(outputs, labels)
            
            val_loss += loss.item() * inputs.size(0)
            _, predicted = torch.max(outputs.data, 1)
            total += labels.size(0)
            correct += (predicted == labels).sum().item()
    
    val_loss = val_loss / len(val_loader.dataset)
    val_acc = correct / total
    
    print(f'Epoch {epoch+1}/{num_epochs}')
    print(f'Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}')

保存微调后的模型

训练完成后保存模型权重:

torch.save(model.state_dict(), 'beit_finetuned.pth')

模型使用:如何用微调后的模型进行预测

加载微调后的模型并进行图像分类:

from PIL import Image
import torch

# 加载模型
model = timm.create_model('beit_base_patch16_224.in22k_ft_in22k', num_classes=num_classes)
model.load_state_dict(torch.load('beit_finetuned.pth'))
model.eval()
model.to(device)

# 获取预处理转换
data_config = timm.data.resolve_model_data_config(model)
transforms = timm.data.create_transform(**data_config, is_training=False)

# 加载并预处理图像
img = Image.open('test_image.jpg').convert('RGB')
img_tensor = transforms(img).unsqueeze(0).to(device)

# 预测
with torch.no_grad():
    output = model(img_tensor)
    probabilities = torch.nn.functional.softmax(output[0], dim=0)
    top5_prob, top5_catid = torch.topk(probabilities, 5)

# 输出结果
for i in range(top5_prob.size(0)):
    print(f"{train_dataset.classes[top5_catid[i]]}: {top5_prob[i].item():.4f}")

微调技巧:提升模型性能的关键策略

1. 学习率调整

  • 初始学习率建议设置在1e-5到1e-4之间
  • 使用学习率调度器(如余弦退火)可以获得更好的效果
  • 可以对不同层使用不同的学习率(分层学习率)

2. 数据增强

  • 适当增加数据增强可以防止过拟合
  • 常用增强方法:随机裁剪、翻转、旋转、色彩抖动等
  • 注意不要过度增强导致数据失真

3. 批处理大小

  • 尽可能使用大的批处理大小
  • 如果GPU内存不足,可以使用梯度累积

4. 早停策略

  • 监控验证集性能,当性能不再提升时停止训练
  • 保存验证集性能最佳的模型

常见问题解决:微调过程中的挑战

过拟合问题

如果模型在训练集上表现良好但在验证集上表现不佳:

  • 增加数据增强
  • 使用正则化技术(如 dropout)
  • 减少训练轮次
  • 尝试更小的学习率

训练速度慢

  • 使用更大的批处理大小
  • 启用混合精度训练
  • 使用更多的GPU进行分布式训练

内存不足

  • 减小批处理大小
  • 使用梯度检查点
  • 降低图像分辨率(需谨慎,可能影响性能)

总结:让beit_base_patch16_224.in22k_ft_in22k为你所用

beit_base_patch16_224.in22k_ft_in22k作为一个预训练的图像Transformer模型,为各种图像分类任务提供了强大的起点。通过本文介绍的微调方法,你可以快速将其适应自己的特定数据集,而无需从头开始训练一个复杂的模型。

记住,迁移学习的关键是找到适合你数据集的微调策略。开始时可以使用本文提供的默认参数,然后根据模型在验证集上的表现进行调整。祝你在项目中取得成功!

引用

如果你的工作中使用了beit_base_patch16_224.in22k_ft_in22k模型,请引用以下论文:

@article{bao2021beit,
  title={Beit: Bert pre-training of image transformers},
  author={Bao, Hangbo and Dong, Li and Piao, Songhao and Wei, Furu},
  journal={arXiv preprint arXiv:2106.08254},
  year={2021}
}

@article{dosovitskiy2020vit,
  title={An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale},
  author={Dosovitskiy, Alexey and Beyer, Lucas and Kolesnikov, Alexander and Weissenborn, Dirk and Zhai, Xiaohua and Unterthiner, Thomas and  Dehghani, Mostafa and Minderer, Matthias and Heigold, Georg and Gelly, Sylvain and Uszkoreit, Jakob and Houlsby, Neil},
  journal={ICLR},
  year={2021}
}

@misc{rw2019timm,
  author = {Ross Wightman},
  title = {PyTorch Image Models},
  year = {2019},
  publisher = {GitHub},
  journal = {GitHub repository},
  doi = {10.5281/zenodo.4414861},
  howpublished = {\url{https://github.com/huggingface/pytorch-image-models}}
}

【免费下载链接】beit_base_patch16_224.in22k_ft_in22k 【免费下载链接】beit_base_patch16_224.in22k_ft_in22k 项目地址: https://ai.gitcode.com/hf_mirrors/timm/beit_base_patch16_224.in22k_ft_in22k

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值