news 2026/9/10 5:49:53

Transformers 图像描述实战:基于 GIT 与 Pokémon BLIP 数据集完成微调与推理

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Transformers 图像描述实战:基于 GIT 与 Pokémon BLIP 数据集完成微调与推理

Transformers 图像描述实战:基于 GIT 与 Pokémon BLIP 数据集完成微调与推理

【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers

本文以 🤗 Transformers 官方任务指南为骨架,完整讲解图像描述(Image Captioning)任务从数据集加载、预处理、模型微调到推理的端到端流程。指南以microsoft/git-base为基座模型、lambdalabs/pokemon-blip-captions为训练数据,训练完成后即可用Trainer训练出的模型对任意图片生成自然语言描述。读完本文,你将掌握图像描述任务的完整数据管线、GIT 模型的加载方式、基于 WER 的评估方法,以及generate推理的调用方式,并了解这些操作在仓库源码中的实现位置。

任务背景与前置准备

图像描述(Image Captioning)是预测给定图片自然语言描述的任务,典型现实应用包括帮助视障人士理解画面内容、提升内容的可访问性等。因此,图像描述模型需要同时理解视觉模态与文本模态。

开始之前,请确保已安装必要的库:

pip install transformers datasets evaluate -q pip install jiwer -q

其中jiwer用于计算 Word Error Rate(WER)评估指标,evaluate是 🤗 Evaluate 评估库,datasets用于加载和预处理数据集。

如果你希望将训练好的模型上传共享给社区,建议先登录 Hugging Face 账号:

from huggingface_hub import notebook_login notebook_login()

按提示输入 token 即可完成登录;Trainer在设置了push_to_hub=True后会自动将模型推送至 Hub。

加载 Pokémon BLIP 图像描述数据集

使用 🤗 Datasets 库加载由 {image, caption} 配对组成的数据集。本指南使用lambdalabs/pokemon-blip-captions(Pokémon 图片及其 BLIP 生成的描述),也可以参考该数据集的创建方式构建自己的图像描述数据集。

from datasets import load_dataset ds = load_dataset("lambdalabs/pokemon-blip-captions") ds

输出如下:

DatasetDict({ train: Dataset({ features: ['image', 'text'], num_rows: 833 }) })

数据集包含imagetext两个特征:image为图片像素数据,text为对应的描述文本。833 行数据中全部位于 train 分区。

提示:许多图像描述数据集为每张图片提供多条候选描述。这种情况下,常见的训练策略是每次训练时从可用描述中随机采样一条,以增强模型的泛化能力。

接着用train_test_split方法将训练集按 10% 比例划分为训练集和测试集:

ds = ds["train"].train_test_split(test_size=0.1) train_ds = ds["train"] test_ds = ds["test"]

可视化训练集中的若干样本,便于直观理解数据内容:

from textwrap import wrap import matplotlib.pyplot as plt import numpy as np def plot_images(images, captions): plt.figure(figsize=(20, 20)) for i in range(len(images)): ax = plt.subplot(1, len(images), i + 1) caption = captions[i] caption = "\n".join(wrap(caption, 12)) plt.title(caption) plt.imshow(images[i]) plt.axis("off") sample_images_to_visualize = [np.array(train_ds[i]["image"]) for i in range(5)] sample_captions = [train_ds[i]["text"] for i in range(5)] plot_images(sample_images_to_visualize, sample_captions)

数据预处理:图像与文本的双模态管线

由于数据集包含图像和文本两种模态,预处理管线需要同时处理二者。做法是加载与待微调模型关联的 processor 类:

from transformers import AutoProcessor checkpoint = "microsoft/git-base" processor = AutoProcessor.from_pretrained(checkpoint)

processor 内部会完成图像的预处理(包括尺寸调整、像素缩放)与描述的 tokenize。在仓库源码中,GIT 对应的 processor 为GitProcessor,它继承自ProcessorMixin,本质上是"图像处理器 + 分词器"的组合体,因此AutoProcessor.from_pretrained会自动加载配套的GitImageProcessor(负责 resize、像素缩放)和BertTokenizer(负责文本 tokenize)。

定义数据变换函数,将图像和文本统一转换为模型输入:

