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 实际数据准备示例

假设你在做一个动物分类任务,包含猫、狗、鸟三个类别。你需要:

  1. 收集每个类别至少100-200张图片
  2. 按照上述格式创建标注文件
  3. 将数据集分为训练集和验证集(建议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星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