Skip to main content

模型微调:基于Llama-Factory

1. 演练目标与架构

本节基于摩尔线程MTT S5000 GPU 完成 Llama-Factory 的微调,覆盖从容器部署到模型产出的全链路操作。

2. 环境准备

机器信息

在MTT S5000 上进行 Llama-Factory 微调。MUSA driver version 为 4.3.4。

使用摩尔线程预训练容器 musa-train:4.3.4_kuae2.1_20260106_alinux,容器内置 MUSA SDK 4.3.4、MT Megatron-LM 0.14.0、MT DeepSpeed 0.17.2、Torch-Musa 2.7.0。

创建并启动微调容器

# 创建容器
sudo docker create \
--privileged \
--env MTHREADS_VISIBLE_DEVICES=all \
--net host \
-v /data/:/data/ \
--name llamafactory \
registry.mthreads.com/public/musa-train:4.3.4_kuae2.1_20260106_alinux \
sleep infinity

# 启动并进入容器
sudo docker start llamafactory
sudo docker exec -it llamafactory bash

# 在容器里,启动ssh服务
service ssh restart

使用容器:musa-train:4.3.4_kuae2.1_20260106_alinux

此容器的 transformers 版本是 4.49,为了支持好 Qwen3 系列模型,需要升级到 >=4.51 版本。

# 升级transformers版本
pip install transformers==4.51.3

GPU 识别验证

进入容器后,验证MTT S5000 GPU 的识别状态与拓扑连接:

# 查看GPU硬件识别与状态
mthreads-gmi

# 验证MCCL通信能力(需已经安装mccl-test,进入目录)
# ./mccl_test

3. 安装 Llama-Factory

在容器里,下载 Llama-Factory 0.9.3,进行 musify,然后安装。

git clone https://github.com/hiyouga/LlamaFactory.git
cd LlamaFactory
git checkout v0.9.3
cd ..
bash /data/playground/tool/musify.sh LlamaFactory
cd LlamaFactory_MUSA
pip install -e .

4. 运行 Llama-Factory WebUI

打开终端执行命令,即可运行 Llama-Factory webui。

# 运行 webui
llamafactory-cli webui

# webui运行在7860端口
Running on local URL: 0.0.0.0: 7860

因为有些场景下无法访问webui,本文下面用命令行操作。

5. 下载模型、数据

下载模型

可以使用 modelscope 下载模型到本地:

# 安装modelscope
pip install modelscope
# 下载模型到本地
modelscope download --model Qwen/Qwen3-4B --local_dir /data/models/Qwen3-4B
modelscope download --model Qwen/Qwen2-0.5B-Instruct --local_dir /data/models/Qwen2-0.5B-Instruct

下载数据

下载示例数据,这个数据集中包含多轮对话格式的训练/验证数据,适合快速测试。

cd Llama-Factory
wget https://atp-modelzoo-sh.oss-cn-shanghai.aliyuncs.com/release/tutorials/llama_factory/data.zip
mv data rawdata && unzip data.zip -d data

6. SFT 微调 Qwen3-0.6B

Qwen3-0.6B 模型比较小,完全可以在单卡MTT S5000 上做全参数 SFT 微调。下面是 SFT 微调脚本,采用 BF16 精度。

MUSA_VISIBLE_DEVICES=0 llamafactory-cli train \
--stage sft \
--do_train True \
--model_name_or_path /data/models/Qwen3-0.6B \
--preprocessing_num_workers 16 \
--finetuning_type full \
--template qwen3 \
--flash_attn auto \
--dataset_dir /data/data \
--dataset train \
--cutoff_len 2048 \
--learning_rate 5e-05 \
--num_train_epochs 3.0 \
--max_samples 100000 \
--per_device_train_batch_size 2 \
--gradient_accumulation_steps 8 \
--lr_scheduler_type cosine \
--max_grad_norm 1.0 \
--logging_steps 5 \
--save_steps 100 \
--warmup_steps 0 \
--packing False \
--enable_thinking True \
--report_to none \
--output_dir saves/Qwen3-0.6B-Thinking/full/train_2025-09-30-09-15-00 \
--bf16 True \
--plot_loss True \
--trust_remote_code True \
--ddp_timeout 180000000 \
--include_num_input_tokens_seen True \
--optim adamw_torch

微调过程很快,训练完成后,Llama-Factory 会显示训练信息:

***** train metrics *****
epoch = 2.0
num_input_tokens_seen = 510192
total_flos = 1255737GF
train_loss = 2.2287
train_runtime = 0:00:28.84
train_samples_per_second = 33.278
train_steps_per_second = 0.208

7. LoRA 微调 Qwen3-32B

Qwen3-32B 在 8 卡MTT S5000 上的微调。

llamafactory-cli train \
--stage sft \
--do_train True \
--model_name_or_path /data/models/Qwen3-32B \
--preprocessing_num_workers 16 \
--finetuning_type lora \
--lora_rank 16 \
--lora_alpha 32 \
--lora_dropout 0.05 \
--lora_target all \
--template qwen3 \
--flash_attn auto \
--dataset_dir /data/data \
--dataset train \
--cutoff_len 2048 \
--learning_rate 5e-05 \
--num_train_epochs 3.0 \
--max_samples 100000 \
--per_device_train_batch_size 1 \
--gradient_accumulation_steps 16 \
--lr_scheduler_type cosine \
--max_grad_norm 1.0 \
--logging_steps 5 \
--save_steps 100 \
--warmup_steps 0 \
--packing False \
--enable_thinking True \
--report_to none \
--output_dir saves/Qwen3-32B-Thinking/lora/train_$(date +%Y%m%d-%H%M%S) \
--bf16 True \
--plot_loss True \
--trust_remote_code True \
--ddp_timeout 180000000 \
--include_num_input_tokens_seen True \
--optim adamw_torch

微调过程大约 3 分钟。训练完成后,Llama-Factory 会显示训练信息:

***** train metrics *****
epoch = 2.0
num_input_tokens_seen = 389392
total_flos = 69886406GF
train_loss = 1.7674
train_runtime = 0:02:24.90
train_samples_per_second = 6.625
train_steps_per_second = 0.041