def transforms(example_batch): images = [x for x in example_batch["image"]] captions = [x for x in example_batch["text"]] inputs = processor(images=images, text=captions, padding="max_length") inputs.update({"labels": inputs["input_ids"]}) return inputs train_ds.set_transform(transforms) test_ds.set_transform(transforms)

这里有两个关键点:

  • processor(images=..., text=..., padding="max_length")返回的input_ids即描述文本的 token 序列,同时内部会对图像做尺寸调整(GIT 默认输入为 224×224)与像素归一化;
  • labels直接复用input_ids,用于计算自回归语言建模损失(next token prediction)。这与GitForCausalLM.forwardlabels的语义一致——标签为(batch_size, sequence_length)的 token 索引,token 值为-100的位置会被忽略(见 modeling_git.py)。

数据集就绪后,即可进入模型微调阶段。

加载基座模型

microsoft/git-base加载为AutoModelForCausalLM对象:

from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained(checkpoint)

从源码看,仓库的自动映射表已将模型类型"git"映射到GitForCausalLM(见 modeling_auto.py)。GitForCausalLMGitModel(视觉编码器 + 文本 Transformer)与一个nn.Linear(config.hidden_size, config.vocab_size)输出层组成,并继承GenerationMixin,因此天然支持generate自回归生成。

GIT 的结构可以从配置类中得到印证(见 configuration_git.py):

  • GitVisionConfig:视觉编码器配置,默认image_size=224patch_size=16、12 层 Transformer、hidden size 768;
  • GitConfig:整体配置,默认vocab_size=30522、6 层文本 Transformer、max_position_embeddings=1024,并内嵌vision_config子配置;num_image_with_embedding参数用于视频描述/VQA 场景追加时间嵌入。

评估:使用 Word Error Rate

图像描述模型通常使用 Rouge Score 或 Word Error Rate(WER)评估。本指南采用 WER。使用 🤗 Evaluate 库实现:

from evaluate import load import torch wer = load("wer") def compute_metrics(eval_pred): logits, labels = eval_pred predicted = logits.argmax(-1) decoded_labels = processor.batch_decode(labels, skip_special_tokens=True) decoded_predictions = processor.batch_decode(predicted, skip_special_tokens=True) wer_score = wer.compute(predictions=decoded_predictions, references=decoded_labels) return {"wer_score": wer_score}

compute_metricsTrainer每次评估时被调用:logits.argmax(-1)取每个位置概率最大的 token id,随后用 processor 的batch_decode(..., skip_special_tokens=True)将 id 序列解码为文本(跳过[CLS][SEP][PAD]等特殊 token),最后与参考文本计算 WER。WER 的潜在局限与注意事项可参考其指标说明(如对同义词、词序变化的敏感性)。

微调训练

使用 🤗Trainer完成微调。首先通过TrainingArguments定义训练参数:

from transformers import TrainingArguments, Trainer model_name = checkpoint.split("/")[1] training_args = TrainingArguments( output_dir=f"{model_name}-pokemon", learning_rate=5e-5, num_train_epochs=50, fp16=True, per_device_train_batch_size=32, per_device_eval_batch_size=32, gradient_accumulation_steps=2, save_total_limit=3, eval_strategy="steps", eval_steps=50, save_strategy="steps", save_steps=50, logging_steps=50, remove_unused_columns=False, push_to_hub=True, label_names=["labels"], load_best_model_at_end=True, )

各参数含义与配置要点:

参数取值说明
output_dirgit-base-pokemon模型检查点输出目录,由model_name拼接而来
learning_rate5e-5峰值学习率
num_train_epochs50训练轮数(Pokémon 数据集规模小,需较多 epoch 收敛)
fp16True启用混合精度训练,需 GPU 支持;无 GPU 时应设为False
per_device_train/eval_batch_size32单设备 batch 大小,需根据显存调整
gradient_accumulation_steps2梯度累积步数,等效放大 batch 为 64
save_total_limit3最多保留 3 个检查点,防止磁盘膨胀
eval_strategy/eval_stepssteps/50每 50 步评估一次
save_strategy/save_stepssteps/50每 50 步保存一次检查点
logging_steps50每 50 步记录一次日志
remove_unused_columnsFalse关键!保留原始image列等非模型输入列,供自定义transforms使用
push_to_hubTrue训练后自动推送模型到 Hub(需已登录)
label_names["labels"]显式指定标签列名,配合load_best_model_at_end在结束时加载最优检查点

