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

BLIP微调问题:生成字幕中特殊令牌始终偏向单一类别

BLIP模型微调问题:特殊令牌生成异常

目标

让BLIP生成包含特殊令牌的分类字幕,示例:

"A greasy pizza box with leftover cheese. [XXY_CONTAM]"
"A clean, dry pizza box. [CTX_CLEAN]"

训练代码(train.py)

from transformers import BlipProcessor, BlipForConditionalGeneration, Trainer, TrainingArguments
from datasets import Dataset
from PIL import Image
import torch, os, json

# === Paths ===
model_path = "D:/models/blip-caption"
data_path = "D:/BLIP/captions.json"
image_root = "D:/BLIP"
output_dir = "D:/BLIP/blip-finetuned-pizza"
log_dir = "D:/BLIP/logs"

# === Load dataset JSON ===
with open(data_path, "r") as f:
    data = json.load(f)
dataset = Dataset.from_list(data)

# === Load processor and model ===
processor = BlipProcessor.from_pretrained(model_path)
model = BlipForConditionalGeneration.from_pretrained(model_path)

# === Add special tokens ===
special_tokens_dict = {"additional_special_tokens": ["[CTX_CLEAN]", "[XXY_CONTAM]"]}
processor.tokenizer.add_special_tokens(special_tokens_dict)
model.resize_token_embeddings(len(processor.tokenizer))

# === Preprocess function ===
def preprocess(example):
    image_path = os.path.join(image_root, example["image"])
    image = Image.open(image_path).convert("RGB")
    inputs = processor(
        images=image,
        text=example["caption"],
        return_tensors="pt",
        padding="max_length",
        truncation=True,
        max_length=32
    )
    inputs = {k: v.squeeze(0) for k, v in inputs.items()}
    inputs["labels"] = inputs["input_ids"]
    return inputs

processed_dataset = dataset.map(preprocess)

# === Training arguments ===
training_args = TrainingArguments(
    output_dir=output_dir,
    per_device_train_batch_size=2,
    num_train_epochs=10,
    logging_dir=log_dir,
    logging_steps=5,
    save_steps=20,
    save_total_limit=1,
    remove_unused_columns=False,
    fp16=torch.cuda.is_available()
)

# === Trainer ===
trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=processed_dataset,
)

# === Train and Save ===
trainer.train()
model.save_pretrained(output_dir)
processor.tokenizer.save_pretrained(output_dir)
processor.save_pretrained(output_dir)

print("Fine-tuned model saved to", output_dir)
print("Special token IDs:", processor.tokenizer.convert_tokens_to_ids(["[CTX_CLEAN]", "[XXY_CONTAM]"]))

推理代码

from transformers import BlipProcessor, BlipForConditionalGeneration
from PIL import Image
import torch

# === Load model and processor ===
model_path = "D:/BLIP/blip-finetuned-pizza"
image_path = "D:/Models/pizzabox.jpg"

processor = BlipProcessor.from_pretrained(model_path)
model = BlipForConditionalGeneration.from_pretrained(model_path)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)

image = Image.open(image_path).convert("RGB")
inputs = processor(images=image, return_tensors="pt").to(device)

with torch.no_grad():
    output = model.generate(**inputs, max_length=32)
    caption = processor.decode(output[0], skip_special_tokens=False)

print("Generated Caption:", caption)

问题现象

  • 数据集均衡,特殊令牌已在分词器和模型中正确注册
  • 模型始终输出additional_special_tokens列表中的第一个令牌:
    • 列表顺序为["[CTX_CLEAN]", "[XXY_CONTAM]"]时,始终生成[CTX_CLEAN]
    • 反转顺序为["[XXY_CONTAM]", "[CTX_CLEAN]"]时,始终生成[XXY_CONTAM]
  • 模型未学会根据图像内容匹配对应令牌,仅偏好列表首位令牌

疑问点

  1. 如何让BLIP依据图像输出正确的特殊令牌?
  2. 特殊令牌、字幕或预处理流程是否存在问题?
  3. 是否需要手动设置令牌嵌入或在生成时强制令牌使用?

解决方案建议

1. 修正标签处理逻辑

当前直接将input_ids作为labels的方式错误,BLIP字幕生成任务中,标签需要忽略padding部分(设为-100),且要分离图像与文本编码:

def preprocess(example):
    image_path = os.path.join(image_root, example["image"])
    image = Image.open(image_path).convert("RGB")
    # 分离图像和文本编码
    image_inputs = processor(images=image, return_tensors="pt")
    text_inputs = processor(text=example["caption"], return_tensors="pt", padding="max_length", truncation=True, max_length=32)
    
    inputs = {
        "pixel_values": image_inputs["pixel_values"].squeeze(0),
        "input_ids": text_inputs["input_ids"].squeeze(0),
        "attention_mask": text_inputs["attention_mask"].squeeze(0)
    }
    # 构建labels:padding部分设为-100,不参与损失计算
    labels = text_inputs["input_ids"].squeeze(0).clone()
    labels[labels == processor.tokenizer.pad_token_id] = -100
    inputs["labels"] = labels
    return inputs

2. 强化特殊令牌训练信号

  • 固定数据集字幕中特殊令牌的位置(比如强制放在句末),让模型学习位置规律
  • 自定义损失函数,提高特殊令牌的损失权重:
from torch.nn import CrossEntropyLoss

class CustomTrainer(Trainer):
    def compute_loss(self, model, inputs, return_outputs=False):
        labels = inputs.pop("labels")
        outputs = model(**inputs)
        logits = outputs.logits
        
        # 给特殊令牌设置更高损失权重
        special_token_ids = processor.tokenizer.convert_tokens_to_ids(["[CTX_CLEAN]", "[XXY_CONTAM]"])
        weight = torch.ones(logits.shape[-1]).to(logits.device)
        for token_id in special_token_ids:
            weight[token_id] = 5.0
        
        loss_fct = CrossEntropyLoss(weight=weight)
        loss = loss_fct(logits.view(-1, logits.shape[-1]), labels.view(-1))
        
        return (loss, outputs) if return_outputs else loss

# 替换原Trainer
trainer = CustomTrainer(
    model=model,
    args=training_args,
    train_dataset=processed_dataset,
)

3. 调整推理生成策略

默认生成策略易陷入局部最优,调整参数提升多样性:

with torch.no_grad():
    output = model.generate(
        **inputs,
        max_length=32,
        num_beams=5,
        temperature=0.7,
        forced_eos_token_id=processor.tokenizer.eos_token_id,
        do_sample=True
    )
    caption = processor.decode(output[0], skip_special_tokens=False)

4. 初始化特殊令牌嵌入

新添加的特殊令牌默认用随机嵌入,可初始化语义相近的令牌嵌入:

# 获取相近语义令牌的嵌入
clean_emb = model.get_input_embeddings()(processor.tokenizer.convert_tokens_to_ids("clean"))
dirty_emb = model.get_input_embeddings()(processor.tokenizer.convert_tokens_to_ids("dirty"))

# 赋值给特殊令牌
special_token_ids = processor.tokenizer.convert_tokens_to_ids(["[CTX_CLEAN]", "[XXY_CONTAM]"])
model.get_input_embeddings().weight.data[special_token_ids[0]] = clean_emb
model.get_input_embeddings().weight.data[special_token_ids[1]] = dirty_emb

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 13:44:55