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

如何在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]>s token?ViLT和BERT逻辑类似,这个token的隐藏状态是对整个图像-文本输入的全局聚合,最适合做分类任务的特征输入。
  • 如果有批量数据,只需要把多个样本的编码张量用torch.cat拼接,就能批量处理。
  • 有GPU的话记得把模型和输入张量移到GPU:model.to("cuda"),pixel_values = pixel_values.to("cuda"),训练速度会快很多。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 10:33:10