然后将模型、数据集与评估函数一起交给Trainer

trainer = Trainer( model=model, args=training_args, train_dataset=train_ds, eval_dataset=test_ds, compute_metrics=compute_metrics, )

启动训练:

trainer.train()

训练过程中可观察到训练损失平滑下降。训练完成后,用push_to_hub将模型共享到 Hub,方便社区直接使用:

trainer.push_to_hub()

推理:生成图像描述

test_ds取一张样本图片测试模型:

from PIL import Image import requests url = "https://huggingface.co/datasets/sayakpaul/sample-datasets/resolve/main/pokemon.png" image = Image.open(requests.get(url, stream=True).raw) image

为模型准备图像输入:

from accelerate import Accelerator device = Accelerator().device inputs = processor(images=image, return_tensors="pt").to(device) pixel_values = inputs.pixel_values

这里使用Accelerator().device自动选择 GPU/CPU 设备。processor 输出的pixel_values即为 GIT 视觉编码器的输入张量——与源码中GitForCausalLM.forwardpixel_values参数对应(见 modeling_git.py)。

调用generate解码预测结果:

generated_ids = model.generate(pixel_values=pixel_values, max_length=50) generated_caption = processor.batch_decode(generated_ids, skip_special_tokens=True)[0] print(generated_caption)

输出示例:

a drawing of a pink and blue pokemon

微调后的模型生成的描述质量相当不错。generateGenerationMixin提供,max_length=50限制生成的最大 token 数;生成的 id 序列经batch_decode去掉特殊 token 后即为最终描述文本。

值得注意的是,同样的GitForCausalLM模型在仓库的源码 docstring 中还展示了两个相关用法(见 modeling_git.py):

  • VQA(视觉问答):传入pixel_values的同时,将问题文本 tokenize 后以input_ids传入generate,即可让模型根据图片回答问题;
  • 视频描述:加载microsoft/git-base-vatex检查点,配合num_image_with_embedding时间嵌入处理多帧输入。

这体现了 GIT 架构"视觉编码器 + 因果语言模型"的通用性:同一套微调流程稍作改动即可迁移到不同多模态生成任务。

小结

本文完整复现了图像描述任务的端到端流程:加载并划分pokemon-blip-captions数据集 → 用AutoProcessor构建图像 + 文本双模态预处理管线 → 加载microsoft/git-baseAutoModelForCausalLM→ 以 WER 为评估指标、Trainer完成 50 个 epoch 的微调 → 用generate对新图片生成描述。对应的全部实现细节均可在仓库的 GIT 模型目录 中找到:GitProcessor组合了图像处理器与分词器,GitConfig/GitVisionConfig定义了双模态配置,GitForCausalLM则基于GenerationMixin提供统一的生成入口。这套方法论不仅适用于图像描述,稍加改动即可扩展到 VQA、视频描述等更多多模态任务。

【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 5:48:20

Wiki与RAG不是二选一:知识存储与调用的协同架构

1. 这不是“选一个”,而是“搭一套”:Wiki 和 RAG 的本质分工错位很多人看到标题“Wiki 和 RAG 如何选择”,第一反应是:我该用 Wiki 做知识库,还是用 RAG 做知识库?——这个提问本身,就踩进了最…

作者头像 李华
网站建设 2026/9/10 5:47:56

AI文本人性化改写:特征检测与自然度优化实战指南

1. 先看清楚:humanizer 要解决的是哪种“AI味” 1.1 AI 文本的指纹到底藏在哪里 我做了两年多的内容工具链开发,接触过大量 AI 生成的初稿。说实话,绝大多数人抱怨“一眼假”,并不是因为内容本身有事实错误,而是文本的…

作者头像 李华
网站建设 2026/9/10 5:47:25

DeepSeek Harness实测:插件化架构如何重塑AI工具生态

最近几天打开技术社区,满屏都是 DeepSeek Harness 的讨论。有人把它捧成"AI 时代的 Chrome",也有人泼冷水说不过是又一轮插件生态圈地。作为一个从命令行时代就开始折腾各种工具链的老玩家,我花了整整一个周末把 Harness 从安装到深…

作者头像 李华
网站建设 2026/9/10 5:46:37

2030年软件测试趋势:AI接管重复劳动,质量工程成核心

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华