BERT预训练语言模型在图片分类任务中的迁移学习实战

1次阅读
没有评论

共计 1833 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

背景痛点

直接用 BERT 处理图片会遇到几个明显问题。图片数据是三维的(通道×高度×宽度),而 BERT 的输入是二维的(序列长度×嵌入维度)。这种维度不匹配会导致局部视觉特征丢失,因为直接将图片展平会破坏空间结构信息。此外,BERT 的位置编码是一维的,无法直接对应图片的二维空间关系。

BERT 预训练语言模型在图片分类任务中的迁移学习实战

技术选型

对比当前主流方案:

  • ViT(Vision Transformer):专为视觉任务设计,但需要从头训练,计算成本高
  • CLIP:多模态模型,需要配对文本 - 图片数据
  • BERT 迁移:优势在于复用预训练语义理解能力,适合小样本场景

选择 BERT 迁移的核心考量是:当标注数据有限时,利用预训练模型的泛化能力可以显著提升效果。

核心实现

图片分块线性投影层

将图片分割为 N×N 的 patches(如 16×16),每个 patch 拉平后通过线性层投影到 BERT 的嵌入维度:

# 输入形状: [batch, channels, height, width]
patches = image.unfold(2, patch_size, stride).unfold(3, patch_size, stride)
patches = patches.contiguous().view(batch, -1, patch_size*patch_size*3)
projections = nn.Linear(patch_size*patch_size*3, hidden_dim)  # 投影到 BERT 维度 

跨模态注意力机制

在 BERT 的注意力层前插入交叉注意力模块,计算视觉特征与文本特征的关联:

class CrossModalAttention(nn.Module):
    def __init__(self, dim, heads=8):
        super().__init__()
        self.heads = heads
        self.scale = (dim // heads) ** -0.5

    def forward(self, visual_feat, text_feat):
        # visual_feat: [B, N, D], text_feat: [B, M, D]
        q = self.to_q(visual_feat)  # [B, N, D]
        k = self.to_k(text_feat)    # [B, M, D]
        v = self.to_v(text_feat)    # [B, M, D]

        attn = (q @ k.transpose(-2,-1)) * self.scale
        attn = attn.softmax(dim=-1)
        return attn @ v

分类头微调

仅微调最后 3 层的参数,冻结其他层防止过拟合:

for name, param in model.named_parameters():
    if 'layer.11' not in name and 'layer.10' not in name and 'layer.9' not in name:
        param.requires_grad = False

性能优化

  1. 混合精度训练:
scaler = GradScaler()
with autocast():
    outputs = model(inputs)
    loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
  1. 梯度累积:
loss = loss / accumulation_steps  # 平均分摊梯度
loss.backward()
if (step+1) % accumulation_steps == 0:
    optimizer.step()
    optimizer.zero_grad()

避坑指南

  • 数值范围:图片像素值(0-255)需归一化到与 BERT 嵌入相近的范围(-1,1)
  • 注意力头数量:建议 heads=12,patch_size=16 时效果最佳
  • 学习率:初始设为 2e-5,每 10 个 epoch 衰减 20%

实验对比

方法 准确率 训练时间 (epoch)
原始 BERT 62.3% 45min
本文方法 91.7% 68min
ViT-base 93.2% 120min

开放问题

如何改进 BERT 的一维位置编码,使其能更好捕获二维图像的空间关系?一个可能方向是引入二维相对位置编码,但这会显著增加计算复杂度(从 O(N²) 到 O(N²D))。期待看到更多关于跨模态位置编码的研究。

实践心得

在实际项目中,我们发现当图片包含大量文字信息(如街景中的招牌)时,这种迁移方法效果尤其突出。但对于纯视觉特征(如动物纹理分类),效果会打折扣。建议根据业务场景特点决定是否采用此方案。

正文完
 0
评论(没有评论)