大模型的监督微调(Supervised Fine-Tuning, SFT)

SWIFT 的全称是 Scalable lightWeight Infrastructure for Fine-Tuning(可扩展的轻量级微调基础设施),提供从模型微调到最终部署的一整套工具,让大模型的定制和落地变得简单、高效。

一、Swift 框架简介

MS-SWIFT 是阿里云 ModelScope 社区开源的大模型微调框架,核心特点是:

  • 支持 600+ 大语言模型和 300+ 多模态模型

  • 内置 LoRA、QLoRA 等参数高效微调方法

  • 单卡 24GB 显存即可微调 7B 级别模型

  • 覆盖训练→评估→量化→部署全流程

二、环境安装

python 复制代码
# 创建虚拟环境
conda create -n swift python=3.10
conda activate swift

# 安装 PyTorch
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118

# 安装 ms-swift
pip install ms-swift -U

# 可选:安装加速依赖
pip install deepspeed flash-attn --no-build-isolation

三、数据格式准备

Swift 支持多种数据格式,推荐使用 JSONL 格式的 messages 结构。我这里使用了一个小数据样本。

1、JSON与JSONAL区别

特性 JSON JSONL
结构 一个完整的对象(通常是数组) 每行一个独立对象
文件扩展名 .json .jsonl.jsonlines
读取方式 一次性加载整个文件到内存 逐行读取,一次处理一条
内存占用 大文件会占用大量内存 内存友好,只加载当前行
编辑便利性 需要解析整个结构才能修改 可以直接追加或修改某一行
流式处理 不支持 支持(边读边处理)
断点续传 困难 容易(记录已处理行数)

四、核心训练命令

根据数据集和路径自行修改代码。

python 复制代码
swift sft \
    --model autodl-tmp/modelscope_cache/models/Qwen/Qwen2.5-VL-7B-Instruct\
    --adapters autodl-tmp/output/continued_training/v2-20260614-155943/checkpoint-816 \
    --dataset /root/autodl-tmp/BUSCoT/DatasetFiles/train_fixed.jsonl\
    --quant_bits 4 \
    --lora_rank 16 \
    --lora_alpha 32 \
    --tuner_type lora \
    --gradient_checkpointing true \
    --per_device_train_batch_size 2 \
    --learning_rate 1e-4\
    --output_dir /root/autodl-tmp/output/continued_training

五、训练过程

六、推理测试

1、生成测试集结果文件

python 复制代码
swift infer \
    --adapters autodl-tmp/output/continued_training/v4-20260614-181229/checkpoint-816 \
    --val_dataset /root/autodl-tmp/BUSCoT/DatasetFiles/test.jsonl \
    --infer_backend pt \
    --max_batch_size 4 \
    --max_new_tokens 512 \
    --result_path ./test_results.jsonl

2、评估推理结果准确率

python 复制代码
import json
import re

def calculate_accuracy(result_file):
    correct = 0
    total = 0
    
    with open(result_file, 'r') as f:
        for line in f:
            data = json.loads(line)
            
            # 模型预测的输出
            predicted = data.get('response', '')
            # 真实标签(需要从测试集中获取,或者结果文件中已包含)
            ground_truth = data.get('ground_truth', '')
            
            # 如果结果文件中没有 ground_truth,可以从 predict 结构中提取
            # 有些格式下,真实值在 labels 字段中
            if not ground_truth:
                ground_truth = data.get('labels', '')
            
            # 提取 answer 标签中的数字
            pred_match = re.search(r'<answer>\s*(\d+)\s*</answer>', predicted)
            true_match = re.search(r'<answer>\s*(\d+)\s*</answer>', ground_truth)
            
            if pred_match and true_match:
                if pred_match.group(1) == true_match.group(1):
                    correct += 1
                total += 1
            elif pred_match and not true_match:
                # 如果真实值没有 answer 标签,尝试直接匹配数字
                total += 1
                # 这里可以根据实际情况处理
    
    if total > 0:
        print(f"准确率: {correct}/{total} = {correct/total*100:.2f}%")
    else:
        print("没有找到有效的评估样本")

calculate_accuracy('test_results.jsonl')

3、评估结果

结果还有待提高,谢谢关注!!!

相关推荐
搜yyzk6682 小时前
智能电话机器人:1天千通电话,高效低成本
人工智能
fthux6 小时前
装闭 RenoPit 源码解析(09):AnalysisEngine装修闭坑分析主流程
人工智能·ai·开源·github·open source·renopit
敏编程6 小时前
开源项目 NextBand:探索融合健康传感、语音交互与 AI Agent 的下一代可穿戴设备
人工智能
ovO6 小时前
我按下“发送”之后:DeepSeek Harness 如何把一次请求变成 1,177 个增量片段
人工智能·开源·deepseek
昨日之日20067 小时前
LTX-2.5:更清晰、更可控、同步音画、多镜头连贯的AI视频神器
人工智能·音视频
九硕智慧建筑一体化厂家7 小时前
数字化管控落地,直流照明构筑安全照明体系
运维·人工智能·笔记·安全·智慧城市
江厌017 小时前
本地部署DeepSeek显卡怎么选?我把显存计算公式和实测记录写下来了
人工智能·百度·联想工作站·代理商推荐·渠道对比·工作站
小席是个热心肠7 小时前
AI相关的自我学习
java·人工智能·学习
FII工业富联科技服务7 小时前
工业AI自治的演进趋势:从内容生成到Agent驱动的工厂运营范式转变
人工智能
云和数据.ChenGuang8 小时前
fastapi的参数剖析
人工智能·深度学习·机器学习·语言模型·状态模式·fastapi