Llava-v1.6-7b模型微调实战:自定义视觉分类任务
Llava-v1.6-7b模型微调实战:自定义视觉分类任务
1. 引言
你是不是遇到过这样的情况:手头有一堆特定领域的图片,想要让AI模型帮你自动分类,但现成的视觉模型总是表现不佳?比如医疗影像分析、工业质检、或者特定场景的物体识别,通用模型往往力不从心。
今天我们就来解决这个问题。我将带你一步步微调Llava-v1.6-7b模型,让它成为你专属的视觉分类专家。不需要深厚的机器学习背景,只要跟着做,你就能让模型学会识别你关心的特定视觉类别。
Llava-v1.6-7b是个多模态模型,既能理解图像又能处理文本,这让我们可以用自然语言来描述分类任务。相比于传统的视觉分类模型,它的优势在于能够理解更复杂的指令和上下文。
2. 环境准备与快速部署
2.1 硬件要求
首先看看你需要什么样的硬件环境。Llava-v1.6-7b是个7B参数的模型,微调时需要一定的计算资源:
- GPU内存:建议24GB以上(如RTX 3090、A10等)
- 系统内存:至少32GB RAM
- 存储空间:需要20-30GB空间存放模型和数据集
如果你的显存不够,也可以使用量化技术,但效果可能会打些折扣。
2.2 软件环境安装
我们来快速搭建环境。建议使用conda来管理环境:
conda create -n llava-finetune python=3.10 -y
conda activate llava-finetune
# 安装主要依赖
pip install torch torchvision torchaudio
pip install transformers accelerate datasets
pip install pillow opencv-python
# 安装Llava相关包
git clone https://github.com/haotian-liu/LLaVA.git
cd LLaVA
pip install -e .
如果你的CUDA版本比较新,可能需要安装对应版本的PyTorch。这些命令应该能在10分钟内完成环境搭建。
3. 数据准备与处理
3.1 数据集格式
微调需要准备标注好的图像数据。Llava使用一种特殊的对话格式,对于分类任务,我们可以这样组织数据:
{
"id": "unique_id_001",
"image": "path/to/image1.jpg",
"conversations": [
{
"from": "human",
"value": "<image>\n请问这张图片属于哪个类别?"
},
{
"from": "gpt",
"value": "这是一张猫的图片"
}
]
}
每个样本包含图片路径和一段对话,人类问问题,模型回答分类结果。
3.2 实际数据准备示例
假设你在做一个动物分类任务,包含猫、狗、鸟三个类别。你需要:
- 收集每个类别至少100-200张图片
- 按照上述格式创建标注文件
- 将数据集分为训练集和验证集(建议8:2)
这里有个简单的Python脚本来帮你生成标注文件:
import json
import os
from pathlib import Path
def create_dataset_json(image_dir, output_path):
dataset = []
categories = ["猫", "狗", "鸟"]
for category in categories:
category_dir = os.path.join(image_dir, category)
if not os.path.exists(category_dir):
continue
for img_file in os.listdir(category_dir):
if img_file.lower().endswith(('.png', '.jpg', '.jpeg')):
sample = {
"id": f"{category}_{img_file}",
"image": os.path.join(category, img_file),
"conversations": [
{
"from": "human",
"value": "<image>\n请问这张图片属于哪个类别?"
},
{
"from": "gpt",
"value": f"这是一张{category}的图片"
}
]
}
dataset.append(sample)
with open(output_path, 'w', encoding='utf-8') as f:
json.dump(dataset, f, ensure_ascii=False, indent=2)
# 使用示例
create_dataset_json("animals_dataset", "train.json")
4. 微调配置与训练
4.1 训练参数设置
现在来到关键部分——训练配置。创建一个配置文件train_config.json:
{
"model_name_or_path": "liuhaotian/llava-v1.6-vicuna-7b",
"version": "v1",
"data_path": "train.json",
"image_folder": "animals_dataset",
"vision_tower": "openai/clip-vit-large-patch14-336",
"mm_vision_select_layer": -2,
"mm_use_im_start_end": false,
"bf16": true,
"output_dir": "./llava-finetuned",
"num_train_epochs": 3,
"per_device_train_batch_size": 4,
"per_device_eval_batch_size": 4,
"gradient_accumulation_steps": 4,
"evaluation_strategy": "steps",
"eval_steps": 100,
"save_strategy": "steps",
"save_steps": 200,
"save_total_limit": 1,
"learning_rate": 2e-5,
"weight_decay": 0.0,
"warmup_ratio": 0.03,
"lr_scheduler_type": "cosine",
"logging_steps": 10,
"tf32": true,
"model_max_length": 2048,
"gradient_checkpointing": true,
"dataloader_pin_memory": false,
"mm_projector_lr": 2e-5,
"group_by_modality_length": false
}
这些参数对新手比较友好,学习率设得不高不低,训练轮数也适中,既不会欠拟合也不会过拟合。
4.2 开始训练
运行训练命令:
cd LLaVA
torchrun --nproc_per_node=1 llava/train/train_mem.py \
--train_config train_config.json \
--deepspeed zero2.json
训练过程中你会看到损失值逐渐下降。如果一切正常,3个epoch后你应该能得到一个不错的模型。
5. 模型评估与测试
5.1 验证集评估
训练完成后,我们需要看看模型表现如何。创建一个简单的评估脚本:
from llava.model.builder import load_pretrained_model
from llava.mm_utils import get_model_name_from_path
from llava.eval.run_llava import eval_model
from PIL import Image
import os
# 加载微调后的模型
model_path = "./llava-finetuned"
tokenizer, model, image_processor, context_len = load_pretrained_model(
model_path=model_path,
model_base=None,
model_name=get_model_name_from_path(model_path)
)
def predict_category(image_path):
# 准备输入
args = type('Args', (), {
"model_path": model_path,
"model_base": None,
"model_name": get_model_name_from_path(model_path),
"query": "<image>\n请问这张图片属于哪个类别?",
"conv_mode": None,
"image_file": image_path,
"sep": ",",
"temperature": 0,
"top_p": None,
"num_beams": 1,
"max_new_tokens": 50
})()
# 获取预测结果
result = eval_model(args)
return result
# 测试一张图片
test_image = "test_dog.jpg"
prediction = predict_category(test_image)
print(f"模型预测: {prediction}")
5.2 性能指标
除了直观感受,我们还需要量化指标:
- 准确率:在所有测试样本中预测正确的比例
- 混淆矩阵:看看模型容易混淆哪些类别
- 推理速度:单张图片的处理时间
你可以在验证集上运行批量测试,计算这些指标。
6. 实际应用与部署
6.1 模型导出
训练好的模型可以导出为更易部署的格式:
python llava/model/apply_delta.py \
--base /path/to/vicuna-7b \
--target ./llava-deploy \
--delta liuhaotian/llava-v1.6-vicuna-7b
6.2 简单推理API
创建一个简单的Flask应用来提供分类服务:
from flask import Flask, request, jsonify
from PIL import Image
import io
from llava.model.builder import load_pretrained_model
from llava.mm_utils import process_images, tokenizer_image_token
import torch
app = Flask(__name__)
# 加载模型
model_path = "./llava-deploy"
tokenizer, model, image_processor, context_len = load_pretrained_model(model_path)
@app.route('/classify', methods=['POST'])
def classify():
if 'image' not in request.files:
return jsonify({'error': 'No image provided'}), 400
image_file = request.files['image']
image = Image.open(io.BytesIO(image_file.read()))
# 预处理图像
image_tensor = process_images([image], image_processor, model.config)
# 准备输入
input_ids = tokenizer_image_token(
"<image>\n请问这张图片属于哪个类别?",
tokenizer,
return_tensors='pt'
)
# 推理
with torch.no_grad():
output = model.generate(
input_ids,
images=image_tensor,
max_new_tokens=50,
use_cache=True
)
# 解码结果
response = tokenizer.decode(output[0], skip_special_tokens=True)
return jsonify({'category': response})
if __name__ == '__main__':
app.run(host='0.0.0.0', port=5000)
这样你就有了一个可以接收图片并返回分类结果的API服务。
7. 常见问题与解决
在微调过程中可能会遇到一些问题,这里列举几个常见的:
问题1:显存不足 解决方法:减小batch size,使用梯度累积,或者启用梯度检查点
问题2:过拟合 解决方法:增加数据量,使用数据增强,减少训练轮数,或者增加dropout
问题3:模型不收敛 解决方法:检查学习率是否合适,确认数据标注是否正确
问题4:推理速度慢 解决方法:使用量化技术,或者考虑模型蒸馏到更小的模型
8. 总结
通过这篇教程,你应该已经掌握了Llava-v1.6-7b模型微调的基本流程。从环境准备、数据预处理,到训练配置和模型评估,我们一步步走完了整个流程。
实际用下来,Llava的微调相对还是比较简单的,特别是用对话格式来处理分类任务,比传统的分类头方法更灵活。你不仅可以用它做简单分类,还可以让模型输出更详细的描述。
如果你刚开始接触模型微调,建议先从小的数据集开始,熟悉整个流程后再扩展到更大的任务。遇到问题也不用担心,多试几次就能掌握窍门。微调后的模型在你的特定领域应该会有不错的表现,毕竟它是专门为你的任务优化的。
获取更多AI镜像
想探索更多AI镜像和应用场景?访问 CSDN星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。
更多推荐
所有评论(0)