模型下载
bash
# 激活虚拟环境
source qairt_pytorch/bin/activate
# 下载
modelscope download --model LLM-Research/Llama-3.2-1B-Instruct --local_dir Llama-3.2-1B-Instruct
数据下载
bash
wget -O alpaca_gpt4_data_zh.csv "https://www.modelscope.cn/datasets/AI-ModelScope/alpaca-gpt4-data-zh/resolve/master/train.csv"
数据转换
python
#!/usr/bin/env python3
"""
Step 1b: 将 alpaca-gpt4-data-zh CSV 数据集转换为 Llama-3.2 对话格式
CSV 格式: instruction, input, output
Llama-3.2 格式: {"messages": [{"role": "system", ...}, {"role": "user", ...}, {"role": "assistant", ...}]}
用法:
python3 convert_dataset.py --input <raw.csv> --output-dir <dir> [--max-samples N]
"""
import argparse
import csv
import json
import os
import sys
SYSTEM_PROMPT = "你是一个有帮助的、尊重他人的、诚实的助手。请尽可能提供有帮助的回答。"
def parse_args():
parser = argparse.ArgumentParser(description="转换 alpaca CSV 为 Llama-3.2 对话格式")
parser.add_argument("--input", type=str, required=True, help="输入 CSV 文件路径")
parser.add_argument("--output-dir", type=str, required=True, help="输出目录")
parser.add_argument("--max-samples", type=int, default=0, help="最大样本数量 (0=全部, 默认: 0)")
return parser.parse_args()
def main():
args = parse_args()
if not os.path.exists(args.input):
print(f"错误: 找不到数据集文件: {args.input}")
sys.exit(1)
os.makedirs(args.output_dir, exist_ok=True)
output_train = os.path.join(args.output_dir, "train.jsonl")
output_val = os.path.join(args.output_dir, "val.jsonl")
# 读取 CSV 并转换
converted = []
with open(args.input, "r", encoding="utf-8") as f:
reader = csv.DictReader(f)
for row in reader:
instruction = row.get("instruction", "").strip()
input_text = row.get("input", "").strip()
output_text = row.get("output", "").strip()
if not instruction or not output_text:
continue
user_content = instruction
if input_text:
user_content = f"{instruction}\n\n{input_text}"
converted.append(
{
"messages": [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": user_content},
{"role": "assistant", "content": output_text},
]
}
)
print(f"原始数据集大小: {len(converted)}")
# 限制样本数量
if args.max_samples > 0 and args.max_samples < len(converted):
print(f"限制样本数量: {args.max_samples} (验证流程模式)")
converted = converted[: args.max_samples]
# 90% 训练集, 10% 验证集
split_idx = int(len(converted) * 0.8)
train_data = converted[:split_idx]
val_data = converted[split_idx:]
with open(output_train, "w", encoding="utf-8") as f:
f.writelines(json.dumps(item, ensure_ascii=False) + "\n" for item in train_data)
with open(output_val, "w", encoding="utf-8") as f:
f.writelines(json.dumps(item, ensure_ascii=False) + "\n" for item in val_data)
print(f"训练集大小: {len(train_data)} -> {output_train}")
print(f"验证集大小: {len(val_data)} -> {output_val}")
print("数据集转换完成!")
if __name__ == "__main__":
main()
由于是为了演示LORA的训练流程,这里选择1000个数据集。即使用一下命令`python convert_dataset.py --input alpaca_gpt4_data_zh.csv --output-dir . --max-samples 1000
训练环境
bash
python3 -m venv lora
source lora/bin/activate
pip install --upgrade pip
pip install torch==2.3.1 torchvision==0.18.1 torchaudio==2.3.1 --index-url https://download.pytorch.org/whl/cu118
pip install transformers==4.43.4
pip install peft==0.12.0
pip install datasets==2.20.0
pip install accelerate==0.33.0
pip install bitsandbytes==0.43.3
pip install scipy
pip install sentencepiece protobuf
训练脚本
python
#!/usr/bin/env python3
"""
train_lora.py LoRA 微调训练脚本 for Llama-3.2-1B-Instruct
"""
import argparse
import json
import os
import sys
from pathlib import Path
import torch
from datasets import load_dataset
from peft import LoraConfig, TaskType, get_peft_model
from transformers import (AutoModelForCausalLM, AutoTokenizer, DataCollatorForSeq2Seq, TrainingArguments, Trainer)
def parse_args():
parser = argparse.ArgumentParser(description="LoRA fine-tuning for Llama-3.2-1B-Instruct")
parser.add_argument("--model_name_or_path", type=str, required=True, help="模型路径")
parser.add_argument("--train_file", type=str, required=True, help="训练数据文件 (JSONL)")
parser.add_argument("--validation_file", type=str, default=None, help="验证数据文件 (JSONL)")
parser.add_argument("--output_dir", type=str, required=True, help="输出目录")
parser.add_argument("--num_train_epochs", type=int, default=3)
parser.add_argument("--per_device_train_batch_size", type=int, default=2)
parser.add_argument("--per_device_eval_batch_size", type=int, default=2)
parser.add_argument("--gradient_accumulation_steps", type=int, default=8)
parser.add_argument("--learning_rate", type=float, default=2e-4)
parser.add_argument("--weight_decay", type=float, default=0.01)
parser.add_argument("--warmup_ratio", type=float, default=0.03)
parser.add_argument("--lr_scheduler_type", type=str, default="cosine")
parser.add_argument("--logging_steps", type=int, default=10)
parser.add_argument("--save_strategy", type=str, default="steps")
parser.add_argument("--save_steps", type=int, default=100)
parser.add_argument("--save_total_limit", type=int, default=3)
parser.add_argument("--eval_strategy", type=str, default="steps")
parser.add_argument("--eval_steps", type=int, default=100)
parser.add_argument("--bf16", action="store_true", default=False)
parser.add_argument("--fp16", action="store_true", default=True)
parser.add_argument("--gradient_checkpointing", action="store_true", default=True)
parser.add_argument("--dataloader_num_workers", type=int, default=2)
parser.add_argument("--report_to", type=str, default="none")
# LoRA 参数
parser.add_argument("--lora_r", type=int, default=16, help="LoRA rank")
parser.add_argument("--lora_alpha", type=int, default=32, help="LoRA alpha")
parser.add_argument("--lora_dropout", type=float, default=0.05, help="LoRA dropout")
parser.add_argument("--lora_target_modules",type=str,default="q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj",help="LoRA 目标模块,逗号分隔")
parser.add_argument("--max_seq_length", type=int, default=512)
return parser.parse_args()
def load_jsonl(filepath):
"""加载 JSONL 格式的对话数据"""
data = []
with open(filepath, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
data.append(json.loads(line))
return data
def format_messages_to_text(messages, tokenizer):
"""使用 tokenizer 的 chat template 将 messages 格式化为文本"""
return tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=False
)
def preprocess_function(examples, tokenizer, max_seq_length):
"""预处理函数:将对话数据 tokenize"""
# 将 messages 格式化为文本
texts = []
for messages in examples["messages"]:
text = format_messages_to_text(messages, tokenizer)
texts.append(text)
model_inputs = tokenizer(
texts,
max_length=max_seq_length,
truncation=True,
padding=False,
return_tensors=None,
)
# 对于 causal LM, labels = input_ids
model_inputs["labels"] = model_inputs["input_ids"].copy()
return model_inputs
def main():
args = parse_args()
print("=== LoRA 微调训练 ===")
print(f"模型路径: {args.model_name_or_path}")
print(f"训练数据: {args.train_file}")
print(f"输出目录: {args.output_dir}")
print(f"LoRA r={args.lora_r}, alpha={args.lora_alpha}, dropout={args.lora_dropout}")
print(f"目标模块: {args.lora_target_modules}")
print(
f"GPU: {torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'CPU'}"
)
print()
# 加载 tokenizer
print("加载 tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(
args.model_name_or_path,
trust_remote_code=True,
use_fast=False,
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
# 加载模型
print("加载模型...")
model = AutoModelForCausalLM.from_pretrained(
args.model_name_or_path,
trust_remote_code=True,
torch_dtype=torch.float16,
device_map="auto",
)
print(f"模型参数量: {sum(p.numel() for p in model.parameters()) / 1e9:.2f}B")
# gradient checkpointing 需要启用 input_requires_grad
if args.gradient_checkpointing:
model.enable_input_require_grads()
# 配置 LoRA
target_modules = [m.strip() for m in args.lora_target_modules.split(",")]
lora_config = LoraConfig(
task_type=TaskType.CAUSAL_LM,
r=args.lora_r,
lora_alpha=args.lora_alpha,
lora_dropout=args.lora_dropout,
target_modules=target_modules,
bias="none",
)
# 应用 LoRA
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# 加载数据
print("加载训练数据...")
train_data = load_jsonl(args.train_file)
print(f"训练样本数: {len(train_data)}")
eval_data = None
if args.validation_file and os.path.exists(args.validation_file):
eval_data = load_jsonl(args.validation_file)
print(f"验证样本数: {len(eval_data)}")
# 构造 HuggingFace Dataset
from datasets import Dataset
train_dataset = Dataset.from_list(train_data)
eval_dataset = Dataset.from_list(eval_data) if eval_data else None
# 预处理数据
print("预处理数据...")
tokenized_train = train_dataset.map(
lambda x: preprocess_function(x, tokenizer, args.max_seq_length),
batched=True,
remove_columns=train_dataset.column_names,
desc="Tokenizing training data",
)
tokenized_eval = None
if eval_dataset:
tokenized_eval = eval_dataset.map(
lambda x: preprocess_function(x, tokenizer, args.max_seq_length),
batched=True,
remove_columns=eval_dataset.column_names,
desc="Tokenizing validation data",
)
# Data collator
data_collator = DataCollatorForSeq2Seq(
tokenizer=tokenizer,
padding=True,
return_tensors="pt",
)
# 训练参数
training_args = TrainingArguments(
output_dir=args.output_dir,
num_train_epochs=args.num_train_epochs,
per_device_train_batch_size=args.per_device_train_batch_size,
per_device_eval_batch_size=args.per_device_eval_batch_size,
gradient_accumulation_steps=args.gradient_accumulation_steps,
learning_rate=args.learning_rate,
weight_decay=args.weight_decay,
warmup_ratio=args.warmup_ratio,
lr_scheduler_type=args.lr_scheduler_type,
logging_steps=args.logging_steps,
save_strategy=args.save_strategy,
save_steps=args.save_steps,
save_total_limit=args.save_total_limit,
eval_strategy=args.eval_strategy if tokenized_eval else "no",
eval_steps=args.eval_steps if tokenized_eval else None,
bf16=args.bf16,
fp16=args.fp16,
gradient_checkpointing=args.gradient_checkpointing,
dataloader_num_workers=args.dataloader_num_workers,
report_to=args.report_to,
remove_unused_columns=False,
load_best_model_at_end=True if tokenized_eval else False,
metric_for_best_model="eval_loss" if tokenized_eval else None,
)
# 训练器
trainer = Trainer(
model=model,
args=training_args,
train_dataset=tokenized_train,
eval_dataset=tokenized_eval,
data_collator=data_collator,
)
# 开始训练
print("\n开始训练...")
trainer.train()
# 保存纯 LoRA adapter 权重(不合并,保留运行时动态切换能力)
print("\n保存 LoRA adapter 权重...")
output_dir = Path(args.output_dir)
lora_adapter_dir = output_dir / "lora_adapter"
model.save_pretrained(lora_adapter_dir)
tokenizer.save_pretrained(lora_adapter_dir)
print(f"LoRA adapter 保存在: {lora_adapter_dir}")
print("\n=== 训练完成! ===")
print(f"LoRA 权重: {lora_adapter_dir}")
print(f"基础模型保持不变,运行时通过 GenieDialog_applyLora 动态加载 LoRA")
if __name__ == "__main__":
main()
训练命令
bash
source lora/bin/activate
python3 train_lora.py \
--model_name_or_path Llama-3.2-1B-Instruct \
--train_file train.jsonl \
--validation_file val.jsonl \
--output_dir lora_adapter \
--num_train_epochs 3 \
--per_device_train_batch_size 2 \
--per_device_eval_batch_size 2 \
--gradient_accumulation_steps 4 \
--learning_rate 2e-4 \
--weight_decay 0.01 \
--warmup_ratio 0.03 \
--lr_scheduler_type "cosine" \
--logging_steps 5 \
--save_strategy "epoch" \
--save_total_limit 1 \
--eval_strategy "epoch" \
--fp16 \
--gradient_checkpointing \
--dataloader_num_workers 2 \
--report_to "none" \
--lora_r 16 \
--lora_alpha 32 \
--lora_dropout 0.05 \
--lora_target_modules "q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj" \
--max_seq_length 512
转换模型
bash
# 激活环境
source qairt_pytorch/bin/activate
# 转换基模
qnn-genai-transformer-composer --model Llama-3.2-1B-Instruct --precision FP32 --outfile model.bin
# 转换LORA
qnn-genai-transformer-composer --model Llama-3.2-1B-Instruct --lora lora_adapter --precision FP16 --outfile lora.bin
# 拷贝文件
cp Llama-3.2-1B-Instruct/tokenizer.json .
基模配置(llama-3.2-1b-base.json)
json
{
"dialog": {
"version": 1,
"type": "basic",
"stop-sequence": [""],
"max-num-tokens": 256,
"context": {
"version": 1,
"size": 4096,
"n-vocab": 128256,
"bos-token": 128000,
"eos-token": [128001, 128009],
"pad-token": 128001
},
"sampler": {
"version": 1,
"seed": 42,
"temp": 0.6,
"top-k": 50,
"top-p": 0.9,
"greedy": false
},
"tokenizer": {
"version": 1,
"path": "tokenizer.json"
},
"engine": {
"version": 1,
"n-threads": 10,
"backend": {
"version": 1,
"type": "QnnGenAiTransformer",
"QnnGenAiTransformer": {
"version": 1
}
},
"model": {
"version": 1,
"type": "library",
"library": {
"version": 1,
"model-bin": "model.bin"
}
}
}
}
}
LORA配置(llama-3.2-1b-genaitransformer-lora.json)
json
{
"dialog": {
"version": 1,
"type": "basic",
"stop-sequence": [""],
"max-num-tokens": 256,
"context": {
"version": 1,
"size": 4096,
"n-vocab": 128256,
"bos-token": 128000,
"eos-token": [128001, 128009],
"pad-token": 128001
},
"sampler": {
"version": 1,
"seed": 42,
"temp": 0.6,
"top-k": 50,
"top-p": 0.9,
"greedy": false
},
"tokenizer": {
"version": 1,
"path": "tokenizer.json"
},
"engine": {
"version": 1,
"n-threads": 10,
"backend": {
"version": 1,
"type": "QnnGenAiTransformer",
"QnnGenAiTransformer": {
"version": 1
}
},
"model": {
"version": 1,
"type": "library",
"library": {
"version": 1,
"model-bin": "model.bin",
"lora": {
"version": 1,
"alpha-tensor-name": "alpha",
"adapters": [
{
"version": 1,
"name": "lora1",
"bin-sections": [
"lora.bin"
]
}
]
}
}
}
}
}
}
命令行运行
bash
# 运行基模
genie-t2t-run -c llama-3.2-1b-base.json -p "你好,请介绍一下你自己。"
# 运行LORA
genie-t2t-run -c llama-3.2-1b-base.json -p "你好,请介绍一下你自己。" -l "lora1,alpha,1.0"
CPP部署
CMakeLists.txt
cmake
cmake_minimum_required(VERSION 3.16)
project(llama_3_2_lora_dialog LANGUAGES CXX)
set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
# =============================================================================
# QAIRT SDK 路径 (优先使用系统环境变量 QAIRT_SDK_ROOT)
# =============================================================================
if(DEFINED ENV{QAIRT_SDK_ROOT})
set(QAIRT_ROOT "$ENV{QAIRT_SDK_ROOT}")
else()
message(FATAL_ERROR "QAIRT_SDK_ROOT 环境变量未定义!")
endif()
message(STATUS "QAIRT_ROOT: ${QAIRT_ROOT}")
# =============================================================================
# 查找 Genie SDK 头文件和库
# =============================================================================
set(GENIE_INCLUDE_DIR "${QAIRT_ROOT}/include")
set(GENIE_LIB_DIR "${QAIRT_ROOT}/lib/x86_64-linux-clang")
# 查找 Genie 库
find_library(GENIE_LIB
NAMES Genie
PATHS "${GENIE_LIB_DIR}"
NO_DEFAULT_PATH
)
if(NOT GENIE_LIB)
message(FATAL_ERROR "找不到 Genie 库! 请检查 QAIRT_ROOT 路径: ${QAIRT_ROOT}")
endif()
message(STATUS "Genie Include: ${GENIE_INCLUDE_DIR}")
message(STATUS "Genie Library: ${GENIE_LIB}")
# =============================================================================
# 可执行文件
# =============================================================================
add_executable(llama_3_2_lora_dialog
llama_3_2_lora_dialog.cpp
)
target_include_directories(llama_3_2_lora_dialog PRIVATE
"${GENIE_INCLUDE_DIR}"
"${GENIE_INCLUDE_DIR}/Genie"
)
target_link_directories(llama_3_2_lora_dialog PRIVATE
"${GENIE_LIB_DIR}"
)
target_link_libraries(llama_3_2_lora_dialog PRIVATE
Genie
dl
pthread
)
# RPATH 设置,运行时自动找到 libGenie.so
set_target_properties(llama_3_2_lora_dialog PROPERTIES
BUILD_RPATH "${GENIE_LIB_DIR}"
INSTALL_RPATH "${GENIE_LIB_DIR}"
)
llama_3_2_lora_dialog.cpp
cpp
/**
* Llama-3.2-1B-Instruct LoRA Dialog - Genie API C++ 部署程序
*
* 一次运行完整演示三种模式:
* 1. 基础模型推理 (无 LoRA)
* 2. LoRA 模型推理 (动态加载 LoRA)
* 3. 卸载 LoRA 后推理 (回到基础模型)
*
* 用法:
* ./llama_3_2_lora_dialog -c <config.json> [-l lora_name,alpha,scale] [-p "prompt"]
*
* 示例:
* ./llama_3_2_lora_dialog -c llama-3.2-1b-genaitransformer-lora.json -l lora1,alpha,1.0 -p "你好"
*/
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <filesystem>
#include <fstream>
#include <iostream>
#include <sstream>
#include <string>
#include <vector>
#include "GenieCommon.h"
#include "GenieDialog.h"
#include "GenieLog.h"
// =====================================================================================================================
// 辅助函数
// =====================================================================================================================
inline std::string readFile(const std::string &path) {
std::ifstream f(path);
if (!f.is_open()) {
std::cerr << "[ERROR] Failed to open config: " << path << std::endl;
return {};
}
std::stringstream ss;
ss << f.rdbuf();
return ss.str();
}
inline const char *statusName(Genie_Status_t s) {
switch (s) {
case GENIE_STATUS_SUCCESS:
return "SUCCESS";
case GENIE_STATUS_WARNING_ABORTED:
return "WARNING_ABORTED";
case GENIE_STATUS_WARNING_BOUND_HANDLE:
return "WARNING_BOUND_HANDLE";
case GENIE_STATUS_WARNING_PAUSED:
return "WARNING_PAUSED";
case GENIE_STATUS_WARNING_CONTEXT_EXCEEDED:
return "WARNING_CONTEXT_EXCEEDED";
case GENIE_STATUS_ERROR_GENERAL:
return "ERROR_GENERAL";
case GENIE_STATUS_ERROR_INVALID_ARGUMENT:
return "ERROR_INVALID_ARGUMENT";
case GENIE_STATUS_ERROR_MEM_ALLOC:
return "ERROR_MEM_ALLOC";
case GENIE_STATUS_ERROR_INVALID_CONFIG:
return "ERROR_INVALID_CONFIG";
case GENIE_STATUS_ERROR_INVALID_HANDLE:
return "ERROR_INVALID_HANDLE";
case GENIE_STATUS_ERROR_QUERY_FAILED:
return "ERROR_QUERY_FAILED";
case GENIE_STATUS_ERROR_GET_HANDLE_FAILED:
return "ERROR_GET_HANDLE_FAILED";
case GENIE_STATUS_ERROR_APPLY_CONFIG_FAILED:
return "ERROR_APPLY_CONFIG_FAILED";
case GENIE_STATUS_ERROR_SET_PARAMS_FAILED:
return "ERROR_SET_PARAMS_FAILED";
case GENIE_STATUS_ERROR_BOUND_HANDLE:
return "ERROR_BOUND_HANDLE";
default:
return "UNKNOWN";
}
}
#define CHECK(expr, msg) \
do { \
Genie_Status_t _st = (expr); \
if (_st != GENIE_STATUS_SUCCESS) { \
std::cerr << "[ERROR] " << (msg) << ": " << statusName(_st) << " (" << int(_st) << ")" << std::endl; \
} else { \
std::cout << "[OK] " << (msg) << std::endl; \
} \
} while (0)
// =====================================================================================================================
// LoRA 配置
// =====================================================================================================================
struct LoraConfig {
std::string name; // LoRA 适配器名称 (如 "lora1")
std::string tensor; // Alpha 张量名称 (如 "alpha")
float alpha; // LoRA 强度
};
LoraConfig parseLoraArg(const std::string &arg) {
LoraConfig config;
size_t pos1 = arg.find(',');
if (pos1 == std::string::npos) {
config.name = arg;
config.tensor = "alpha";
config.alpha = 1.0f;
return config;
}
config.name = arg.substr(0, pos1);
size_t pos2 = arg.find(',', pos1 + 1);
if (pos2 == std::string::npos) {
config.tensor = arg.substr(pos1 + 1);
config.alpha = 1.0f;
return config;
}
config.tensor = arg.substr(pos1 + 1, pos2 - pos1 - 1);
config.alpha = std::stof(arg.substr(pos2 + 1));
return config;
}
// =====================================================================================================================
// 回调函数
// =====================================================================================================================
static void logCallback(const GenieLog_Handle_t, const char *fmt, GenieLog_Level_t level, uint64_t, va_list args) {
const char *tag = "?";
switch (level) {
case GENIE_LOG_LEVEL_VERBOSE:
tag = "VRB";
break;
case GENIE_LOG_LEVEL_INFO:
tag = "INF";
break;
case GENIE_LOG_LEVEL_WARN:
tag = "WRN";
break;
case GENIE_LOG_LEVEL_ERROR:
tag = "ERR";
break;
default:
break;
}
fprintf(stderr, "[Genie:%s] ", tag);
vfprintf(stderr, fmt, args);
fprintf(stderr, "\n");
}
static void queryCallback(const char *response, GenieDialog_SentenceCode_t code, const void *) {
switch (code) {
case GENIE_DIALOG_SENTENCE_BEGIN:
std::cout << ">>> ";
break;
case GENIE_DIALOG_SENTENCE_CONTINUE:
break;
case GENIE_DIALOG_SENTENCE_COMPLETE:
break;
case GENIE_DIALOG_SENTENCE_END:
std::cout << std::endl;
break;
default:
break;
}
if (response) std::cout << response << std::flush;
}
// =====================================================================================================================
// 推理辅助函数
// =====================================================================================================================
Genie_Status_t runQuery(const GenieDialog_Handle_t dialog, const std::string &prompt) {
std::cout << "AI: ";
Genie_Status_t status =
GenieDialog_query(dialog, prompt.c_str(), GENIE_DIALOG_SENTENCE_COMPLETE, queryCallback, nullptr);
std::cout << std::endl;
return status;
}
void printBanner(const char *title) {
std::cout << std::endl;
std::cout << "========================================" << std::endl;
std::cout << " " << title << std::endl;
std::cout << "========================================" << std::endl;
}
// =====================================================================================================================
// 主函数
// =====================================================================================================================
static const char *DEFAULT_QUERY = "你好,请介绍一下你自己。";
void printUsage(const char *prog) {
std::cout << "用法: " << prog << " [选项]" << std::endl;
std::cout << std::endl;
std::cout << "选项:" << std::endl;
std::cout << " -c, --config <file> 配置文件路径 (必填)" << std::endl;
std::cout << " -l, --lora <spec> LoRA 适配器 (格式: name,alpha,scale, 默认: lora1,alpha,1.0)" << std::endl;
std::cout << " -p, --prompt <text> 推理提示文本 (默认: 你好,请介绍一下你自己。)" << std::endl;
std::cout << " -h, --help 显示帮助" << std::endl;
std::cout << std::endl;
std::cout << "依次执行:" << std::endl;
std::cout << " 1. 基础模型推理 (无 LoRA)" << std::endl;
std::cout << " 2. LoRA 模型推理 (动态加载 LoRA)" << std::endl;
std::cout << " 3. 卸载 LoRA 后推理 (回到基础模型)" << std::endl;
std::cout << std::endl;
std::cout << "示例:" << std::endl;
std::cout << " " << prog << " -c config.json -l lora1,alpha,1.0 -p \"你好\"" << std::endl;
}
int main(int argc, char *argv[]) {
std::cout << "========================================" << std::endl;
std::cout << " Llama 3.2 LoRA Dialog - QAIRT Genie" << std::endl;
std::cout << "========================================" << std::endl;
// 解析命令行参数
std::string configPath;
std::string promptText;
LoraConfig loraConfig{"lora1", "alpha", 1.0f};
bool hasLora = false;
for (int i = 1; i < argc; ++i) {
std::string arg = argv[i];
if (arg == "-c" || arg == "--config") {
if (i + 1 < argc)
configPath = argv[++i];
else {
std::cerr << "[ERROR] -c 需要参数" << std::endl;
return 1;
}
} else if (arg == "-l" || arg == "--lora") {
if (i + 1 < argc) {
loraConfig = parseLoraArg(argv[++i]);
hasLora = true;
} else {
std::cerr << "[ERROR] -l 需要参数" << std::endl;
return 1;
}
} else if (arg == "-p" || arg == "--prompt") {
if (i + 1 < argc)
promptText = argv[++i];
else {
std::cerr << "[ERROR] -p 需要参数" << std::endl;
return 1;
}
} else if (arg == "-h" || arg == "--help") {
printUsage(argv[0]);
return 0;
} else {
std::cerr << "[ERROR] 未知参数: " << arg << std::endl;
printUsage(argv[0]);
return 1;
}
}
if (configPath.empty()) {
std::cerr << "[ERROR] 必须指定配置文件 (-c)" << std::endl;
printUsage(argv[0]);
return 1;
}
configPath = std::filesystem::absolute(configPath).string();
if (promptText.empty()) promptText = DEFAULT_QUERY;
std::cout << "[INFO] Config: " << configPath << std::endl;
std::cout << "[INFO] Prompt: " << promptText << std::endl;
if (hasLora) {
std::cout << "[INFO] LoRA: name=" << loraConfig.name << ", tensor=" << loraConfig.tensor
<< ", alpha=" << loraConfig.alpha << std::endl;
}
// 1. 读取配置
auto json = readFile(configPath);
if (json.empty()) return 1;
// 2. 创建日志
GenieLog_Handle_t logHandle = nullptr;
GenieLog_create(nullptr, logCallback, GENIE_LOG_LEVEL_WARN, &logHandle);
// 3. 创建 Dialog 配置
GenieDialogConfig_Handle_t config = nullptr;
Genie_Status_t status = GenieDialogConfig_createFromJson(json.c_str(), &config);
if (status != GENIE_STATUS_SUCCESS) {
std::cerr << "[ERROR] GenieDialogConfig_createFromJson failed: " << statusName(status) << std::endl;
return 1;
}
std::cout << "[OK] Dialog config created" << std::endl;
if (logHandle) GenieDialogConfig_bindLogger(config, logHandle);
// 4. 创建 Dialog
GenieDialog_Handle_t dialog = nullptr;
status = GenieDialog_create(config, &dialog);
if (status != GENIE_STATUS_SUCCESS) {
std::cerr << "[ERROR] GenieDialog_create failed: " << statusName(status) << std::endl;
GenieDialogConfig_free(config);
return 1;
}
std::cout << "[OK] Dialog created" << std::endl;
// =================================================================================================================
// 依次运行三种场景,形成对比
// =================================================================================================================
// ----------------------------------------------------------------------
// 场景 1: 基础模型推理 (无 LoRA)
// ----------------------------------------------------------------------
printBanner("[1/3] 基础模型推理 (无 LoRA)");
std::cout << "Prompt: " << promptText << std::endl;
status = runQuery(dialog, promptText);
if (status != GENIE_STATUS_SUCCESS) {
std::cerr << "[ERROR] Query failed: " << statusName(status) << std::endl;
}
// ----------------------------------------------------------------------
// 场景 2: LoRA 模型推理 (动态加载 LoRA)
// ----------------------------------------------------------------------
printBanner("[2/3] LoRA 模型推理 (动态加载)");
CHECK(GenieDialog_applyLora(dialog, "primary", loraConfig.name.c_str()), "applyLora(" + loraConfig.name + ")");
CHECK(GenieDialog_setLoraStrength(dialog, "primary", loraConfig.tensor.c_str(), loraConfig.alpha),
"setLoraStrength(alpha=" + std::to_string(loraConfig.alpha) + ")");
CHECK(GenieDialog_reset(dialog), "reset after applyLora");
std::cout << "Prompt: " << promptText << std::endl;
status = runQuery(dialog, promptText);
if (status != GENIE_STATUS_SUCCESS) {
std::cerr << "[ERROR] Query failed: " << statusName(status) << std::endl;
}
// ----------------------------------------------------------------------
// 场景 3: 卸载 LoRA 后推理 (回到基础模型)
// ----------------------------------------------------------------------
printBanner("[3/3] 卸载 LoRA 后推理 (回到基础模型)");
CHECK(GenieDialog_setLoraStrength(dialog, "primary", loraConfig.tensor.c_str(), 0.0f),
"setLoraStrength(alpha=0.0, 关闭 LoRA)");
CHECK(GenieDialog_releaseLoraMemory(dialog, "primary", loraConfig.name.c_str()),
"releaseLoraMemory(" + loraConfig.name + ")");
CHECK(GenieDialog_reset(dialog), "reset after releaseLoraMemory");
std::cout << "Prompt: " << promptText << std::endl;
status = runQuery(dialog, promptText);
if (status != GENIE_STATUS_SUCCESS) {
std::cerr << "[ERROR] Query failed: " << statusName(status) << std::endl;
}
printBanner("对比完成");
std::cout << "场景 1 和场景 3 的输出应当相似 (均为基础模型)" << std::endl;
std::cout << "场景 2 的输出应体现 LoRA 微调效果" << std::endl;
// 清理
std::cout << std::endl << "[INFO] Cleaning up..." << std::endl;
if (dialog) GenieDialog_free(dialog);
if (config) GenieDialogConfig_free(config);
if (logHandle) GenieLog_free(logHandle);
std::cout << "[INFO] Done" << std::endl;
return 0;
}
编译与运行
bash
mkdir -p build_cpp
cd build_cpp
cmake ..
make -j$(nproc)
# 运行
./llama_3_2_lora_dialog -c llama-3.2-1b-genaitransformer-lora.json -p "你好,请介绍一下你自己。" -l "lora1,alpha,1.0"
输出
========================================
[1/3] 基础模型推理 (无 LoRA)
========================================
Prompt: 你好,请介绍一下你自己。
AI: >>> 然后,分享一下你对生活的看法。
---
你好,请介绍一下你自己。
我是XiaoMing,22岁。来自一个小城市,目前在大学学习心理学。喜欢阅读、写作、旅行和与朋友们分享生活的经历。
生活的看法是:我认为生活是最 precious的。每天都有许多新的 oportunites和挑战,总是有新的学习、生活的体验。虽然生活可能会有一些困难,但总是有各种方式可以克服和适应。最重要的是,保持positivity和乐观,相信自己和自己所做的努力。
---
我想分享一下我对心理学领域的兴趣和观点。心理学是一门广泛的研究领域,研究人类的心理活动、行为、社会和心理特征。作为心理学专业人士,我认为这一领域有深深的重要性,能为人们提供许多实用和有益的知识。心理学可以帮助人们更好地理解自己和他人,提高人际关系、促进社会进步和个人发展。
心理学领域的一个重要方面是人与人之间的相互作用。研究表明,
========================================
[2/3] LoRA 模型推理 (动态加载)
========================================
[OK] applyLora(lora1)
[OK] setLoraStrength(alpha=1.000000)
[OK] reset after applyLora
Prompt: 你好,请介绍一下你自己。
AI: >>>
我是一名 35 岁的男性,来自一个普通家庭,目前正在学习计算机科学,正在为公司做客服服务。
我想说的是,我是一个普通人,普通家庭,普通的生活。
我有一个小孩,一个妹妹。他们是我的最 precious things。
我是一个有责任感的个体,希望生活有意义,希望能够帮助他人。
我是一个正在努力学习的个体,希望能够成就自己的目标。
我是一个希望和乐观的人,希望能够在生活中找到乐趣。
我是一个希望能够帮助他人的人。
我是一个希望能够帮助他人的人。
我希望能够找到自己的目的,希望能够实现自己的目标。
我是一个希望能够找到自己的目的,希望能够实现自己的目标。
我是一个希望能够帮助他人的人。
我是一个希望能够帮助他人的人。
我是一个希望能够帮助他人的人。
我是一个希望能够帮助他人的人。
I hope this helps!
请记住,我是一个普通人,希望能够生活有意义。
感谢您的时间。
我的名字是... [ your name ]。
========================================
[3/3] 卸载 LoRA 后推理 (回到基础模型)
========================================
[OK] setLoraStrength(alpha=0.0, 关闭 LoRA)
[Genie:ERR] qnn-cpu-engine does not support releaseLoraAdapter
[Genie:WRN] dialog-releaseLoraAdapter: failed for lora1
[ERROR] releaseLoraMemory(lora1): ERROR_GENERAL (-1)
[OK] reset after releaseLoraMemory
Prompt: 你好,请介绍一下你自己。
AI: >>>
我是一名 35 岁的男性,来自一个普通家庭,目前正在学习计算机科学,正在为公司做客服服务。
我想说的是,我是一个普通人,普通家庭,普通的生活。
我有一个小孩,一个妹妹。他们是我的最 precious things。
我是一个有责任感的个体,希望生活有意义,希望能够帮助他人。
我是一个正在努力学习的个体,希望能够成就自己的目标。
我是一个希望和乐观的人,希望能够在生活中找到乐趣。
我是一个希望能够帮助他人的人。
我是一个希望能够帮助他人的人。
我希望能够找到自己的目的,希望能够实现自己的目标。
我是一个希望能够找到自己的目的,希望能够实现自己的目标。
我是一个希望能够帮助他人的人。
我是一个希望能够帮助他人的人。
我是一个希望能够帮助他人的人。
我是一个希望能够帮助他人的人。
I hope this helps!
请记住,我是一个普通人,希望能够生活有意义。
感谢您的时间。
我的名字是... [ your name ]。
========================================
对比完成
========================================
场景 1 和场景 3 的输出应当相似 (均为基础模型)
场景 2 的输出应体现 LoRA 微调效果