如何在HuggingFace ViLT模型之上添加PyTorch分类头
在ViLT上添加分类层实现文本标签分类
下面是针对你的需求的完整实现方案,从依赖安装到推理的代码都写得很清楚,新手也能跟着跑:
1. 安装必要依赖
先确保你装好这些基础库:
pip install transformers torch pillow
2. 加载预训练ViLT模型和处理器
我们用预训练的ViLT-VQA模型做基础,处理器负责把图像和文本转换成模型能识别的张量:
import torch import torch.nn as nn from transformers import ViltProcessor, ViltModel # 加载预训练处理器和模型 processor = ViltProcessor.from_pretrained("dandelin/vilt-base-vqa2") vilt_model = ViltModel.from_pretrained("dandelin/vilt-base-vqa2")
3. 定义带自定义分类层的模型
ViLT输出里的<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>s token隐藏状态是整个图像-文本对的聚合表示,我们基于它添加分类层:
class ViltClassifier(nn.Module): def __init__(self, vilt_model, num_labels): super().__init__() self.vilt = vilt_model # 分类层:输入是ViLT的隐藏维度(默认768),输出是你的标签总数 self.classifier = nn.Linear(vilt_model.config.hidden_size, num_labels) def forward(self, pixel_values, input_ids, attention_mask, token_type_ids): # 前向传播得到ViLT的原始输出 outputs = self.vilt( pixel_values=pixel_values, input_ids=input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids ) # 取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>s token的隐藏状态作为全局特征 cls_hidden_state = outputs.last_hidden_state[:, 0, :] # 过分类层得到各标签的logit值 logits = self.classifier(cls_hidden_state) return logits # 替换成你实际的标签数量,比如你的标签集有8个类别就写8 num_labels = 10 model = ViltClassifier(vilt_model, num_labels)
4. 数据准备示例
拿一张图片和一个问题,用处理器转换成模型需要的输入格式:
from PIL import Image # 替换成你的图片路径和问题 image = Image.open("your_image.jpg") question = "这张图里的物体是什么?" # 用处理器处理输入,返回PyTorch张量 encoding = processor(image, question, return_tensors="pt") # 提取模型需要的各个输入张量 pixel_values = encoding["pixel_values"] input_ids = encoding["input_ids"] attention_mask = encoding["attention_mask"] token_type_ids = encoding["token_type_ids"]
5. 基础训练循环
用交叉熵损失和AdamW优化器做简单训练:
# 定义损失函数和优化器 loss_fn = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=5e-5) # 示例标签(假设真实标签是第3个类别,替换成你的真实标签索引) label = torch.tensor([2]) # 训练步骤 model.train() optimizer.zero_grad() # 前向传播得到logits logits = model(pixel_values, input_ids, attention_mask, token_type_ids) # 计算损失 loss = loss_fn(logits, label) # 反向传播+更新参数 loss.backward() optimizer.step() print(f"当前训练损失: {loss.item()}")
6. 推理示例
训练完成后,用模型预测最可能的标签:
model.eval() with torch.no_grad(): logits = model(pixel_values, input_ids, attention_mask, token_type_ids) # 取概率最高的标签索引 predicted_label_idx = torch.argmax(logits, dim=1).item() # 替换成你的真实标签列表,比如["猫", "狗", "杯子"...] labels = ["类别1", "类别2", "类别3", "类别4", "类别5", "类别6", "类别7", "类别8", "类别9", "类别10"] print(f"预测结果: {labels[predicted_label_idx]}")
关键说明
- 为什么用
<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>stoken?ViLT和BERT逻辑类似,这个token的隐藏状态是对整个图像-文本输入的全局聚合,最适合做分类任务的特征输入。 - 如果有批量数据,只需要把多个样本的编码张量用
torch.cat拼接,就能批量处理。 - 有GPU的话记得把模型和输入张量移到GPU:
model.to("cuda"),pixel_values = pixel_values.to("cuda"),训练速度会快很多。
内容的提问来源于stack exchange,提问作者user10418143
相关产品推荐
相关产品推荐

