【DNN】Llama 3.2 1B模型LORA微调与QAIRT部署

模型下载

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 微调效果
相关推荐
Leslie1651 小时前
模型调价之后,如何把 Token 账单变成可回滚的预算策略
人工智能
企鹅的企1 小时前
2027北京AI健康科技与智慧医疗展官方链接产业资源
人工智能·科技
独隅1 小时前
KMP 全栈进化:Koog 框架打造纯 Kotlin AI Agent 实战效果
开发语言·人工智能·kotlin
腾视科技-AIoT1 小时前
私有云时代来临:AI NAS如何重塑你的数字生活?
人工智能·ai·生活·nas·ai算力模组·ainas·腾视科技
fthux1 小时前
装闭 RenoPit 源码解析(12):从AI分析结果到React避坑报告
人工智能·ai·开源·github·open source·renopit
老余说AI1 小时前
TikTok Shop东南亚上线“内容授权工具“,搬运内容可合法化
人工智能
o_insist1 小时前
从 Vue 生命周期与 Spring AOP 理解 LangChain Middleware
人工智能·agent
o_insist1 小时前
AI Agent 如何动态选择工具:Skill 匹配与三种筛选模式
人工智能·agent
l1258651 小时前
# RAG重排序实战:硅基流动bge-reranker-v2-m3在线API vs 本地CrossEncoder,一篇讲透两种方案
数据库·人工智能·python·深度学习·算法·机器学习·langchain