mirror of
https://github.com/DrHo1y/ezrknn-llm.git
synced 2026-10-01 15:46:38 +07:00
release v1.1.0
This commit is contained in:
19
CHANGELOG.md
19
CHANGELOG.md
@@ -1,4 +1,17 @@
|
||||
# CHANGELOG
|
||||
## v1.1.0
|
||||
- Support group-wise quantization (w4a16 group sizes of 32/64/128, w8a8 group sizes of 128/256/512).
|
||||
- Support joint inference with LoRA model loading
|
||||
- Support storage and preloading of prompt cache.
|
||||
- Support gguf model conversion (currently only support q4_0 and fp16).
|
||||
- Optimize initialization, prefill, and decode time.
|
||||
- Support four input types: prompt, embedding, token, and multimodal.
|
||||
- Add PC-based simulation accuracy testing and inference interface support for rkllm-toolkit.
|
||||
- Add gdq algorithm to improve 4-bit quantization accuracy.
|
||||
- Add mixed quantization algorithm, supporting a combination of grouped and non-grouped quantization based on specified ratios.
|
||||
- Add support for models such as Llama3, Gemma2, and MiniCPM3.
|
||||
- Resolve catastrophic forgetting issue when the number of tokens exceeds max_context.
|
||||
|
||||
## v1.0.1
|
||||
- Optimize model conversion memory occupation
|
||||
- Optimize inference memory occupation
|
||||
@@ -11,7 +24,7 @@
|
||||
- Add logprob and token_id to the return value
|
||||
|
||||
## v1.0.0
|
||||
- Supports the conversion and deployment of LLM models on RK3588/RK3576 platforms
|
||||
- Support the conversion and deployment of LLM models on RK3588/RK3576 platforms
|
||||
- Compatible with Hugging Face model architectures
|
||||
- Currently supports the models Llama, Qwen, Qwen2, and Phi-2
|
||||
- Supports quantization with w8a8 and w4a16 precision
|
||||
- Currently support the models Llama, Qwen, Qwen2, and Phi-2
|
||||
- Support quantization with w8a8 and w4a16 precision
|
||||
45
README.md
Normal file → Executable file
45
README.md
Normal file → Executable file
@@ -18,18 +18,21 @@
|
||||
- RK3576 Series
|
||||
|
||||
# Support Models
|
||||
- [X] [TinyLLAMA 1.1B](https://huggingface.co/TinyLlama/TinyLlama-1.1B-Chat-v1.0/tree/fe8a4ea1ffedaf415f4da2f062534de366a451e6)
|
||||
- [X] [Qwen 1.8B](https://huggingface.co/Qwen/Qwen-1_8B-Chat/tree/1d0f68de57b88cfde81f3c3e537f24464d889081)
|
||||
- [X] [Qwen2 0.5B](https://huggingface.co/Qwen/Qwen1.5-0.5B/tree/8f445e3628f3500ee69f24e1303c9f10f5342a39)
|
||||
- [X] [Phi-2 2.7B](https://hf-mirror.com/microsoft/phi-2/tree/834565c23f9b28b96ccbeabe614dd906b6db551a)
|
||||
- [X] [Phi-3 3.8B](https://huggingface.co/microsoft/Phi-3-mini-4k-instruct/tree/291e9e30e38030c23497afa30f3af1f104837aa6)
|
||||
- [X] [ChatGLM3 6B](https://huggingface.co/THUDM/chatglm3-6b/tree/103caa40027ebfd8450289ca2f278eac4ff26405)
|
||||
- [X] [Gemma 2B](https://huggingface.co/google/gemma-2b-it/tree/de144fb2268dee1066f515465df532c05e699d48)
|
||||
- [X] [InternLM2 1.8B](https://huggingface.co/internlm/internlm2-chat-1_8b/tree/ecccbb5c87079ad84e5788baa55dd6e21a9c614d)
|
||||
- [X] [MiniCPM 2B](https://huggingface.co/openbmb/MiniCPM-2B-sft-bf16/tree/79fbb1db171e6d8bf77cdb0a94076a43003abd9e)
|
||||
- [X] [LLAMA models](https://huggingface.co/meta-llama)
|
||||
- [X] [TinyLLAMA models](https://huggingface.co/TinyLlama)
|
||||
- [X] [Qwen models](https://huggingface.co/models?search=Qwen/Qwen)
|
||||
- [X] [Phi models](https://huggingface.co/models?search=microsoft/phi)
|
||||
- [X] [ChatGLM3-6B](https://huggingface.co/THUDM/chatglm3-6b/tree/103caa40027ebfd8450289ca2f278eac4ff26405)
|
||||
- [X] [Gemma models](https://huggingface.co/collections/google/gemma-2-release-667d6600fd5220e7b967f315)
|
||||
- [X] [InternLM2 models](https://huggingface.co/collections/internlm/internlm2-65b0ce04970888799707893c)
|
||||
- [X] [MiniCPM models](https://huggingface.co/collections/openbmb/minicpm-65d48bf958302b9fd25b698f)
|
||||
|
||||
# Download
|
||||
- You can also download all packages, docker image, examples, docs and platform-tools from [RKLLM_SDK](https://console.zbox.filez.com/l/RJJDmB), fetch code: rkllm
|
||||
You can download the latest package, docker image, example, documentation, and platform-tool from [RKLLM_SDK](https://console.zbox.filez.com/l/RJJDmB), fetch code: rkllm
|
||||
|
||||
# Note
|
||||
|
||||
The modifications in version 1.1.0 are significant, making it incompatible with older version models. Please use the latest toolchain for model conversion and inference.
|
||||
|
||||
# RKNN Toolkit2
|
||||
If you want to deploy additional AI model, we have introduced a SDK called RKNN-Toolkit2. For details, please refer to:
|
||||
@@ -37,15 +40,17 @@ If you want to deploy additional AI model, we have introduced a SDK called RKNN-
|
||||
https://github.com/airockchip/rknn-toolkit2
|
||||
|
||||
# CHANGELOG
|
||||
## v1.0.1
|
||||
- Optimize model conversion memory occupation
|
||||
- Optimize inference memory occupation
|
||||
- Increase prefill speed
|
||||
- Reduce initialization time
|
||||
- Improve quantization accuracy
|
||||
- Add support for Gemma, ChatGLM3, MiniCPM, InternLM2, and Phi-3
|
||||
- Add Server invocation
|
||||
- Add inference interruption interface
|
||||
- Add logprob and token_id to the return value
|
||||
## v1.1.0
|
||||
- Support group-wise quantization (w4a16 group sizes of 32/64/128, w8a8 group sizes of 128/256/512).
|
||||
- Support joint inference with LoRA model loading
|
||||
- Support storage and preloading of prompt cache.
|
||||
- Support gguf model conversion (currently only support q4_0 and fp16).
|
||||
- Optimize initialization, prefill, and decode time.
|
||||
- Support four input types: prompt, embedding, token, and multimodal.
|
||||
- Add PC-based simulation accuracy testing and inference interface support for rkllm-toolkit.
|
||||
- Add gdq algorithm to improve 4-bit quantization accuracy.
|
||||
- Add mixed quantization algorithm, supporting a combination of grouped and non-grouped quantization based on specified ratios.
|
||||
- Add support for models such as Llama3, Gemma2, and MiniCPM3.
|
||||
- Resolve catastrophic forgetting issue when the number of tokens exceeds max_context.
|
||||
|
||||
for older version, please refer [CHANGELOG](CHANGELOG.md)
|
||||
Binary file not shown.
BIN
doc/Rockchip_RKLLM_SDK_CN_1.1.0.pdf
Executable file
BIN
doc/Rockchip_RKLLM_SDK_CN_1.1.0.pdf
Executable file
Binary file not shown.
Binary file not shown.
BIN
doc/Rockchip_RKLLM_SDK_EN_1.1.0.pdf
Executable file
BIN
doc/Rockchip_RKLLM_SDK_EN_1.1.0.pdf
Executable file
Binary file not shown.
@@ -1,21 +1,26 @@
|
||||
cmake_minimum_required(VERSION 3.10)
|
||||
project(llm_demo)
|
||||
project(rkllm_demo)
|
||||
|
||||
set(CMAKE_CXX_STANDARD 11)
|
||||
set(CMAKE_CXX_STANDARD_REQUIRED ON)
|
||||
|
||||
set(SOURCE_FILES src/main.cpp)
|
||||
set(SOURCE_FILES_1 src/llm_demo.cpp)
|
||||
add_executable(llm_demo ${SOURCE_FILES_1})
|
||||
|
||||
add_executable(${PROJECT_NAME} ${SOURCE_FILES})
|
||||
set(SOURCE_FILES_2 src/multimodel_demo.cpp)
|
||||
add_executable(multimodel_demo ${SOURCE_FILES_2})
|
||||
|
||||
set(RKLLM_API_PATH "${CMAKE_SOURCE_DIR}/../../runtime/${CMAKE_SYSTEM_NAME}/librkllm_api")
|
||||
include_directories(${RKLLM_API_PATH}/include)
|
||||
if(CMAKE_SYSTEM_NAME STREQUAL "Android")
|
||||
set(RKLLM_RT_LIB ${RKLLM_API_PATH}/${CMAKE_ANDROID_ARCH_ABI}/librkllmrt.so)
|
||||
target_link_libraries(${PROJECT_NAME} ${RKLLM_RT_LIB} log)
|
||||
find_package(OpenMP REQUIRED)
|
||||
target_link_libraries(llm_demo ${RKLLM_RT_LIB} log OpenMP::OpenMP_CXX)
|
||||
target_link_libraries(multimodel_demo ${RKLLM_RT_LIB} log OpenMP::OpenMP_CXX)
|
||||
elseif(CMAKE_SYSTEM_NAME STREQUAL "Linux")
|
||||
set(RKLLM_RT_LIB ${RKLLM_API_PATH}/aarch64/librkllmrt.so)
|
||||
target_link_libraries(${PROJECT_NAME} ${RKLLM_RT_LIB})
|
||||
target_link_libraries(llm_demo ${RKLLM_RT_LIB})
|
||||
target_link_libraries(multimodel_demo ${RKLLM_RT_LIB})
|
||||
endif()
|
||||
|
||||
|
||||
|
||||
194
rkllm-runtime/examples/rkllm_api_demo/src/llm_demo.cpp
Normal file
194
rkllm-runtime/examples/rkllm_api_demo/src/llm_demo.cpp
Normal file
File diff suppressed because one or more lines are too long
@@ -1,128 +0,0 @@
|
||||
// Copyright (c) 2024 by Rockchip Electronics Co., Ltd. All Rights Reserved.
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#include <string.h>
|
||||
#include <unistd.h>
|
||||
#include <string>
|
||||
#include "rkllm.h"
|
||||
#include <fstream>
|
||||
#include <iostream>
|
||||
#include <csignal>
|
||||
#include <vector>
|
||||
|
||||
#define PROMPT_TEXT_PREFIX "<|im_start|>system You are a helpful assistant. <|im_end|> <|im_start|>user"
|
||||
#define PROMPT_TEXT_POSTFIX "<|im_end|><|im_start|>assistant"
|
||||
|
||||
using namespace std;
|
||||
LLMHandle llmHandle = nullptr;
|
||||
|
||||
void exit_handler(int signal)
|
||||
{
|
||||
if (llmHandle != nullptr)
|
||||
{
|
||||
{
|
||||
cout << "程序即将退出" << endl;
|
||||
LLMHandle _tmp = llmHandle;
|
||||
llmHandle = nullptr;
|
||||
rkllm_destroy(_tmp);
|
||||
}
|
||||
exit(signal);
|
||||
}
|
||||
}
|
||||
|
||||
void callback(RKLLMResult *result, void *userdata, LLMCallState state)
|
||||
{
|
||||
if (state == LLM_RUN_FINISH)
|
||||
{
|
||||
printf("\n");
|
||||
}
|
||||
else if (state == LLM_RUN_ERROR)
|
||||
{
|
||||
printf("\\run error\n");
|
||||
}
|
||||
else
|
||||
{
|
||||
printf("%s", result->text);
|
||||
}
|
||||
}
|
||||
|
||||
int main(int argc, char **argv)
|
||||
{
|
||||
if (argc != 2)
|
||||
{
|
||||
printf("Usage:%s [rkllm_model_path]\n", argv[0]);
|
||||
return -1;
|
||||
}
|
||||
signal(SIGINT, exit_handler);
|
||||
string rkllm_model(argv[1]);
|
||||
printf("rkllm init start\n");
|
||||
|
||||
//设置参数及初始化
|
||||
RKLLMParam param = rkllm_createDefaultParam();
|
||||
param.model_path = rkllm_model.c_str();
|
||||
param.num_npu_core = 2;
|
||||
param.top_k = 1;
|
||||
param.max_new_tokens = 256;
|
||||
param.max_context_len = 512;
|
||||
param.logprobs = false;
|
||||
param.top_logprobs = 5;
|
||||
param.use_gpu = false;
|
||||
rkllm_init(&llmHandle, param, callback);
|
||||
printf("rkllm init success\n");
|
||||
|
||||
vector<string> pre_input;
|
||||
pre_input.push_back("把下面的现代文翻译成文言文:到了春风和煦,阳光明媚的时候,湖面平静,没有惊涛骇浪,天色湖光相连,一片碧绿,广阔无际;沙洲上的鸥鸟,时而飞翔,时而停歇,美丽的鱼游来游去,岸上与小洲上的花草,青翠欲滴。");
|
||||
pre_input.push_back("以咏梅为题目,帮我写一首古诗,要求包含梅花、白雪等元素。");
|
||||
pre_input.push_back("上联: 江边惯看千帆过");
|
||||
pre_input.push_back("把这句话翻译成中文:Knowledge can be acquired from many sources. These include books, teachers and practical experience, and each has its own advantages. The knowledge we gain from books and formal education enables us to learn about things that we have no opportunity to experience in daily life. We can also develop our analytical skills and learn how to view and interpret the world around us in different ways. Furthermore, we can learn from the past by reading books. In this way, we won't repeat the mistakes of others and can build on their achievements.");
|
||||
pre_input.push_back("把这句话翻译成英文:RK3588是新一代高端处理器,具有高算力、低功耗、超强多媒体、丰富数据接口等特点");
|
||||
cout << "\n**********************可输入以下问题对应序号获取回答/或自定义输入********************\n"
|
||||
<< endl;
|
||||
for (int i = 0; i < (int)pre_input.size(); i++)
|
||||
{
|
||||
cout << "[" << i << "] " << pre_input[i] << endl;
|
||||
}
|
||||
cout << "\n*************************************************************************\n"
|
||||
<< endl;
|
||||
|
||||
string text;
|
||||
while (true)
|
||||
{
|
||||
std::string input_str;
|
||||
printf("\n");
|
||||
printf("user: ");
|
||||
std::getline(std::cin, input_str);
|
||||
if (input_str == "exit")
|
||||
{
|
||||
break;
|
||||
}
|
||||
for (int i = 0; i < (int)pre_input.size(); i++)
|
||||
{
|
||||
if (input_str == to_string(i))
|
||||
{
|
||||
input_str = pre_input[i];
|
||||
cout << input_str << endl;
|
||||
}
|
||||
}
|
||||
// string text = PROMPT_TEXT_PREFIX + input_str + PROMPT_TEXT_POSTFIX;
|
||||
string text = input_str;
|
||||
|
||||
printf("robot: ");
|
||||
rkllm_run(llmHandle, text.c_str(), NULL);
|
||||
}
|
||||
|
||||
rkllm_destroy(llmHandle);
|
||||
|
||||
return 0;
|
||||
}
|
||||
191
rkllm-runtime/examples/rkllm_api_demo/src/multimodel_demo.cpp
Normal file
191
rkllm-runtime/examples/rkllm_api_demo/src/multimodel_demo.cpp
Normal file
@@ -0,0 +1,191 @@
|
||||
// Copyright (c) 2024 by Rockchip Electronics Co., Ltd. All Rights Reserved.
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#include <string.h>
|
||||
#include <unistd.h>
|
||||
#include <string>
|
||||
#include "rkllm.h"
|
||||
#include <fstream>
|
||||
#include <iostream>
|
||||
#include <csignal>
|
||||
#include <vector>
|
||||
|
||||
#define PROMPT_TEXT_PREFIX "<用户>"
|
||||
#define PROMPT_TEXT_POSTFIX "<AI>"
|
||||
|
||||
|
||||
using namespace std;
|
||||
LLMHandle llmHandle = nullptr;
|
||||
|
||||
void exit_handler(int signal)
|
||||
{
|
||||
if (llmHandle != nullptr)
|
||||
{
|
||||
{
|
||||
cout << "程序即将退出" << endl;
|
||||
LLMHandle _tmp = llmHandle;
|
||||
llmHandle = nullptr;
|
||||
rkllm_destroy(_tmp);
|
||||
}
|
||||
}
|
||||
exit(signal);
|
||||
}
|
||||
|
||||
void callback(RKLLMResult *result, void *userdata, LLMCallState state)
|
||||
{
|
||||
if (state == RKLLM_RUN_FINISH)
|
||||
{
|
||||
printf("\n");
|
||||
} else if (state == RKLLM_RUN_ERROR) {
|
||||
printf("\\run error\n");
|
||||
} else if (state == RKLLM_RUN_GET_LAST_HIDDEN_LAYER) {
|
||||
/* ================================================================================================================
|
||||
若使用GET_LAST_HIDDEN_LAYER功能,callback接口会回传内存指针:last_hidden_layer,token数量:num_tokens与隐藏层大小:embd_size
|
||||
通过这三个参数可以取得last_hidden_layer中的数据
|
||||
注:需要在当前callback中获取,若未及时获取,下一次callback会将该指针释放
|
||||
===============================================================================================================*/
|
||||
if (result->last_hidden_layer.embd_size != 0 && result->last_hidden_layer.num_tokens != 0) {
|
||||
int data_size = result->last_hidden_layer.embd_size * result->last_hidden_layer.num_tokens * sizeof(float);
|
||||
printf("\ndata_size:%d",data_size);
|
||||
std::ofstream outFile("last_hidden_layer.bin", std::ios::binary);
|
||||
if (outFile.is_open()) {
|
||||
outFile.write(reinterpret_cast<const char*>(result->last_hidden_layer.hidden_states), data_size);
|
||||
outFile.close();
|
||||
std::cout << "Data saved to output.bin successfully!" << std::endl;
|
||||
} else {
|
||||
std::cerr << "Failed to open the file for writing!" << std::endl;
|
||||
}
|
||||
}
|
||||
} else if (state == RKLLM_RUN_NORMAL) {
|
||||
printf("%s", result->text);
|
||||
// for(int i=0; i<result->num; i++)
|
||||
// {
|
||||
// printf("%d token_id: %d logprob: %f\n", i, result->tokens[i].id, result->tokens[i].logprob);
|
||||
// }
|
||||
}
|
||||
}
|
||||
|
||||
int main(int argc, char **argv)
|
||||
{
|
||||
if (argc < 4) {
|
||||
std::cerr << "Usage: " << argv[0] << " model_path max_new_tokens max_context_len\n";
|
||||
return 1;
|
||||
}
|
||||
|
||||
signal(SIGINT, exit_handler);
|
||||
printf("rkllm init start\n");
|
||||
|
||||
//设置参数及初始化
|
||||
RKLLMParam param = rkllm_createDefaultParam();
|
||||
param.model_path = argv[1];
|
||||
param.top_k = 1;
|
||||
param.max_new_tokens = std::atoi(argv[2]);
|
||||
param.max_context_len = std::atoi(argv[3]);
|
||||
|
||||
// if use multimodel mode, need to set img_start,img_end and img_content
|
||||
param.img_start = "<image>";
|
||||
param.img_end = "</image>\n";
|
||||
param.img_content = "<unk>";
|
||||
|
||||
param.skip_special_token = true;
|
||||
|
||||
int ret = rkllm_init(&llmHandle, ¶m, callback);
|
||||
if (ret == 0){
|
||||
printf("rkllm init success\n");
|
||||
} else {
|
||||
printf("rkllm init failed\n");
|
||||
exit_handler(-1);
|
||||
}
|
||||
|
||||
vector<string> pre_input;
|
||||
pre_input.push_back("<image>What is in the image?");
|
||||
cout << "\n**********************可输入以下问题对应序号获取回答/或自定义输入********************\n"
|
||||
<< endl;
|
||||
for (int i = 0; i < (int)pre_input.size(); i++)
|
||||
{
|
||||
cout << "[" << i << "] " << pre_input[i] << endl;
|
||||
}
|
||||
cout << "\n*************************************************************************\n"
|
||||
<< endl;
|
||||
|
||||
string text;
|
||||
RKLLMInput rkllm_input;
|
||||
|
||||
// 初始化 infer 参数结构体
|
||||
RKLLMInferParam rkllm_infer_params;
|
||||
memset(&rkllm_infer_params, 0, sizeof(RKLLMInferParam)); // 将所有内容初始化为 0
|
||||
|
||||
// // 1. 初始化并设置 LoRA 参数(如果需要使用 LoRA)
|
||||
// RKLLMLoraParam lora_params;
|
||||
// lora_params.lora_adapter_name = "test"; // 指定用于推理的 lora 名称
|
||||
|
||||
// // 2. 初始化并设置 Prompt Cache 参数(如果需要使用 prompt cache)
|
||||
// RKLLMPromptCacheParam prompt_cache_params;
|
||||
// prompt_cache_params.save_prompt_cache = true; // 是否保存 prompt cache
|
||||
// prompt_cache_params.prompt_cache_path = "./prompt_cache.bin"; // 若需要保存prompt cache, 指定 cache 文件路径
|
||||
|
||||
rkllm_infer_params.mode = RKLLM_INFER_GENERATE;
|
||||
// rkllm_infer_params.lora_params = &lora_params;
|
||||
// rkllm_infer_params.prompt_cache_params = &prompt_cache_params;
|
||||
|
||||
// rkllm_load_prompt_cache(llmHandle, "./prompt_cache.bin");
|
||||
while (true)
|
||||
{
|
||||
std::string input_str;
|
||||
printf("\n");
|
||||
printf("user: ");
|
||||
std::getline(std::cin, input_str);
|
||||
if (input_str == "exit")
|
||||
{
|
||||
break;
|
||||
}
|
||||
for (int i = 0; i < (int)pre_input.size(); i++)
|
||||
{
|
||||
if (input_str == to_string(i))
|
||||
{
|
||||
input_str = pre_input[i];
|
||||
cout << input_str << endl;
|
||||
}
|
||||
}
|
||||
|
||||
float * img_embed_data_ptr = (float *)malloc(64 * 2304 * sizeof(float));
|
||||
// 打开输入文件流(以二进制模式)
|
||||
std::ifstream inFile("./img_vec.bin", std::ios::binary);
|
||||
// 检查文件是否成功打开
|
||||
if (!inFile) {
|
||||
printf("Failed to open file for reading: ");
|
||||
}
|
||||
// 读取数据
|
||||
float temp_data;
|
||||
int idx = 0;
|
||||
while (inFile.read(reinterpret_cast<char*>(&temp_data), sizeof(float))){
|
||||
img_embed_data_ptr[idx] = temp_data;
|
||||
idx = idx + 1;
|
||||
}
|
||||
// 关闭文件
|
||||
inFile.close();
|
||||
|
||||
text = PROMPT_TEXT_PREFIX + input_str + PROMPT_TEXT_POSTFIX;
|
||||
rkllm_input.input_type = RKLLM_INPUT_MULTIMODAL;
|
||||
rkllm_input.multimodal_input.prompt = (char *)text.c_str();
|
||||
rkllm_input.multimodal_input.image_embed = img_embed_data_ptr;
|
||||
rkllm_input.multimodal_input.n_image_tokens = 64;
|
||||
printf("robot: ");
|
||||
rkllm_run(llmHandle, &rkllm_input, &rkllm_infer_params, NULL);
|
||||
free(img_embed_data_ptr);
|
||||
}
|
||||
rkllm_destroy(llmHandle);
|
||||
|
||||
return 0;
|
||||
}
|
||||
@@ -8,8 +8,8 @@ Before running the demo, you need to prepare the following files:
|
||||
### Build
|
||||
You can run the demo with the only command:
|
||||
```bash
|
||||
# ./build_rkllm_server_flask.sh [target_platform:rk3588/rk3576] [RKLLM-Server workshop] [transformed_rkllm_model_path in borad]
|
||||
./build_rkllm_server_flask.sh rk3588 /user/data/rkllm_server /user/data/rkllm_server/model.rkllm
|
||||
# Usage: ./build_rkllm_server_flask.sh --workshop [RKLLM-Server Working Path] --model_path [Absolute Path of Converted RKLLM Model on Board] --platform [Target Platform: rk3588/rk3576] --npu_num [NPU Core Count] [--lora_model_path [Lora Model Path]] [--prompt_cache_path [Prompt Cache File Path]]
|
||||
./build_rkllm_server_flask.sh --workshop /user/data --model_path /user/data/model.rkllm --platform rk3588 --npu_num 3
|
||||
```
|
||||
### Access with API
|
||||
After building the RKLLM-Server-Flask, You can use ‘chat_api_flask.py’ to access the RKLLM-Server-Flask and get the answser of RKLLM models.
|
||||
@@ -20,8 +20,8 @@ Attention: you should check the IP address of the board with 'ifconfig' command
|
||||
### Build
|
||||
You can run the demo with the only command:
|
||||
```bash
|
||||
# ./build_rkllm_server_gradio.sh [target_platform:rk3588/rk3576] [RKLLM-Server workshop] [transformed_rkllm_model_path in borad]
|
||||
./build_rkllm_server_gradio.sh rk3588 /user/data/rkllm_server /user/data/rkllm_server/model.rkllm
|
||||
# Usage: ./build_rkllm_server_gradio.sh --workshop [RKLLM-Server Working Path] --model_path [Absolute Path of Converted RKLLM Model on Board] --platform [Target Platform: rk3588/rk3576] --npu_num [NPU Core Count] [--lora_model_path [Lora Model Path]] [--prompt_cache_path [Prompt Cache File Path]]
|
||||
./build_rkllm_server_gradio.sh --workshop /user/data --model_path /user/data/model.rkllm --platform rk3588 --npu_num 3
|
||||
```
|
||||
### Access the Server
|
||||
After running the demo, You can access the RKLLM-Server-Gradio with two ways:
|
||||
|
||||
@@ -1,61 +1,110 @@
|
||||
#!/bin/bash
|
||||
|
||||
#*****************************************************************************************#
|
||||
# 该脚本为 RKLLM-Server-Flask 服务的一键设置脚本
|
||||
# 用户可以运行该脚本实现Linux板端的 RKLLM-Server-Flask 服务的自动化部署。
|
||||
# 使用说明: ./build_rkllm_server_flask.sh [目标平台:rk3588/rk3576] [RKLLM-Server工作路径] [已转换的rkllm模型在板端的绝对路径]
|
||||
# example: ./build_rkllm_server_flask.sh rk3588 /user/data/rkllm_server /user/data/rkllm_server/model.rkllm
|
||||
# This script is an automated setup script for the RKLLM-Server-Flask service.
|
||||
# Users can run this script to automate the deployment of the RKLLM-Server-Flask service on a Linux board.
|
||||
# Usage: ./build_rkllm_server_flask.sh --workshop [RKLLM-Server Working Path] --model_path [Absolute Path of Converted RKLLM Model on Board] --platform [Target Platform: rk3588/rk3576] --npu_num [NPU Core Count] [--lora_model_path [Lora Model Path]] [--prompt_cache_path [Prompt Cache File Path]]
|
||||
# example: ./build_rkllm_server_flask.sh --workshop /user/data --model_path /user/data/model.rkllm --platform rk3588 --npu_num 3
|
||||
#*****************************************************************************************#
|
||||
|
||||
#################### 检查板端是否已经安装了 pip/gradio 库 ####################
|
||||
# 1.准备板端的gradio环境
|
||||
LORA_PATH=""
|
||||
PROMPT_FILE_PATH=""
|
||||
|
||||
# Function to display help
|
||||
function show_help {
|
||||
echo "Usage: ./build_rkllm_server_flask.sh --workshop [RKLLM-Server Working Path] --model_path [Absolute Path of Converted RKLLM Model on Board] --platform [Target Platform: rk3588/rk3576] --npu_num [NPU Core Count] [--lora_path [Lora Model Path]] [--prompt_cache_path [Prompt Cache File Path]]"
|
||||
}
|
||||
|
||||
# Parse command-line options
|
||||
while [[ $# -gt 0 ]]; do
|
||||
case "$1" in
|
||||
--workshop)
|
||||
WORKING_PATH="$2"
|
||||
shift 2
|
||||
;;
|
||||
--model_path)
|
||||
MODEL_PATH="$2"
|
||||
shift 2
|
||||
;;
|
||||
--platform)
|
||||
TARGET_PLATFORM="$2"
|
||||
shift 2
|
||||
;;
|
||||
--npu_num)
|
||||
NPU_CORE_COUNT="$2"
|
||||
shift 2
|
||||
;;
|
||||
--lora_model_path)
|
||||
LORA_PATH="$2"
|
||||
shift 2
|
||||
;;
|
||||
--prompt_cache_path)
|
||||
PROMPT_FILE_PATH="$2"
|
||||
shift 2
|
||||
;;
|
||||
--help)
|
||||
show_help
|
||||
exit 0
|
||||
;;
|
||||
*)
|
||||
echo "无效的选项: $1" 1>&2
|
||||
show_help
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
done
|
||||
|
||||
#################### Check if pip and the Flask library are already installed on the board. ####################
|
||||
adb shell << EOF
|
||||
|
||||
# 检查是否安装了 pip3
|
||||
if ! command -v pip3 &> /dev/null; then
|
||||
echo "-------- pip3 未安装,将进行安装... --------"
|
||||
# 安装 pip3
|
||||
echo "-------- pip3 is not installed. Installing it now... --------"
|
||||
sudo apt update
|
||||
sudo apt install python3-pip -y
|
||||
else
|
||||
echo "-------- pip3 已经安装 --------"
|
||||
echo "-------- pip3 is already installed. --------"
|
||||
fi
|
||||
|
||||
# 检查是否安装了 flask
|
||||
if ! python3 -c "import flask" &> /dev/null; then
|
||||
echo "-------- flask 未安装,将进行安装... --------"
|
||||
echo "-------- flask is not installed. Installing it now... --------"
|
||||
# 安装 flask
|
||||
pip install flask==2.2.2 Werkzeug==2.2.2 -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
pip install flask==2.2.2 Werkzeug==2.2.2 -i https://pypi.tuna.tsinghua.edu.cn/simple --break-system-packages
|
||||
else
|
||||
echo "-------- flask 已经安装 --------"
|
||||
echo "-------- flask is already installed. --------"
|
||||
fi
|
||||
|
||||
exit
|
||||
|
||||
EOF
|
||||
|
||||
#################### 推送 server 运行的相关文件进入板端 ####################
|
||||
# 2.检查需要推送进板端的路径是否存在
|
||||
adb shell ls $2 > /dev/null 2>&1
|
||||
#################### Push the relevant files for the server to the board. ####################
|
||||
adb shell ls $WORKING_PATH > /dev/null 2>&1
|
||||
|
||||
if [ $? -ne 0 ]; then
|
||||
# 如果路径不存在,则创建路径
|
||||
adb shell mkdir -p $2
|
||||
echo "-------- rkllm_server 工作目录不存在,已创建目录 --------"
|
||||
adb shell mkdir -p $WORKING_PATH
|
||||
echo "-------- The rkllm_server working directory does not exist, so it has been created. --------"
|
||||
else
|
||||
echo "-------- rkllm_server 工作目录已存在 --------"
|
||||
echo "-------- The rkllm_server working directory already exists. --------"
|
||||
fi
|
||||
|
||||
# 3.更新 ./rkllm_server/lib 中的 librkllmrt.so 文件
|
||||
# Update the `librkllmrt.so` file in the `./rkllm_server/lib` directory.
|
||||
cp ../../runtime/Linux/librkllm_api/aarch64/librkllmrt.so ./rkllm_server/lib/
|
||||
|
||||
# 4.推送文件到 Linux 板端
|
||||
adb push ./rkllm_server $2
|
||||
adb push ./rkllm_server $WORKING_PATH
|
||||
|
||||
#################### Enter the board terminal and start the server service. ####################
|
||||
CMD="python3 flask_server.py --rkllm_model_path $MODEL_PATH --target_platform $TARGET_PLATFORM --num_npu_core $NPU_CORE_COUNT"
|
||||
if [[ -n "$LORA_PATH" ]]; then
|
||||
CMD="$CMD --lora_model_path $LORA_PATH"
|
||||
fi
|
||||
|
||||
if [[ -n "$PROMPT_FILE_PATH" ]]; then
|
||||
CMD="$CMD --prompt_cache_path $PROMPT_FILE_PATH"
|
||||
fi
|
||||
|
||||
#################### 进入板端并启动 server 服务 ####################
|
||||
# 5.进入板端启动 server 服务
|
||||
adb shell << EOF
|
||||
|
||||
cd $2/rkllm_server/
|
||||
python3 flask_server.py --target_platform $1 --rkllm_model_path $3
|
||||
cd $WORKING_PATH/rkllm_server/
|
||||
$CMD
|
||||
|
||||
EOF
|
||||
|
||||
@@ -1,61 +1,110 @@
|
||||
#!/bin/bash
|
||||
|
||||
#*****************************************************************************************#
|
||||
# 该脚本为 RKLLM-Server-Gradio 服务的一键设置脚本
|
||||
# 用户可以运行该脚本实现Linux板端的 RKLLM-Server-Gradio 服务的自动化部署。
|
||||
# 使用说明: ./build_rkllm_server_gradio.sh [目标平台:rk3588/rk3576] [RKLLM-Server工作路径] [已转换的rkllm模型在板端的绝对路径]
|
||||
# example: ./build_rkllm_server_gradio.sh rk3588 /user/data/rkllm_server /user/data/rkllm_server/model.rkllm
|
||||
# This script is an automated setup script for the RKLLM-Server-Gradio service.
|
||||
# Users can run this script to automate the deployment of the RKLLM-Server-Gradio service on a Linux board.
|
||||
# Usage: ./build_rkllm_server_gradio.sh --workshop [RKLLM-Server Working Path] --model_path [Absolute Path of Converted RKLLM Model on Board] --platform [Target Platform: rk3588/rk3576] --npu_num [NPU Core Count] [--lora_model_path [Lora Model Path]] [--prompt_cache_path [Prompt Cache File Path]]
|
||||
# example: ./build_rkllm_server_gradio.sh --workshop /user/data --model_path /user/data/model.rkllm --platform rk3588 --npu_num 3
|
||||
#*****************************************************************************************#
|
||||
|
||||
#################### 检查板端是否已经安装了 pip/gradio 库 ####################
|
||||
# 1.准备板端的gradio环境
|
||||
LORA_PATH=""
|
||||
PROMPT_FILE_PATH=""
|
||||
|
||||
# Function to display help
|
||||
function show_help {
|
||||
echo "Usage: ./build_rkllm_server_gradio.sh --workshop [RKLLM-Server Working Path] --model_path [Absolute Path of Converted RKLLM Model on Board] --platform [Target Platform: rk3588/rk3576] --npu_num [NPU Core Count] [--lora_path [Lora Model Path]] [--prompt_cache_path [Prompt Cache File Path]]"
|
||||
}
|
||||
|
||||
# Parse command-line options
|
||||
while [[ $# -gt 0 ]]; do
|
||||
case "$1" in
|
||||
--workshop)
|
||||
WORKING_PATH="$2"
|
||||
shift 2
|
||||
;;
|
||||
--model_path)
|
||||
MODEL_PATH="$2"
|
||||
shift 2
|
||||
;;
|
||||
--platform)
|
||||
TARGET_PLATFORM="$2"
|
||||
shift 2
|
||||
;;
|
||||
--npu_num)
|
||||
NPU_CORE_COUNT="$2"
|
||||
shift 2
|
||||
;;
|
||||
--lora_model_path)
|
||||
LORA_PATH="$2"
|
||||
shift 2
|
||||
;;
|
||||
--prompt_cache_path)
|
||||
PROMPT_FILE_PATH="$2"
|
||||
shift 2
|
||||
;;
|
||||
--help)
|
||||
show_help
|
||||
exit 0
|
||||
;;
|
||||
*)
|
||||
echo "无效的选项: $1" 1>&2
|
||||
show_help
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
done
|
||||
|
||||
#################### Check if pip and the Gradio library are already installed on the board. ####################
|
||||
adb shell << EOF
|
||||
|
||||
# 检查是否安装了 pip3
|
||||
if ! command -v pip3 &> /dev/null; then
|
||||
echo "-------- pip3 未安装,将进行安装... --------"
|
||||
# 安装 pip3
|
||||
echo "-------- pip3 is not installed. Installing it now... --------"
|
||||
sudo apt update
|
||||
sudo apt install python3-pip -y
|
||||
else
|
||||
echo "-------- pip3 已经安装 --------"
|
||||
echo "-------- pip3 is already installed. --------"
|
||||
fi
|
||||
|
||||
# 检查是否安装了 gradio
|
||||
if ! python3 -c "import gradio" &> /dev/null; then
|
||||
echo "-------- Gradio 未安装,将进行安装... --------"
|
||||
# 安装 Gradio
|
||||
pip3 install gradio>=4.24.0 -i https://pypi.tuna.tsinghua.edu.cn/simple/
|
||||
echo "-------- gradio is not installed. Installing it now... --------"
|
||||
pip3 install gradio>=4.24.0 -i https://pypi.tuna.tsinghua.edu.cn/simple/ --break-system-packages
|
||||
else
|
||||
echo "-------- Gradio 已经安装 --------"
|
||||
echo "-------- gradio is already installed. --------"
|
||||
fi
|
||||
|
||||
exit
|
||||
|
||||
EOF
|
||||
|
||||
#################### 推送 server 运行的相关文件进入板端 ####################
|
||||
# 2.检查需要推送进板端的路径是否存在
|
||||
adb shell ls $2 > /dev/null 2>&1
|
||||
#################### Push the relevant files for the server to the board. ####################
|
||||
adb shell ls $WORKING_PATH > /dev/null 2>&1
|
||||
|
||||
if [ $? -ne 0 ]; then
|
||||
# 如果路径不存在,则创建路径
|
||||
adb shell mkdir -p $2
|
||||
echo "-------- rkllm_server 工作目录不存在,已创建目录 --------"
|
||||
adb shell mkdir -p $WORKING_PATH
|
||||
echo "-------- The rkllm_server working directory does not exist, so it has been created. --------"
|
||||
else
|
||||
echo "-------- rkllm_server 工作目录已存在 --------"
|
||||
echo "-------- The rkllm_server working directory already exists. --------"
|
||||
fi
|
||||
|
||||
# 3.更新 ./rkllm_server/lib 中的 librkllmrt.so 文件
|
||||
# Update the `librkllmrt.so` file in the `./rkllm_server/lib` directory.
|
||||
cp ../../runtime/Linux/librkllm_api/aarch64/librkllmrt.so ./rkllm_server/lib/
|
||||
|
||||
# 4.推送文件到 Linux 板端
|
||||
adb push ./rkllm_server $2
|
||||
adb push ./rkllm_server $WORKING_PATH
|
||||
|
||||
#################### Enter the board terminal and start the server service. ####################
|
||||
CMD="python3 gradio_server.py --rkllm_model_path $MODEL_PATH --target_platform $TARGET_PLATFORM --num_npu_core $NPU_CORE_COUNT"
|
||||
|
||||
if [[ -n "$LORA_PATH" ]]; then
|
||||
CMD="$CMD --lora_model_path $LORA_PATH"
|
||||
fi
|
||||
|
||||
if [[ -n "$PROMPT_FILE_PATH" ]]; then
|
||||
CMD="$CMD --prompt_cache_path $PROMPT_FILE_PATH"
|
||||
fi
|
||||
|
||||
#################### 进入板端并启动 server 服务 ####################
|
||||
# 5.进入板端启动 server 服务
|
||||
adb shell << EOF
|
||||
|
||||
cd $2/rkllm_server/
|
||||
python3 gradio_server.py --target_platform $1 --rkllm_model_path $3
|
||||
cd $WORKING_PATH/rkllm_server/
|
||||
$CMD
|
||||
|
||||
EOF
|
||||
|
||||
@@ -2,53 +2,53 @@ import sys
|
||||
import requests
|
||||
import json
|
||||
|
||||
# 设置 Server 服务器的地址
|
||||
server_url = 'http://172.16.10.102:8080/rkllm_chat'
|
||||
# 设置是否开启流式对话
|
||||
# Set the address of the Server.
|
||||
server_url = 'http://172.16.10.79:8080/rkllm_chat'
|
||||
# Set whether to enable streaming mode.
|
||||
is_streaming = True
|
||||
|
||||
# 创建一个会话对象
|
||||
# Create a session object.
|
||||
session = requests.Session()
|
||||
session.keep_alive = False # 关闭连接池,保持长连接
|
||||
session.keep_alive = False # Close the connection pool to maintain a long connection.
|
||||
adapter = requests.adapters.HTTPAdapter(max_retries=5)
|
||||
session.mount('https://', adapter)
|
||||
session.mount('http://', adapter)
|
||||
|
||||
if __name__ == '__main__':
|
||||
print("============================")
|
||||
print("在终端中输入您的问题,即可与 RKLLM 模型进行对话....")
|
||||
print("Input your question in the terminal to start a conversation with the RKLLM model...")
|
||||
print("============================")
|
||||
# 进入循环,持续获取用户输入,并与RKLLM模型进行对话
|
||||
# Enter a loop to continuously get user input and converse with the RKLLM model.
|
||||
while True:
|
||||
try:
|
||||
user_message = input("请输入您的问题:")
|
||||
user_message = input("\n*Please enter your question:")
|
||||
if user_message == "exit":
|
||||
print("============================")
|
||||
print("程序正在退出......")
|
||||
print("The RKLLM Server is stopping......")
|
||||
print("============================")
|
||||
break
|
||||
else:
|
||||
# 设置请求头,此处的请求头实际并无作用,仅为模拟OpenAI接口设计
|
||||
# Set the request headers; in this case, the headers have no actual effect and are only used to simulate the OpenAI interface design.
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': 'not_required'
|
||||
}
|
||||
|
||||
# 准备要发送的数据
|
||||
# model: 为用户在设置RKLLM-Server时定义的模型,此处并无作用
|
||||
# messages: 用户输入的问题,RKLLM-Server将会把它作为输入,并返回模型的回复;支持在 messags 加入多个问题
|
||||
# stream: 是否开启流式对话,与OpenAI接口相同
|
||||
# Prepare the data to be sent
|
||||
# model: The model defined by the user when setting up RKLLM-Server; this has no effect here
|
||||
# messages: The user's input question, which RKLLM-Server will use as input and return the model's reply; multiple questions can be added to messages
|
||||
# stream: Whether to enable streaming conversation, similar to the OpenAI interface
|
||||
data = {
|
||||
"model": 'your_model_deploy_with_RKLLM_Server',
|
||||
"messages": [{"role": "user", "content": user_message}],
|
||||
"stream": is_streaming
|
||||
}
|
||||
|
||||
# 发送 POST 请求
|
||||
# Send a POST request
|
||||
responses = session.post(server_url, json=data, headers=headers, stream=is_streaming, verify=False)
|
||||
|
||||
if not is_streaming:
|
||||
# 解析响应
|
||||
# Parse the response
|
||||
if responses.status_code == 200:
|
||||
print("Q:", data["messages"][-1]["content"])
|
||||
print("A:", json.loads(responses.text)["choices"][-1]["message"]["content"])
|
||||
@@ -66,16 +66,13 @@ if __name__ == '__main__':
|
||||
sys.stdout.flush()
|
||||
else:
|
||||
print('Error:', responses.text)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
except KeyboardInterrupt:
|
||||
# 捕获 Ctrl-C 信号,关闭会话
|
||||
# Capture Ctrl-C signal to close the session
|
||||
session.close()
|
||||
|
||||
print("\n")
|
||||
print("============================")
|
||||
print("程序正在退出......")
|
||||
print("The RKLLM Server is stopping......")
|
||||
print("============================")
|
||||
break
|
||||
|
||||
@@ -1,44 +1,43 @@
|
||||
from gradio_client import Client
|
||||
|
||||
# 该函数通过调用Gradio Client API与RKLLM模型进行交互
|
||||
# This function interacts with the RKLLM model by calling the Gradio Client API.
|
||||
def chat_with_rkllm(user_message, history=[]):
|
||||
# 实例化Gradio Client,用户需要根据自己部署的具体网址进行修改
|
||||
client = Client("http://172.16.10.102:8080/")
|
||||
# Instantiate the Gradio Client. Users need to modify according to their specific deployment URL.
|
||||
client = Client("http://172.16.10.79:8080")
|
||||
|
||||
# 调用Gradio Client API进行交互,内部的API主要包括:
|
||||
# /get_user_input:模型获取用户输入,并将输入添加至历史记录history
|
||||
# /get_RKLLM_output:RKLLM利用已包含输入的历史记录history生成回复
|
||||
# Call the Gradio Client API for interaction. The internal APIs mainly include:
|
||||
# get_user_input: The model retrieves user input and adds it to the history record 'history'.
|
||||
# get_RKLLM_output: RKLLM generates a response using the historical record 'history' that contains the input.
|
||||
_, history = client.predict(user_message=user_message, history=history, api_name="/get_user_input")
|
||||
result_history = client.predict(history=history, api_name="/get_RKLLM_output")
|
||||
return result_history
|
||||
|
||||
if __name__ == '__main__':
|
||||
#初始化聊天记录
|
||||
result_history = []
|
||||
|
||||
print("============================")
|
||||
print("在终端中输入您的问题,即可与 RKLLM 模型进行对话....")
|
||||
print("Enter your question in the terminal to have a conversation with the RKLLM model...")
|
||||
print("============================")
|
||||
# 进入循环,持续获取用户输入,并与RKLLM模型进行对话
|
||||
# Enter a loop to continuously receive user input and have a conversation with the RKLLM model...
|
||||
while True:
|
||||
try:
|
||||
user_message = input("请输入您的问题:")
|
||||
user_message = input("Please enter your question:")
|
||||
if user_message == "exit":
|
||||
print("============================")
|
||||
print("程序正在退出......")
|
||||
print("The RKLLM Server is stopping......")
|
||||
print("============================")
|
||||
break
|
||||
else:
|
||||
# 调用chat_with_rkllm函数,获取模型的回复
|
||||
# Call the `chat_with_rkllm` function to get the model's response.
|
||||
result_history = chat_with_rkllm(user_message, result_history)
|
||||
|
||||
# 打印模型输出
|
||||
# Print the history of chatting
|
||||
print("Q:", result_history[-1][0])
|
||||
print("A:", result_history[-1][1])
|
||||
except KeyboardInterrupt:
|
||||
print("\n")
|
||||
print("============================")
|
||||
print("程序正在退出......")
|
||||
print("The RKLLM Server is stopping......")
|
||||
print("============================")
|
||||
break
|
||||
|
||||
@@ -11,68 +11,41 @@ from flask import Flask, request, jsonify, Response
|
||||
|
||||
app = Flask(__name__)
|
||||
|
||||
# 创建一个锁,用于控制多人访问Server
|
||||
lock = threading.Lock()
|
||||
PROMPT_TEXT_PREFIX = "<|im_start|>system You are a helpful assistant. <|im_end|> <|im_start|>user"
|
||||
PROMPT_TEXT_POSTFIX = "<|im_end|><|im_start|>assistant"
|
||||
|
||||
# 创建一个全局变量,用于标识服务器当前是否处于阻塞状态
|
||||
is_blocking = False
|
||||
|
||||
# 设置动态库路径
|
||||
# Set the dynamic library path
|
||||
rkllm_lib = ctypes.CDLL('lib/librkllmrt.so')
|
||||
|
||||
# 定义全局变量,用于保存回调函数的输出,便于在gradio界面中输出
|
||||
global_text = []
|
||||
global_state = -1
|
||||
split_byte_data = bytes(b"") # 用于保存分割的字节数据
|
||||
# Define the structures from the library
|
||||
RKLLM_Handle_t = ctypes.c_void_p
|
||||
userdata = ctypes.c_void_p(None)
|
||||
|
||||
# 定义动态库中的结构体
|
||||
class Token(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("logprob", ctypes.c_float),
|
||||
("id", ctypes.c_int32)
|
||||
]
|
||||
LLMCallState = ctypes.c_int
|
||||
LLMCallState.RKLLM_RUN_NORMAL = 0
|
||||
LLMCallState.RKLLM_RUN_FINISH = 1
|
||||
LLMCallState.RKLLM_RUN_ERROR = 2
|
||||
LLMCallState.RKLLM_RUN_GET_LAST_HIDDEN_LAYER = 3
|
||||
|
||||
class RKLLMResult(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("text", ctypes.c_char_p),
|
||||
("tokens", ctypes.POINTER(Token)),
|
||||
("num", ctypes.c_int32)
|
||||
]
|
||||
RKLLMInputMode = ctypes.c_int
|
||||
RKLLMInputMode.RKLLM_INPUT_PROMPT = 0
|
||||
RKLLMInputMode.RKLLM_INPUT_TOKEN = 1
|
||||
RKLLMInputMode.RKLLM_INPUT_EMBED = 2
|
||||
RKLLMInputMode.RKLLM_INPUT_MULTIMODAL = 3
|
||||
|
||||
RKLLMInferMode = ctypes.c_int
|
||||
RKLLMInferMode.RKLLM_INFER_GENERATE = 0
|
||||
RKLLMInferMode.RKLLM_INFER_GET_LAST_HIDDEN_LAYER = 1
|
||||
|
||||
# 定义回调函数
|
||||
def callback(result, userdata, state):
|
||||
global global_text, global_state, split_byte_data
|
||||
if state == 0:
|
||||
# 保存输出的token文本及RKLLM运行状态
|
||||
global_state = state
|
||||
# 需要监控当前的字节数据是否完整,不完整则进行记录,后续进行解析
|
||||
try:
|
||||
global_text.append((split_byte_data + result.contents.text).decode('utf-8'))
|
||||
print((split_byte_data + result.contents.text).decode('utf-8'), end='')
|
||||
split_byte_data = bytes(b"")
|
||||
except:
|
||||
split_byte_data += result.contents.text
|
||||
sys.stdout.flush()
|
||||
elif state == 1:
|
||||
# 保存RKLLM运行状态
|
||||
global_state = state
|
||||
print("\n")
|
||||
sys.stdout.flush()
|
||||
else:
|
||||
print("run error")
|
||||
|
||||
# Python端与C++端的回调函数连接
|
||||
callback_type = ctypes.CFUNCTYPE(None, ctypes.POINTER(RKLLMResult), ctypes.c_void_p, ctypes.c_int)
|
||||
c_callback = callback_type(callback)
|
||||
|
||||
# 定义动态库中的结构体
|
||||
class RKNNllmParam(ctypes.Structure):
|
||||
class RKLLMParam(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("model_path", ctypes.c_char_p),
|
||||
("license_path", ctypes.c_char_p),
|
||||
("num_npu_core", ctypes.c_int32),
|
||||
("max_context_len", ctypes.c_int32),
|
||||
("n_prefill_batch", ctypes.c_int32),
|
||||
("max_new_tokens", ctypes.c_int32),
|
||||
("skip_special_token", ctypes.c_bool),
|
||||
("top_k", ctypes.c_int32),
|
||||
("top_p", ctypes.c_float),
|
||||
("temperature", ctypes.c_float),
|
||||
@@ -84,58 +57,234 @@ class RKNNllmParam(ctypes.Structure):
|
||||
("mirostat_eta", ctypes.c_float),
|
||||
("logprobs", ctypes.c_bool),
|
||||
("top_logprobs", ctypes.c_int32),
|
||||
("use_gpu", ctypes.c_bool)
|
||||
("is_async", ctypes.c_bool),
|
||||
("img_start", ctypes.c_char_p),
|
||||
("img_end", ctypes.c_char_p),
|
||||
("img_content", ctypes.c_char_p),
|
||||
]
|
||||
|
||||
# 定义RKLLM_Handle_t和userdata
|
||||
RKLLM_Handle_t = ctypes.c_void_p
|
||||
userdata = ctypes.c_void_p(None)
|
||||
class RKLLMLoraAdapter(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("lora_adapter_path", ctypes.c_char_p),
|
||||
("lora_adapter_name", ctypes.c_char_p),
|
||||
("scale", ctypes.c_float)
|
||||
]
|
||||
|
||||
# 设置提示文本
|
||||
PROMPT_TEXT_PREFIX = "<|im_start|>system You are a helpful assistant. <|im_end|> <|im_start|>user"
|
||||
PROMPT_TEXT_POSTFIX = "<|im_end|><|im_start|>assistant"
|
||||
class RKLLMLoraParam(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("lora_adapter_name", ctypes.c_char_p)
|
||||
]
|
||||
|
||||
# 定义Python端的RKLLM类,其中包括了对动态库中RKLLM模型的初始化、推理及释放操作
|
||||
class RKLLMPromptCacheParam(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("save_prompt_cache", ctypes.c_int),
|
||||
("prompt_cache_path", ctypes.c_char_p)
|
||||
]
|
||||
|
||||
class RKLLMEmbedInput(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("embed", ctypes.POINTER(ctypes.c_float)),
|
||||
("n_tokens", ctypes.c_size_t)
|
||||
]
|
||||
|
||||
class RKLLMTokenInput(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("input_ids", ctypes.POINTER(ctypes.c_int32)),
|
||||
("n_tokens", ctypes.c_size_t)
|
||||
]
|
||||
|
||||
class RKLLMMultiModelInput(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("prompt", ctypes.c_char_p),
|
||||
("image_embed", ctypes.POINTER(ctypes.c_float)),
|
||||
("n_image_tokens", ctypes.c_size_t)
|
||||
]
|
||||
|
||||
class RKLLMInputUnion(ctypes.Union):
|
||||
_fields_ = [
|
||||
("prompt_input", ctypes.c_char_p),
|
||||
("embed_input", RKLLMEmbedInput),
|
||||
("token_input", RKLLMTokenInput),
|
||||
("multimodal_input", RKLLMMultiModelInput)
|
||||
]
|
||||
|
||||
class RKLLMInput(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("input_mode", ctypes.c_int),
|
||||
("input_data", RKLLMInputUnion)
|
||||
]
|
||||
class RKLLMInferParam(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("mode", RKLLMInferMode),
|
||||
("lora_params", ctypes.POINTER(RKLLMLoraParam)),
|
||||
("prompt_cache_params", ctypes.POINTER(RKLLMPromptCacheParam))
|
||||
]
|
||||
|
||||
class RKLLMResultLastHiddenLayer(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("hidden_states", ctypes.POINTER(ctypes.c_float)),
|
||||
("embd_size", ctypes.c_int),
|
||||
("num_tokens", ctypes.c_int)
|
||||
]
|
||||
|
||||
class RKLLMResult(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("text", ctypes.c_char_p),
|
||||
("size", ctypes.c_int),
|
||||
("last_hidden_layer", RKLLMResultLastHiddenLayer)
|
||||
]
|
||||
|
||||
|
||||
# Create a lock to control multi-user access to the server.
|
||||
lock = threading.Lock()
|
||||
|
||||
# Create a global variable to indicate whether the server is currently in a blocked state.
|
||||
is_blocking = False
|
||||
|
||||
# Define global variables to store the callback function output for displaying in the Gradio interface
|
||||
global_text = []
|
||||
global_state = -1
|
||||
split_byte_data = bytes(b"") # Used to store the segmented byte data
|
||||
|
||||
# Define the callback function
|
||||
def callback_impl(result, userdata, state):
|
||||
global global_text, global_state, split_byte_data
|
||||
if state == LLMCallState.RKLLM_RUN_FINISH:
|
||||
global_state = state
|
||||
print("\n")
|
||||
sys.stdout.flush()
|
||||
elif state == LLMCallState.RKLLM_RUN_ERROR:
|
||||
global_state = state
|
||||
print("run error")
|
||||
sys.stdout.flush()
|
||||
elif state == LLMCallState.RKLLM_RUN_GET_LAST_HIDDEN_LAYER:
|
||||
'''
|
||||
If using the GET_LAST_HIDDEN_LAYER function, the callback interface will return the memory pointer: last_hidden_layer, the number of tokens: num_tokens, and the size of the hidden layer: embd_size.
|
||||
With these three parameters, you can retrieve the data from last_hidden_layer.
|
||||
Note: The data needs to be retrieved during the current callback; if not obtained in time, the pointer will be released by the next callback.
|
||||
'''
|
||||
if result.last_hidden_layer.embd_size != 0 and result.last_hidden_layer.num_tokens != 0:
|
||||
data_size = result.last_hidden_layer.embd_size * result.last_hidden_layer.num_tokens * ctypes.sizeof(ctypes.c_float)
|
||||
print(f"data_size: {data_size}")
|
||||
global_text.append(f"data_size: {data_size}\n")
|
||||
output_path = os.getcwd() + "/last_hidden_layer.bin"
|
||||
with open(output_path, "wb") as outFile:
|
||||
data = ctypes.cast(result.last_hidden_layer.hidden_states, ctypes.POINTER(ctypes.c_float))
|
||||
float_array_type = ctypes.c_float * (data_size // ctypes.sizeof(ctypes.c_float))
|
||||
float_array = float_array_type.from_address(ctypes.addressof(data.contents))
|
||||
outFile.write(bytearray(float_array))
|
||||
print(f"Data saved to {output_path} successfully!")
|
||||
global_text.append(f"Data saved to {output_path} successfully!")
|
||||
else:
|
||||
print("Invalid hidden layer data.")
|
||||
global_text.append("Invalid hidden layer data.")
|
||||
global_state = state
|
||||
time.sleep(0.05) # Delay for 0.05 seconds to wait for the output result
|
||||
sys.stdout.flush()
|
||||
else:
|
||||
# Save the output token text and the RKLLM running state
|
||||
global_state = state
|
||||
# Monitor if the current byte data is complete; if incomplete, record it for later parsing
|
||||
try:
|
||||
global_text.append((split_byte_data + result.contents.text).decode('utf-8'))
|
||||
print((split_byte_data + result.contents.text).decode('utf-8'), end='')
|
||||
split_byte_data = bytes(b"")
|
||||
except:
|
||||
split_byte_data += result.contents.text
|
||||
sys.stdout.flush()
|
||||
|
||||
# Connect the callback function between the Python side and the C++ side
|
||||
callback_type = ctypes.CFUNCTYPE(None, ctypes.POINTER(RKLLMResult), ctypes.c_void_p, ctypes.c_int)
|
||||
callback = callback_type(callback_impl)
|
||||
|
||||
# Define the RKLLM class, which includes initialization, inference, and release operations for the RKLLM model in the dynamic library
|
||||
class RKLLM(object):
|
||||
def __init__(self, model_path, target_platform):
|
||||
rknnllm_param = RKNNllmParam()
|
||||
rknnllm_param.model_path = bytes(model_path, 'utf-8')
|
||||
if target_platform == "rk3588":
|
||||
rknnllm_param.num_npu_core = 3
|
||||
elif target_platform == "rk3576":
|
||||
rknnllm_param.num_npu_core = 1
|
||||
rknnllm_param.max_context_len = 320
|
||||
rknnllm_param.max_new_tokens = 512
|
||||
rknnllm_param.top_k = 1
|
||||
rknnllm_param.top_p = 0.9
|
||||
rknnllm_param.temperature = 0.8
|
||||
rknnllm_param.repeat_penalty = 1.1
|
||||
rknnllm_param.frequency_penalty = 0.0
|
||||
rknnllm_param.presence_penalty = 0.0
|
||||
rknnllm_param.mirostat = 0
|
||||
rknnllm_param.mirostat_tau = 5.0
|
||||
rknnllm_param.mirostat_eta = 0.1
|
||||
rknnllm_param.logprobs = False
|
||||
rknnllm_param.top_logprobs = 5
|
||||
rknnllm_param.use_gpu = True
|
||||
def __init__(self, model_path, num_npu_core, lora_model_path = None, prompt_cache_path = None):
|
||||
rkllm_param = RKLLMParam()
|
||||
rkllm_param.model_path = bytes(model_path, 'utf-8')
|
||||
rkllm_param.license_path = None
|
||||
rkllm_param.num_npu_core = num_npu_core
|
||||
|
||||
rkllm_param.max_context_len = 512
|
||||
rkllm_param.n_prefill_batch = 512
|
||||
rkllm_param.max_new_tokens = -1
|
||||
rkllm_param.skip_special_token = True
|
||||
|
||||
rkllm_param.top_k = 1
|
||||
rkllm_param.top_p = 0.9
|
||||
rkllm_param.temperature = 0.8
|
||||
rkllm_param.repeat_penalty = 1.1
|
||||
rkllm_param.frequency_penalty = 0.0
|
||||
rkllm_param.presence_penalty = 0.0
|
||||
|
||||
rkllm_param.mirostat = 0
|
||||
rkllm_param.mirostat_tau = 5.0
|
||||
rkllm_param.mirostat_eta = 0.1
|
||||
|
||||
rkllm_param.logprobs = False
|
||||
rkllm_param.top_logprobs = 5
|
||||
rkllm_param.is_async = False
|
||||
|
||||
rkllm_param.img_start = "".encode('utf-8')
|
||||
rkllm_param.img_end = "".encode('utf-8')
|
||||
rkllm_param.img_content = "".encode('utf-8')
|
||||
|
||||
self.handle = RKLLM_Handle_t()
|
||||
|
||||
self.rkllm_init = rkllm_lib.rkllm_init
|
||||
self.rkllm_init.argtypes = [ctypes.POINTER(RKLLM_Handle_t), ctypes.POINTER(RKNNllmParam), callback_type]
|
||||
self.rkllm_init.argtypes = [ctypes.POINTER(RKLLM_Handle_t), RKLLMParam, callback_type]
|
||||
self.rkllm_init.restype = ctypes.c_int
|
||||
self.rkllm_init(ctypes.byref(self.handle), rknnllm_param, c_callback)
|
||||
self.rkllm_init(ctypes.byref(self.handle), rkllm_param, callback)
|
||||
|
||||
self.rkllm_run = rkllm_lib.rkllm_run
|
||||
self.rkllm_run.argtypes = [RKLLM_Handle_t, ctypes.POINTER(ctypes.c_char), ctypes.c_void_p]
|
||||
self.rkllm_run.argtypes = [RKLLM_Handle_t, ctypes.POINTER(RKLLMInput), ctypes.POINTER(RKLLMInferParam), ctypes.c_void_p]
|
||||
self.rkllm_run.restype = ctypes.c_int
|
||||
|
||||
self.rkllm_destroy = rkllm_lib.rkllm_destroy
|
||||
self.rkllm_destroy.argtypes = [RKLLM_Handle_t]
|
||||
self.rkllm_destroy.restype = ctypes.c_int
|
||||
|
||||
self.lora_adapter_path = None
|
||||
self.lora_model_name = None
|
||||
if lora_model_path:
|
||||
self.lora_adapter_path = lora_model_path
|
||||
self.lora_adapter_name = "test"
|
||||
|
||||
lora_adapter = RKLLMLoraAdapter()
|
||||
ctypes.memset(ctypes.byref(lora_adapter), 0, ctypes.sizeof(RKLLMLoraAdapter))
|
||||
lora_adapter.lora_adapter_path = ctypes.c_char_p((self.lora_adapter_path).encode('utf-8'))
|
||||
lora_adapter.lora_adapter_name = ctypes.c_char_p((self.lora_adapter_name).encode('utf-8'))
|
||||
lora_adapter.scale = 1.0
|
||||
|
||||
rkllm_load_lora = rkllm_lib.rkllm_load_lora
|
||||
rkllm_load_lora.argtypes = [RKLLM_Handle_t, ctypes.POINTER(RKLLMLoraAdapter)]
|
||||
rkllm_load_lora.restype = ctypes.c_int
|
||||
rkllm_load_lora(self.handle, ctypes.byref(lora_adapter))
|
||||
|
||||
self.prompt_cache_path = None
|
||||
if prompt_cache_path:
|
||||
self.prompt_cache_path = prompt_cache_path
|
||||
|
||||
rkllm_load_prompt_cache = rkllm_lib.rkllm_load_prompt_cache
|
||||
rkllm_load_prompt_cache.argtypes = [RKLLM_Handle_t, ctypes.c_char_p]
|
||||
rkllm_load_prompt_cache.restype = ctypes.c_int
|
||||
rkllm_load_prompt_cache(self.handle, ctypes.c_char_p((prompt_cache_path).encode('utf-8')))
|
||||
|
||||
def run(self, prompt):
|
||||
prompt = bytes(PROMPT_TEXT_PREFIX + prompt + PROMPT_TEXT_POSTFIX, 'utf-8')
|
||||
self.rkllm_run(self.handle, prompt, ctypes.byref(userdata))
|
||||
rkllm_lora_params = None
|
||||
if self.lora_model_name:
|
||||
rkllm_lora_params = RKLLMLoraParam()
|
||||
rkllm_lora_params.lora_adapter_name = ctypes.c_char_p((self.lora_model_name).encode('utf-8'))
|
||||
|
||||
rkllm_infer_params = RKLLMInferParam()
|
||||
ctypes.memset(ctypes.byref(rkllm_infer_params), 0, ctypes.sizeof(RKLLMInferParam))
|
||||
rkllm_infer_params.mode = RKLLMInferMode.RKLLM_INFER_GENERATE
|
||||
rkllm_infer_params.lora_params = ctypes.byref(rkllm_lora_params) if rkllm_lora_params else None
|
||||
|
||||
rkllm_input = RKLLMInput()
|
||||
rkllm_input.input_mode = RKLLMInputMode.RKLLM_INPUT_PROMPT
|
||||
rkllm_input.input_data.prompt_input = ctypes.c_char_p((PROMPT_TEXT_PREFIX + prompt + PROMPT_TEXT_POSTFIX).encode('utf-8'))
|
||||
self.rkllm_run(self.handle, ctypes.byref(rkllm_input), ctypes.byref(rkllm_infer_params), None)
|
||||
return
|
||||
|
||||
def release(self):
|
||||
@@ -143,62 +292,81 @@ class RKLLM(object):
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--target_platform', help='目标平台: 如rk3588/rk3576;')
|
||||
parser.add_argument('--rkllm_model_path', help='Linux板端上已转换好的rkllm模型的绝对路径')
|
||||
parser.add_argument('--rkllm_model_path', type=str, required=True, help='Absolute path of the converted RKLLM model on the Linux board;')
|
||||
parser.add_argument('--target_platform', type=str, required=True, help='Target platform: e.g., rk3588/rk3576;')
|
||||
parser.add_argument('--num_npu_core', type=int, required=True, help='Target num_npu_core;')
|
||||
parser.add_argument('--lora_model_path', type=str, help='Absolute path of the lora_model on the Linux board;')
|
||||
parser.add_argument('--prompt_cache_path', type=str, help='Absolute path of the prompt_cache file on the Linux board;')
|
||||
args = parser.parse_args()
|
||||
|
||||
if not (args.target_platform in ["rk3588", "rk3576"]):
|
||||
print("====== Error: 请指定正确的目标平台: rk3588/rk3576 ======")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
if not os.path.exists(args.rkllm_model_path):
|
||||
print("====== Error: 请给出准确的rkllm模型路径,需注意是板端的绝对路径 ======")
|
||||
print("Error: Please provide the correct rkllm model path, and ensure it is the absolute path on the board.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
# 定频设置
|
||||
if not (args.target_platform in ["rk3588", "rk3576"]):
|
||||
print("Error: Please specify the correct target platform: rk3588/rk3576.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
if not (args.num_npu_core in [1, 2, 3]):
|
||||
print("Error: rk3576 supports 1/2 cores, rk3588 supports 1/2/3 cores, please specify the correct number of cores.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
if args.lora_model_path:
|
||||
if not os.path.exists(args.lora_model_path):
|
||||
print("Error: Please provide the correct lora_model path, and advise it is the absolute path on the board.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
if args.prompt_cache_path:
|
||||
if not os.path.exists(args.prompt_cache_path):
|
||||
print("Error: Please provide the correct prompt_cache_file path, and advise it is the absolute path on the board.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
# Fix frequency
|
||||
command = "sudo bash fix_freq_{}.sh".format(args.target_platform)
|
||||
subprocess.run(command, shell=True)
|
||||
|
||||
# 设置文件描述符限制
|
||||
# Set resource limit
|
||||
resource.setrlimit(resource.RLIMIT_NOFILE, (102400, 102400))
|
||||
|
||||
# 初始化RKLLM模型
|
||||
# Initialize RKLLM model
|
||||
print("=========init....===========")
|
||||
sys.stdout.flush()
|
||||
target_platform = args.target_platform
|
||||
model_path = args.rkllm_model_path
|
||||
rkllm_model = RKLLM(model_path, target_platform)
|
||||
print("RKLLM初始化成功!")
|
||||
num_npu_core = args.num_npu_core
|
||||
rkllm_model = RKLLM(model_path, num_npu_core, args.lora_model_path, args.prompt_cache_path)
|
||||
print("RKLLM Model has been initialized successfully!")
|
||||
print("==============================")
|
||||
sys.stdout.flush()
|
||||
|
||||
# 创建一个函数用于接受用户使用 request 发送的数据
|
||||
# Create a function to receive data sent by the user using a request
|
||||
@app.route('/rkllm_chat', methods=['POST'])
|
||||
def receive_message():
|
||||
# 链接全局变量,获取回调函数的输出信息
|
||||
# Link global variables to retrieve the output information from the callback function
|
||||
global global_text, global_state
|
||||
global is_blocking
|
||||
|
||||
# 如果服务器正在阻塞状态,则返回特定响应
|
||||
# If the server is in a blocking state, return a specific response.
|
||||
if is_blocking or global_state==0:
|
||||
return jsonify({'status': 'error', 'message': 'RKLLM_Server is busy! Maybe you can try again later.'}), 503
|
||||
|
||||
# 加锁
|
||||
lock.acquire()
|
||||
try:
|
||||
# 设置服务器为阻塞状态
|
||||
# Set the server to a blocking state.
|
||||
is_blocking = True
|
||||
|
||||
# 获取 POST 请求中的 JSON 数据
|
||||
# Get JSON data from the POST request.
|
||||
data = request.json
|
||||
if data and 'messages' in data:
|
||||
# 重置全局变量
|
||||
# Reset global variables.
|
||||
global_text = []
|
||||
global_state = -1
|
||||
|
||||
# 定义返回的结构体
|
||||
# Define the structure for the returned response.
|
||||
rkllm_responses = {
|
||||
"id": "rkllm_chat",
|
||||
"object": "rkllm_chat",
|
||||
@@ -212,18 +380,18 @@ if __name__ == "__main__":
|
||||
}
|
||||
|
||||
if not "stream" in data.keys() or data["stream"] == False:
|
||||
# 在这里处理收到的数据
|
||||
# Process the received data here.
|
||||
messages = data['messages']
|
||||
print("Received messages:", messages)
|
||||
for index, message in enumerate(messages):
|
||||
input_prompt = message['content']
|
||||
rkllm_output = ""
|
||||
|
||||
# 创建模型推理的线程
|
||||
# Create a thread for model inference.
|
||||
model_thread = threading.Thread(target=rkllm_model.run, args=(input_prompt,))
|
||||
model_thread.start()
|
||||
|
||||
# 等待模型运行完成,定时检查模型的推理线程
|
||||
# Wait for the model to finish running and periodically check the inference thread of the model.
|
||||
model_thread_finished = False
|
||||
while not model_thread_finished:
|
||||
while len(global_text) > 0:
|
||||
@@ -245,7 +413,6 @@ if __name__ == "__main__":
|
||||
)
|
||||
return jsonify(rkllm_responses), 200
|
||||
else:
|
||||
# 在这里处理收到的数据
|
||||
messages = data['messages']
|
||||
print("Received messages:", messages)
|
||||
for index, message in enumerate(messages):
|
||||
@@ -253,11 +420,9 @@ if __name__ == "__main__":
|
||||
rkllm_output = ""
|
||||
|
||||
def generate():
|
||||
# 创建模型推理的线程
|
||||
model_thread = threading.Thread(target=rkllm_model.run, args=(input_prompt,))
|
||||
model_thread.start()
|
||||
|
||||
# 等待模型运行完成,定时检查模型的推理线程
|
||||
model_thread_finished = False
|
||||
while not model_thread_finished:
|
||||
while len(global_text) > 0:
|
||||
@@ -282,16 +447,14 @@ if __name__ == "__main__":
|
||||
else:
|
||||
return jsonify({'status': 'error', 'message': 'Invalid JSON data!'}), 400
|
||||
finally:
|
||||
# 释放锁
|
||||
lock.release()
|
||||
# 将服务器状态设置为非阻塞
|
||||
is_blocking = False
|
||||
|
||||
# 启动 Flask 应用程序
|
||||
# Start the Flask application.
|
||||
# app.run(host='0.0.0.0', port=8080)
|
||||
app.run(host='0.0.0.0', port=8080, threaded=True, debug=False)
|
||||
|
||||
print("====================")
|
||||
print("RKLLM模型推理结束, 释放RKLLM模型资源...")
|
||||
print("RKLLM model inference completed, releasing RKLLM model resources...")
|
||||
rkllm_model.release()
|
||||
print("====================")
|
||||
|
||||
@@ -8,65 +8,45 @@ import time
|
||||
import gradio as gr
|
||||
import argparse
|
||||
|
||||
# 设定环境变量
|
||||
PROMPT_TEXT_PREFIX = "<|im_start|>system You are a helpful assistant. <|im_end|> <|im_start|>user"
|
||||
PROMPT_TEXT_POSTFIX = "<|im_end|><|im_start|>assistant"
|
||||
|
||||
# Set environment variables
|
||||
os.environ["GRADIO_SERVER_NAME"] = "0.0.0.0"
|
||||
os.environ["GRADIO_SERVER_PORT"] = "8080"
|
||||
|
||||
# 设置动态库路径
|
||||
# Set the dynamic library path
|
||||
rkllm_lib = ctypes.CDLL('lib/librkllmrt.so')
|
||||
|
||||
# 定义全局变量,用于保存回调函数的输出,便于在gradio界面中输出
|
||||
global_text = []
|
||||
global_state = -1
|
||||
split_byte_data = bytes(b"") # 用于保存分割的字节数据
|
||||
# Define the structures from the library
|
||||
RKLLM_Handle_t = ctypes.c_void_p
|
||||
userdata = ctypes.c_void_p(None)
|
||||
|
||||
# 定义动态库中的结构体
|
||||
class Token(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("logprob", ctypes.c_float),
|
||||
("id", ctypes.c_int32)
|
||||
]
|
||||
LLMCallState = ctypes.c_int
|
||||
LLMCallState.RKLLM_RUN_NORMAL = 0
|
||||
LLMCallState.RKLLM_RUN_FINISH = 1
|
||||
LLMCallState.RKLLM_RUN_ERROR = 2
|
||||
LLMCallState.RKLLM_RUN_GET_LAST_HIDDEN_LAYER = 3
|
||||
|
||||
class RKLLMResult(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("text", ctypes.c_char_p),
|
||||
("tokens", ctypes.POINTER(Token)),
|
||||
("num", ctypes.c_int32)
|
||||
]
|
||||
RKLLMInputMode = ctypes.c_int
|
||||
RKLLMInputMode.RKLLM_INPUT_PROMPT = 0
|
||||
RKLLMInputMode.RKLLM_INPUT_TOKEN = 1
|
||||
RKLLMInputMode.RKLLM_INPUT_EMBED = 2
|
||||
RKLLMInputMode.RKLLM_INPUT_MULTIMODAL = 3
|
||||
|
||||
# 定义回调函数
|
||||
def callback(result, userdata, state):
|
||||
global global_text, global_state, split_byte_data
|
||||
if state == 0:
|
||||
# 保存输出的token文本及RKLLM运行状态
|
||||
global_state = state
|
||||
# 需要监控当前的字节数据是否完整,不完整则进行记录,后续进行解析
|
||||
try:
|
||||
global_text.append((split_byte_data + result.contents.text).decode('utf-8'))
|
||||
print((split_byte_data + result.contents.text).decode('utf-8'), end='')
|
||||
split_byte_data = bytes(b"")
|
||||
except:
|
||||
split_byte_data += result.contents.text
|
||||
sys.stdout.flush()
|
||||
elif state == 1:
|
||||
# 保存RKLLM运行状态
|
||||
global_state = state
|
||||
print("\n")
|
||||
sys.stdout.flush()
|
||||
else:
|
||||
print("run error")
|
||||
RKLLMInferMode = ctypes.c_int
|
||||
RKLLMInferMode.RKLLM_INFER_GENERATE = 0
|
||||
RKLLMInferMode.RKLLM_INFER_GET_LAST_HIDDEN_LAYER = 1
|
||||
|
||||
# Python端与C++端的回调函数连接
|
||||
callback_type = ctypes.CFUNCTYPE(None, ctypes.POINTER(RKLLMResult), ctypes.c_void_p, ctypes.c_int)
|
||||
c_callback = callback_type(callback)
|
||||
|
||||
# 定义动态库中的结构体
|
||||
class RKNNllmParam(ctypes.Structure):
|
||||
class RKLLMParam(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("model_path", ctypes.c_char_p),
|
||||
("license_path", ctypes.c_char_p),
|
||||
("num_npu_core", ctypes.c_int32),
|
||||
("max_context_len", ctypes.c_int32),
|
||||
("n_prefill_batch", ctypes.c_int32),
|
||||
("max_new_tokens", ctypes.c_int32),
|
||||
("skip_special_token", ctypes.c_bool),
|
||||
("top_k", ctypes.c_int32),
|
||||
("top_p", ctypes.c_float),
|
||||
("temperature", ctypes.c_float),
|
||||
@@ -78,58 +58,227 @@ class RKNNllmParam(ctypes.Structure):
|
||||
("mirostat_eta", ctypes.c_float),
|
||||
("logprobs", ctypes.c_bool),
|
||||
("top_logprobs", ctypes.c_int32),
|
||||
("use_gpu", ctypes.c_bool)
|
||||
("is_async", ctypes.c_bool),
|
||||
("img_start", ctypes.c_char_p),
|
||||
("img_end", ctypes.c_char_p),
|
||||
("img_content", ctypes.c_char_p),
|
||||
]
|
||||
|
||||
# 定义RKLLM_Handle_t和userdata
|
||||
RKLLM_Handle_t = ctypes.c_void_p
|
||||
userdata = ctypes.c_void_p(None)
|
||||
class RKLLMLoraAdapter(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("lora_adapter_path", ctypes.c_char_p),
|
||||
("lora_adapter_name", ctypes.c_char_p),
|
||||
("scale", ctypes.c_float)
|
||||
]
|
||||
|
||||
# 设置提示文本
|
||||
PROMPT_TEXT_PREFIX = "<|im_start|>system You are a helpful assistant. <|im_end|> <|im_start|>user"
|
||||
PROMPT_TEXT_POSTFIX = "<|im_end|><|im_start|>assistant"
|
||||
class RKLLMLoraParam(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("lora_adapter_name", ctypes.c_char_p)
|
||||
]
|
||||
|
||||
# 定义Python端的RKLLM类,其中包括了对动态库中RKLLM模型的初始化、推理及释放操作
|
||||
class RKLLMPromptCacheParam(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("save_prompt_cache", ctypes.c_int),
|
||||
("prompt_cache_path", ctypes.c_char_p)
|
||||
]
|
||||
|
||||
class RKLLMEmbedInput(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("embed", ctypes.POINTER(ctypes.c_float)),
|
||||
("n_tokens", ctypes.c_size_t)
|
||||
]
|
||||
|
||||
class RKLLMTokenInput(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("input_ids", ctypes.POINTER(ctypes.c_int32)),
|
||||
("n_tokens", ctypes.c_size_t)
|
||||
]
|
||||
|
||||
class RKLLMMultiModelInput(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("prompt", ctypes.c_char_p),
|
||||
("image_embed", ctypes.POINTER(ctypes.c_float)),
|
||||
("n_image_tokens", ctypes.c_size_t)
|
||||
]
|
||||
|
||||
class RKLLMInputUnion(ctypes.Union):
|
||||
_fields_ = [
|
||||
("prompt_input", ctypes.c_char_p),
|
||||
("embed_input", RKLLMEmbedInput),
|
||||
("token_input", RKLLMTokenInput),
|
||||
("multimodal_input", RKLLMMultiModelInput)
|
||||
]
|
||||
|
||||
class RKLLMInput(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("input_mode", ctypes.c_int),
|
||||
("input_data", RKLLMInputUnion)
|
||||
]
|
||||
class RKLLMInferParam(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("mode", RKLLMInferMode),
|
||||
("lora_params", ctypes.POINTER(RKLLMLoraParam)),
|
||||
("prompt_cache_params", ctypes.POINTER(RKLLMPromptCacheParam))
|
||||
]
|
||||
|
||||
class RKLLMResultLastHiddenLayer(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("hidden_states", ctypes.POINTER(ctypes.c_float)),
|
||||
("embd_size", ctypes.c_int),
|
||||
("num_tokens", ctypes.c_int)
|
||||
]
|
||||
|
||||
class RKLLMResult(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("text", ctypes.c_char_p),
|
||||
("size", ctypes.c_int),
|
||||
("last_hidden_layer", RKLLMResultLastHiddenLayer)
|
||||
]
|
||||
|
||||
# Define global variables to store the callback function output for displaying in the Gradio interface
|
||||
global_text = []
|
||||
global_state = -1
|
||||
split_byte_data = bytes(b"") # Used to store the segmented byte data
|
||||
|
||||
# Define the callback function
|
||||
def callback_impl(result, userdata, state):
|
||||
global global_text, global_state, split_byte_data
|
||||
if state == LLMCallState.RKLLM_RUN_FINISH:
|
||||
global_state = state
|
||||
print("\n")
|
||||
sys.stdout.flush()
|
||||
elif state == LLMCallState.RKLLM_RUN_ERROR:
|
||||
global_state = state
|
||||
print("run error")
|
||||
sys.stdout.flush()
|
||||
elif state == LLMCallState.RKLLM_RUN_GET_LAST_HIDDEN_LAYER:
|
||||
'''
|
||||
If using the GET_LAST_HIDDEN_LAYER function, the callback interface will return the memory pointer: last_hidden_layer, the number of tokens: num_tokens, and the size of the hidden layer: embd_size.
|
||||
With these three parameters, you can retrieve the data from last_hidden_layer.
|
||||
Note: The data needs to be retrieved during the current callback; if not obtained in time, the pointer will be released by the next callback.
|
||||
'''
|
||||
if result.last_hidden_layer.embd_size != 0 and result.last_hidden_layer.num_tokens != 0:
|
||||
data_size = result.last_hidden_layer.embd_size * result.last_hidden_layer.num_tokens * ctypes.sizeof(ctypes.c_float)
|
||||
print(f"data_size: {data_size}")
|
||||
global_text.append(f"data_size: {data_size}\n")
|
||||
output_path = os.getcwd() + "/last_hidden_layer.bin"
|
||||
with open(output_path, "wb") as outFile:
|
||||
data = ctypes.cast(result.last_hidden_layer.hidden_states, ctypes.POINTER(ctypes.c_float))
|
||||
float_array_type = ctypes.c_float * (data_size // ctypes.sizeof(ctypes.c_float))
|
||||
float_array = float_array_type.from_address(ctypes.addressof(data.contents))
|
||||
outFile.write(bytearray(float_array))
|
||||
print(f"Data saved to {output_path} successfully!")
|
||||
global_text.append(f"Data saved to {output_path} successfully!")
|
||||
else:
|
||||
print("Invalid hidden layer data.")
|
||||
global_text.append("Invalid hidden layer data.")
|
||||
global_state = state
|
||||
time.sleep(0.05) # Delay for 0.05 seconds to wait for the output result
|
||||
sys.stdout.flush()
|
||||
else:
|
||||
# Save the output token text and the RKLLM running state
|
||||
global_state = state
|
||||
# Monitor if the current byte data is complete; if incomplete, record it for later parsing
|
||||
try:
|
||||
global_text.append((split_byte_data + result.contents.text).decode('utf-8'))
|
||||
print((split_byte_data + result.contents.text).decode('utf-8'), end='')
|
||||
split_byte_data = bytes(b"")
|
||||
except:
|
||||
split_byte_data += result.contents.text
|
||||
sys.stdout.flush()
|
||||
|
||||
# Connect the callback function between the Python side and the C++ side
|
||||
callback_type = ctypes.CFUNCTYPE(None, ctypes.POINTER(RKLLMResult), ctypes.c_void_p, ctypes.c_int)
|
||||
callback = callback_type(callback_impl)
|
||||
|
||||
# Define the RKLLM class, which includes initialization, inference, and release operations for the RKLLM model in the dynamic library
|
||||
class RKLLM(object):
|
||||
def __init__(self, model_path, target_platform):
|
||||
rknnllm_param = RKNNllmParam()
|
||||
rknnllm_param.model_path = bytes(model_path, 'utf-8')
|
||||
if target_platform == "rk3588":
|
||||
rknnllm_param.num_npu_core = 3
|
||||
elif target_platform == "rk3576":
|
||||
rknnllm_param.num_npu_core = 2
|
||||
rknnllm_param.max_context_len = 320
|
||||
rknnllm_param.max_new_tokens = 512
|
||||
rknnllm_param.top_k = 1
|
||||
rknnllm_param.top_p = 0.9
|
||||
rknnllm_param.temperature = 0.8
|
||||
rknnllm_param.repeat_penalty = 1.1
|
||||
rknnllm_param.frequency_penalty = 0.0
|
||||
rknnllm_param.presence_penalty = 0.0
|
||||
rknnllm_param.mirostat = 0
|
||||
rknnllm_param.mirostat_tau = 5.0
|
||||
rknnllm_param.mirostat_eta = 0.1
|
||||
rknnllm_param.logprobs = False
|
||||
rknnllm_param.top_logprobs = 5
|
||||
rknnllm_param.use_gpu = True
|
||||
def __init__(self, model_path, num_npu_core, lora_model_path = None, prompt_cache_path = None):
|
||||
rkllm_param = RKLLMParam()
|
||||
rkllm_param.model_path = bytes(model_path, 'utf-8')
|
||||
rkllm_param.license_path = None
|
||||
rkllm_param.num_npu_core = num_npu_core
|
||||
|
||||
rkllm_param.max_context_len = 512
|
||||
rkllm_param.n_prefill_batch = 512
|
||||
rkllm_param.max_new_tokens = -1
|
||||
rkllm_param.skip_special_token = True
|
||||
|
||||
rkllm_param.top_k = 1
|
||||
rkllm_param.top_p = 0.9
|
||||
rkllm_param.temperature = 0.8
|
||||
rkllm_param.repeat_penalty = 1.1
|
||||
rkllm_param.frequency_penalty = 0.0
|
||||
rkllm_param.presence_penalty = 0.0
|
||||
|
||||
rkllm_param.mirostat = 0
|
||||
rkllm_param.mirostat_tau = 5.0
|
||||
rkllm_param.mirostat_eta = 0.1
|
||||
|
||||
rkllm_param.logprobs = False
|
||||
rkllm_param.top_logprobs = 5
|
||||
rkllm_param.is_async = False
|
||||
|
||||
rkllm_param.img_start = "".encode('utf-8')
|
||||
rkllm_param.img_end = "".encode('utf-8')
|
||||
rkllm_param.img_content = "".encode('utf-8')
|
||||
|
||||
self.handle = RKLLM_Handle_t()
|
||||
|
||||
self.rkllm_init = rkllm_lib.rkllm_init
|
||||
self.rkllm_init.argtypes = [ctypes.POINTER(RKLLM_Handle_t), ctypes.POINTER(RKNNllmParam), callback_type]
|
||||
self.rkllm_init.argtypes = [ctypes.POINTER(RKLLM_Handle_t), RKLLMParam, callback_type]
|
||||
self.rkllm_init.restype = ctypes.c_int
|
||||
self.rkllm_init(ctypes.byref(self.handle), rknnllm_param, c_callback)
|
||||
self.rkllm_init(ctypes.byref(self.handle), rkllm_param, callback)
|
||||
|
||||
self.rkllm_run = rkllm_lib.rkllm_run
|
||||
self.rkllm_run.argtypes = [RKLLM_Handle_t, ctypes.POINTER(ctypes.c_char), ctypes.c_void_p]
|
||||
self.rkllm_run.argtypes = [RKLLM_Handle_t, ctypes.POINTER(RKLLMInput), ctypes.POINTER(RKLLMInferParam), ctypes.c_void_p]
|
||||
self.rkllm_run.restype = ctypes.c_int
|
||||
|
||||
self.rkllm_destroy = rkllm_lib.rkllm_destroy
|
||||
self.rkllm_destroy.argtypes = [RKLLM_Handle_t]
|
||||
self.rkllm_destroy.restype = ctypes.c_int
|
||||
|
||||
self.lora_adapter_path = None
|
||||
self.lora_model_name = None
|
||||
if lora_model_path:
|
||||
self.lora_adapter_path = lora_model_path
|
||||
self.lora_adapter_name = "test"
|
||||
|
||||
lora_adapter = RKLLMLoraAdapter()
|
||||
ctypes.memset(ctypes.byref(lora_adapter), 0, ctypes.sizeof(RKLLMLoraAdapter))
|
||||
lora_adapter.lora_adapter_path = ctypes.c_char_p((self.lora_adapter_path).encode('utf-8'))
|
||||
lora_adapter.lora_adapter_name = ctypes.c_char_p((self.lora_adapter_name).encode('utf-8'))
|
||||
lora_adapter.scale = 1.0
|
||||
|
||||
rkllm_load_lora = rkllm_lib.rkllm_load_lora
|
||||
rkllm_load_lora.argtypes = [RKLLM_Handle_t, ctypes.POINTER(RKLLMLoraAdapter)]
|
||||
rkllm_load_lora.restype = ctypes.c_int
|
||||
rkllm_load_lora(self.handle, ctypes.byref(lora_adapter))
|
||||
|
||||
self.prompt_cache_path = None
|
||||
if prompt_cache_path:
|
||||
self.prompt_cache_path = prompt_cache_path
|
||||
|
||||
rkllm_load_prompt_cache = rkllm_lib.rkllm_load_prompt_cache
|
||||
rkllm_load_prompt_cache.argtypes = [RKLLM_Handle_t, ctypes.c_char_p]
|
||||
rkllm_load_prompt_cache.restype = ctypes.c_int
|
||||
rkllm_load_prompt_cache(self.handle, ctypes.c_char_p((prompt_cache_path).encode('utf-8')))
|
||||
|
||||
def run(self, prompt):
|
||||
prompt = bytes(PROMPT_TEXT_PREFIX + prompt + PROMPT_TEXT_POSTFIX, 'utf-8')
|
||||
self.rkllm_run(self.handle, prompt, ctypes.byref(userdata))
|
||||
rkllm_lora_params = None
|
||||
if self.lora_model_name:
|
||||
rkllm_lora_params = RKLLMLoraParam()
|
||||
rkllm_lora_params.lora_adapter_name = ctypes.c_char_p((self.lora_model_name).encode('utf-8'))
|
||||
|
||||
rkllm_infer_params = RKLLMInferParam()
|
||||
ctypes.memset(ctypes.byref(rkllm_infer_params), 0, ctypes.sizeof(RKLLMInferParam))
|
||||
rkllm_infer_params.mode = RKLLMInferMode.RKLLM_INFER_GENERATE
|
||||
rkllm_infer_params.lora_params = ctypes.byref(rkllm_lora_params) if rkllm_lora_params else None
|
||||
|
||||
rkllm_input = RKLLMInput()
|
||||
rkllm_input.input_mode = RKLLMInputMode.RKLLM_INPUT_PROMPT
|
||||
rkllm_input.input_data.prompt_input = ctypes.c_char_p((PROMPT_TEXT_PREFIX + prompt + PROMPT_TEXT_POSTFIX).encode('utf-8'))
|
||||
self.rkllm_run(self.handle, ctypes.byref(rkllm_input), ctypes.byref(rkllm_infer_params), None)
|
||||
return
|
||||
|
||||
def release(self):
|
||||
@@ -137,92 +286,112 @@ class RKLLM(object):
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--target_platform', help='目标平台: 如rk3588/rk3576;')
|
||||
parser.add_argument('--rkllm_model_path', help='Linux板端上已转换好的rkllm模型的绝对路径')
|
||||
parser.add_argument('--rkllm_model_path', type=str, required=True, help='Absolute path of the converted RKLLM model on the Linux board;')
|
||||
parser.add_argument('--target_platform', type=str, required=True, help='Target platform: e.g., rk3588/rk3576;')
|
||||
parser.add_argument('--num_npu_core', type=int, required=True, help='Target num_npu_core;')
|
||||
parser.add_argument('--lora_model_path', type=str, help='Absolute path of the lora_model on the Linux board;')
|
||||
parser.add_argument('--prompt_cache_path', type=str, help='Absolute path of the prompt_cache file on the Linux board;')
|
||||
args = parser.parse_args()
|
||||
|
||||
if not (args.target_platform in ["rk3588", "rk3576"]):
|
||||
print("====== Error: 请指定正确的目标平台: rk3588/rk3576 ======")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
if not os.path.exists(args.rkllm_model_path):
|
||||
print("====== Error: 请给出准确的rkllm模型路径,需注意是板端的绝对路径 ======")
|
||||
print("Error: Please provide the correct rkllm model path, and ensure it is the absolute path on the board.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
# 定频设置
|
||||
if not (args.target_platform in ["rk3588", "rk3576"]):
|
||||
print("Error: Please specify the correct target platform: rk3588/rk3576.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
if not (args.num_npu_core in [1, 2, 3]):
|
||||
print("Error: rk3576 supports 1/2 cores, rk3588 supports 1/2/3 cores, please specify the correct number of cores.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
if args.lora_model_path:
|
||||
if not os.path.exists(args.lora_model_path):
|
||||
print("Error: Please provide the correct lora_model path, and advise it is the absolute path on the board.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
if args.prompt_cache_path:
|
||||
if not os.path.exists(args.prompt_cache_path):
|
||||
print("Error: Please provide the correct prompt_cache_file path, and advise it is the absolute path on the board.")
|
||||
sys.stdout.flush()
|
||||
exit()
|
||||
|
||||
# Fix frequency
|
||||
command = "sudo bash fix_freq_{}.sh".format(args.target_platform)
|
||||
subprocess.run(command, shell=True)
|
||||
|
||||
# 设置文件描述符限制
|
||||
# Set resource limit
|
||||
resource.setrlimit(resource.RLIMIT_NOFILE, (102400, 102400))
|
||||
|
||||
# 初始化RKLLM模型
|
||||
# Initialize RKLLM model
|
||||
print("=========init....===========")
|
||||
sys.stdout.flush()
|
||||
target_platform = args.target_platform
|
||||
model_path = args.rkllm_model_path
|
||||
rkllm_model = RKLLM(model_path, target_platform)
|
||||
print("RKLLM初始化成功!")
|
||||
num_npu_core = args.num_npu_core
|
||||
rkllm_model = RKLLM(model_path, num_npu_core, args.lora_model_path, args.prompt_cache_path)
|
||||
print("RKLLM Model has been initialized successfully!")
|
||||
print("==============================")
|
||||
sys.stdout.flush()
|
||||
|
||||
# 记录用户输入的prompt
|
||||
# Record the user's input prompt
|
||||
def get_user_input(user_message, history):
|
||||
history = history + [[user_message, None]]
|
||||
return "", history
|
||||
|
||||
# 获取RKLLM模型的输出并进行流式打印
|
||||
# Retrieve the output from the RKLLM model and print it in a streaming manner
|
||||
def get_RKLLM_output(history):
|
||||
# 链接全局变量,获取回调函数的输出信息
|
||||
# Link global variables to retrieve the output information from the callback function
|
||||
global global_text, global_state
|
||||
global_text = []
|
||||
global_state = -1
|
||||
|
||||
# 创建模型推理的线程
|
||||
# Create a thread for model inference
|
||||
model_thread = threading.Thread(target=rkllm_model.run, args=(history[-1][0],))
|
||||
model_thread.start()
|
||||
|
||||
# history[-1][1]表示当前的输出对话
|
||||
# history[-1][1] represents the current dialogue
|
||||
history[-1][1] = ""
|
||||
|
||||
# 等待模型运行完成,定时检查模型的推理线程
|
||||
# Wait for the model to finish running and periodically check the inference thread of the model
|
||||
model_thread_finished = False
|
||||
while not model_thread_finished:
|
||||
while len(global_text) > 0:
|
||||
history[-1][1] += global_text.pop(0)
|
||||
time.sleep(0.005)
|
||||
# gradio在调用then方法式自动将yield返回的结果推进行输出
|
||||
# Gradio automatically pushes the result returned by the yield statement when calling the then method
|
||||
yield history
|
||||
|
||||
model_thread.join(timeout=0.005)
|
||||
model_thread_finished = not model_thread.is_alive()
|
||||
|
||||
# 创建gradio界面
|
||||
# Create a Gradio interface
|
||||
with gr.Blocks(title="Chat with RKLLM") as chatRKLLM:
|
||||
gr.Markdown("<div align='center'><font size='70'> Chat with RKLLM </font></div>")
|
||||
gr.Markdown("### 在 inputTextBox 输入您的问题,按下 Enter 键,即可与 RKLLM 模型进行对话。")
|
||||
# 创建一个Chatbot组件,用于显示对话历史
|
||||
gr.Markdown("### Enter your question in the inputTextBox and press the Enter key to chat with the RKLLM model.")
|
||||
# Create a Chatbot component to display conversation history
|
||||
rkllmServer = gr.Chatbot(height=600)
|
||||
# 创建一个Textbox组件,让用户输入消息
|
||||
# Create a Textbox component for user message input
|
||||
msg = gr.Textbox(placeholder="Please input your question here...", label="inputTextBox")
|
||||
# 创建一个Button组件,用于清除聊天历史
|
||||
clear = gr.Button("清除")
|
||||
# Create a Button component to clear the chat history.
|
||||
clear = gr.Button("Clear")
|
||||
|
||||
# 将用户输入的消息提交给get_user_input函数,并且立即更新聊天历史
|
||||
# 然后调用get_RKLLM_output函数,进一步更新聊天历史
|
||||
# queue=False参数确保这些更新不会被放入队列,而是立即执行
|
||||
# Submit the user's input message to the get_user_input function and immediately update the chat history.
|
||||
# Then call the get_RKLLM_output function to further update the chat history.
|
||||
# The queue=False parameter ensures that these updates are not queued, but executed immediately.
|
||||
msg.submit(get_user_input, [msg, rkllmServer], [msg, rkllmServer], queue=False).then(get_RKLLM_output, rkllmServer, rkllmServer)
|
||||
# 当点击清除按钮时,执行一个空操作(lambda: None),并且立即清除聊天历史
|
||||
# When the clear button is clicked, perform a no-operation (lambda: None) and immediately clear the chat history.
|
||||
clear.click(lambda: None, None, rkllmServer, queue=False)
|
||||
|
||||
# 启用事件队列系统
|
||||
# Enable the event queue system.
|
||||
chatRKLLM.queue()
|
||||
# 启动Gradio应用程序
|
||||
# Start the Gradio application.
|
||||
chatRKLLM.launch()
|
||||
|
||||
print("====================")
|
||||
print("RKLLM模型推理结束, 释放RKLLM模型资源...")
|
||||
print("RKLLM model inference completed, releasing RKLLM model resources...")
|
||||
rkllm_model.release()
|
||||
print("====================")
|
||||
Binary file not shown.
@@ -1,119 +1,271 @@
|
||||
#ifndef _LLM_H_
|
||||
#define _LLM_H_
|
||||
#ifndef _RKLLM_H_
|
||||
#define _RKLLM_H_
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
typedef void* LLMHandle; /* Handle for an instance of a language model. */
|
||||
/**
|
||||
* @typedef LLMHandle
|
||||
* @brief A handle used to manage and interact with the large language model.
|
||||
*/
|
||||
typedef void* LLMHandle;
|
||||
|
||||
/**
|
||||
* @brief Structure for possible states of an inference call.
|
||||
*
|
||||
* @enum LLMCallState
|
||||
* @brief Describes the possible states of an LLM call.
|
||||
*/
|
||||
typedef enum {
|
||||
LLM_RUN_NORMAL = 0, /* Inference status is normal and inference has not yet finished. */
|
||||
LLM_RUN_FINISH = 1, /* Inference status is normal and inference has finished. */
|
||||
LLM_RUN_ERROR = 2 /* Inference status is abnormal. */
|
||||
RKLLM_RUN_NORMAL = 0, /**< The LLM call is in a normal running state. */
|
||||
RKLLM_RUN_WAITING = 1, /**< The LLM call is waiting for complete UTF-8 encoded character. */
|
||||
RKLLM_RUN_FINISH = 2, /**< The LLM call has finished execution. */
|
||||
RKLLM_RUN_ERROR = 3, /**< An error occurred during the LLM call. */
|
||||
RKLLM_RUN_GET_LAST_HIDDEN_LAYER = 4 /**< Retrieve the last hidden layer during inference. */
|
||||
} LLMCallState;
|
||||
|
||||
/**
|
||||
* @brief Structure for setting up parameters for the language model
|
||||
*
|
||||
* @enum RKLLMInputType
|
||||
* @brief Defines the types of inputs that can be fed into the LLM.
|
||||
*/
|
||||
typedef enum {
|
||||
RKLLM_INPUT_PROMPT = 0, /**< Input is a text prompt. */
|
||||
RKLLM_INPUT_TOKEN = 1, /**< Input is a sequence of tokens. */
|
||||
RKLLM_INPUT_EMBED = 2, /**< Input is an embedding vector. */
|
||||
RKLLM_INPUT_MULTIMODAL = 3, /**< Input is multimodal (e.g., text and image). */
|
||||
} RKLLMInputType;
|
||||
|
||||
/**
|
||||
* @enum RKLLMInferMode
|
||||
* @brief Specifies the inference modes of the LLM.
|
||||
*/
|
||||
typedef enum {
|
||||
RKLLM_INFER_GENERATE = 0, /**< The LLM generates text based on input. */
|
||||
RKLLM_INFER_GET_LAST_HIDDEN_LAYER = 1, /**< The LLM retrieves the last hidden layer for further processing. */
|
||||
} RKLLMInferMode;
|
||||
|
||||
/**
|
||||
* @struct RKLLMExtendParam
|
||||
* @brief The extend parameters for configuring an LLM instance.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* model_path; /* Path where the model file is located. */
|
||||
int32_t num_npu_core; /* Number of NPU cores used for model inference. */
|
||||
int32_t max_context_len; /* Maximum size of the context. */
|
||||
int32_t max_new_tokens; /* Maximum number of tokens to generate during model inference. */
|
||||
int32_t top_k; /* The number of highest probability tokens to consider for generation. */
|
||||
float top_p; /* Nucleus sampling: cumulative probability cutoff to use for token selection. */
|
||||
float temperature; /* Hyperparameter to control the randomness of predictions by scaling the logits before applying softmax. */
|
||||
float repeat_penalty; /* Penalty applied to the logits of previously generated tokens, helps prevent repetitive or monotonic text. */
|
||||
float frequency_penalty; /* Penalty for repeating the same word or phrase, reducing the likelihood of repeated content. */
|
||||
float presence_penalty; /* Penalty or reward for introducing new tokens into the generated text. */
|
||||
int32_t mirostat; /* Enables mirostat algorithm, where 0 = off, 1 = use mirostat algorithm, 2 = use mirostat 2.0 algorithm. */
|
||||
float mirostat_tau; /* Target entropy (perplexity) for mirostat algorithm, setting the desired complexity of the generated text. */
|
||||
float mirostat_eta; /* Learning rate for the mirostat algorithm. */
|
||||
bool logprobs; /* Whether to return the log probabilities for each output token along with their token ids. */
|
||||
int32_t top_logprobs; /* The number of top tokens for which to return log probabilities, along with their token ids. */
|
||||
bool use_gpu; /* Flag to indicate whether to use GPU for inference. */
|
||||
int32_t base_domain_id; /**< base_domain_id */
|
||||
uint8_t reserved[112]; /**< reserved */
|
||||
} RKLLMExtendParam;
|
||||
|
||||
/**
|
||||
* @struct RKLLMParam
|
||||
* @brief Defines the parameters for configuring an LLM instance.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* model_path; /**< Path to the model file. */
|
||||
int32_t max_context_len; /**< Maximum number of tokens in the context window. */
|
||||
int32_t max_new_tokens; /**< Maximum number of new tokens to generate. */
|
||||
int32_t top_k; /**< Top-K sampling parameter for token generation. */
|
||||
float top_p; /**< Top-P (nucleus) sampling parameter. */
|
||||
float temperature; /**< Sampling temperature, affecting the randomness of token selection. */
|
||||
float repeat_penalty; /**< Penalty for repeating tokens in generation. */
|
||||
float frequency_penalty; /**< Penalizes frequent tokens during generation. */
|
||||
float presence_penalty; /**< Penalizes tokens based on their presence in the input. */
|
||||
int32_t mirostat; /**< Mirostat sampling strategy flag (0 to disable). */
|
||||
float mirostat_tau; /**< Tau parameter for Mirostat sampling. */
|
||||
float mirostat_eta; /**< Eta parameter for Mirostat sampling. */
|
||||
bool skip_special_token; /**< Whether to skip special tokens during generation. */
|
||||
bool is_async; /**< Whether to run inference asynchronously. */
|
||||
const char* img_start; /**< Starting position of an image in multimodal input. */
|
||||
const char* img_end; /**< Ending position of an image in multimodal input. */
|
||||
const char* img_content; /**< Pointer to the image content. */
|
||||
RKLLMExtendParam extend_param; /**< Extend parameters. */
|
||||
} RKLLMParam;
|
||||
|
||||
/**
|
||||
* @brief Structure representing a token with its associated log probability.
|
||||
*
|
||||
* @struct RKLLMLoraAdapter
|
||||
* @brief Defines parameters for a Lora adapter used in model fine-tuning.
|
||||
*/
|
||||
typedef struct {
|
||||
float logprob; /* Log probability corresponding to the token ID. */
|
||||
int id; /* Token ID. */
|
||||
} Token;
|
||||
const char* lora_adapter_path; /**< Path to the Lora adapter file. */
|
||||
const char* lora_adapter_name; /**< Name of the Lora adapter. */
|
||||
float scale; /**< Scaling factor for applying the Lora adapter. */
|
||||
} RKLLMLoraAdapter;
|
||||
|
||||
/**
|
||||
* @brief Structure to hold the results from the language model inference, including text and token details.
|
||||
*
|
||||
* @struct RKLLMEmbedInput
|
||||
* @brief Represents an embedding input to the LLM.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* text; /* Decoded text from the inference output. */
|
||||
Token* tokens; /* Array of Token structures, each containing a log probability and a token ID. */
|
||||
int num; /* Number of top tokens returned, typically those with the highest probabilities. */
|
||||
float* embed; /**< Pointer to the embedding vector (of size n_tokens * n_embed). */
|
||||
size_t n_tokens; /**< Number of tokens represented in the embedding. */
|
||||
} RKLLMEmbedInput;
|
||||
|
||||
/**
|
||||
* @struct RKLLMTokenInput
|
||||
* @brief Represents token input to the LLM.
|
||||
*/
|
||||
typedef struct {
|
||||
int32_t* input_ids; /**< Array of token IDs. */
|
||||
size_t n_tokens; /**< Number of tokens in the input. */
|
||||
} RKLLMTokenInput;
|
||||
|
||||
/**
|
||||
* @struct RKLLMMultiModelInput
|
||||
* @brief Represents multimodal input (e.g., text and image).
|
||||
*/
|
||||
typedef struct {
|
||||
char* prompt; /**< Text prompt input. */
|
||||
float* image_embed; /**< Embedding of the image (of size n_image_tokens * n_image_embed). */
|
||||
size_t n_image_tokens; /**< Number of image tokens. */
|
||||
} RKLLMMultiModelInput;
|
||||
|
||||
/**
|
||||
* @struct RKLLMInput
|
||||
* @brief Represents different types of input to the LLM via a union.
|
||||
*/
|
||||
typedef struct {
|
||||
RKLLMInputType input_type; /**< Specifies the type of input provided (e.g., prompt, token, embed, multimodal). */
|
||||
union {
|
||||
const char* prompt_input; /**< Text prompt input if input_type is RKLLM_INPUT_PROMPT. */
|
||||
RKLLMEmbedInput embed_input; /**< Embedding input if input_type is RKLLM_INPUT_EMBED. */
|
||||
RKLLMTokenInput token_input; /**< Token input if input_type is RKLLM_INPUT_TOKEN. */
|
||||
RKLLMMultiModelInput multimodal_input; /**< Multimodal input if input_type is RKLLM_INPUT_MULTIMODAL. */
|
||||
};
|
||||
} RKLLMInput;
|
||||
|
||||
/**
|
||||
* @struct RKLLMLoraParam
|
||||
* @brief Structure defining parameters for Lora adapters.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* lora_adapter_name; /**< Name of the Lora adapter. */
|
||||
} RKLLMLoraParam;
|
||||
|
||||
/**
|
||||
* @struct RKLLMPromptCacheParam
|
||||
* @brief Structure to define parameters for caching prompts.
|
||||
*/
|
||||
typedef struct {
|
||||
int save_prompt_cache; /**< Flag to indicate whether to save the prompt cache (0 = don't save, 1 = save). */
|
||||
const char* prompt_cache_path; /**< Path to the prompt cache file. */
|
||||
} RKLLMPromptCacheParam;
|
||||
|
||||
/**
|
||||
* @struct RKLLMInferParam
|
||||
* @brief Structure for defining parameters during inference.
|
||||
*/
|
||||
typedef struct {
|
||||
RKLLMInferMode mode; /**< Inference mode (e.g., generate or get last hidden layer). */
|
||||
RKLLMLoraParam* lora_params; /**< Pointer to Lora adapter parameters. */
|
||||
RKLLMPromptCacheParam* prompt_cache_params; /**< Pointer to prompt cache parameters. */
|
||||
} RKLLMInferParam;
|
||||
|
||||
/**
|
||||
* @struct RKLLMResultLastHiddenLayer
|
||||
* @brief Structure to hold the hidden states from the last layer.
|
||||
*/
|
||||
typedef struct {
|
||||
const float* hidden_states; /**< Pointer to the hidden states (of size num_tokens * embd_size). */
|
||||
int embd_size; /**< Size of the embedding vector. */
|
||||
int num_tokens; /**< Number of tokens for which hidden states are stored. */
|
||||
} RKLLMResultLastHiddenLayer;
|
||||
|
||||
/**
|
||||
* @struct RKLLMResult
|
||||
* @brief Structure to represent the result of LLM inference.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* text; /**< Generated text result. */
|
||||
int32_t token_id; /**< ID of the generated token. */
|
||||
RKLLMResultLastHiddenLayer last_hidden_layer; /**< Hidden states of the last layer (if requested). */
|
||||
} RKLLMResult;
|
||||
|
||||
/**
|
||||
* @brief Callback function for handling inference results.
|
||||
*
|
||||
* @param result A pointer to an RKLLMResult struct containing the inference results.
|
||||
* @param userdata A pointer to user-defined function or null if no user function was provided.
|
||||
* @param state The state of the inference process, indicating success, failure, or completion.
|
||||
* @typedef LLMResultCallback
|
||||
* @brief Callback function to handle LLM results.
|
||||
* @param result Pointer to the LLM result.
|
||||
* @param userdata Pointer to user data for the callback.
|
||||
* @param state State of the LLM call (e.g., finished, error).
|
||||
*/
|
||||
typedef void(*LLMResultCallback)(RKLLMResult* result, void* userdata, LLMCallState state);
|
||||
|
||||
/**
|
||||
* @brief Initializes RKLLMParam with default settings.
|
||||
*
|
||||
* @return RKLLMParam An RKLLMParam struct with default values set.
|
||||
* @brief Creates a default RKLLMParam structure with preset values.
|
||||
* @return A default RKLLMParam structure.
|
||||
*/
|
||||
RKLLMParam rkllm_createDefaultParam();
|
||||
|
||||
/**
|
||||
* @brief Initializes the model with specified parameters.
|
||||
*
|
||||
* @param handle Pointer to a handle for the language model, which will be initialized by this function.
|
||||
* @param param An RKLLMParam struct containing all the parameters needed for the model.
|
||||
* @param callback A function pointer to the callback that handles the results of the inference.
|
||||
* @return int Returns 0 on success, or a negative error code on failure.
|
||||
* @brief Initializes the LLM with the given parameters.
|
||||
* @param handle Pointer to the LLM handle.
|
||||
* @param param Configuration parameters for the LLM.
|
||||
* @param callback Callback function to handle LLM results.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_init(LLMHandle* handle, RKLLMParam param, LLMResultCallback callback);
|
||||
int rkllm_init(LLMHandle* handle, RKLLMParam* param, LLMResultCallback callback);
|
||||
|
||||
/**
|
||||
* @brief Releases the model resources.
|
||||
*
|
||||
* @param handle The handle to the language model to be destroyed.
|
||||
* @return int Returns 0 on successful release, or a negative error code if an error occurs.
|
||||
* @brief Loads a Lora adapter into the LLM.
|
||||
* @param handle LLM handle.
|
||||
* @param lora_adapter Pointer to the Lora adapter structure.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_load_lora(LLMHandle handle, RKLLMLoraAdapter* lora_adapter);
|
||||
|
||||
/**
|
||||
* @brief Loads a prompt cache from a file.
|
||||
* @param handle LLM handle.
|
||||
* @param prompt_cache_path Path to the prompt cache file.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_load_prompt_cache(LLMHandle handle, const char* prompt_cache_path);
|
||||
|
||||
/**
|
||||
* @brief Releases the prompt cache from memory.
|
||||
* @param handle LLM handle.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_release_prompt_cache(LLMHandle handle);
|
||||
|
||||
/**
|
||||
* @brief Destroys the LLM instance and releases resources.
|
||||
* @param handle LLM handle.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_destroy(LLMHandle handle);
|
||||
|
||||
/**
|
||||
* @brief Runs model inference on the given prompt.
|
||||
*
|
||||
* @param handle The handle to the initialized language model.
|
||||
* @param prompt The text prompt on which to perform inference.
|
||||
* @param userdata Optional user-defined function that will be passed to the callback.
|
||||
* @return int Returns 0 on success, or a negative error code if an error occurs during inference.
|
||||
* @brief Runs an LLM inference task synchronously.
|
||||
* @param handle LLM handle.
|
||||
* @param rkllm_input Input data for the LLM.
|
||||
* @param rkllm_infer_params Parameters for the inference task.
|
||||
* @param userdata Pointer to user data for the callback.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_run(LLMHandle handle, const char* prompt, void* userdata);
|
||||
int rkllm_run(LLMHandle handle, RKLLMInput* rkllm_input, RKLLMInferParam* rkllm_infer_params, void* userdata);
|
||||
|
||||
/**
|
||||
* @brief Aborts the current inference process.
|
||||
*
|
||||
* @param handle The handle to the language model whose inference is to be aborted.
|
||||
* @return int Returns 0 if the process is successfully aborted, or a negative error code
|
||||
* if no process was running or if the abort fails.
|
||||
* @brief Runs an LLM inference task asynchronously.
|
||||
* @param handle LLM handle.
|
||||
* @param rkllm_input Input data for the LLM.
|
||||
* @param rkllm_infer_params Parameters for the inference task.
|
||||
* @param userdata Pointer to user data for the callback.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_run_async(LLMHandle handle, RKLLMInput* rkllm_input, RKLLMInferParam* rkllm_infer_params, void* userdata);
|
||||
|
||||
/**
|
||||
* @brief Aborts an ongoing LLM task.
|
||||
* @param handle LLM handle.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_abort(LLMHandle handle);
|
||||
|
||||
/**
|
||||
* @brief Checks if an LLM task is currently running.
|
||||
* @param handle LLM handle.
|
||||
* @return Status code (0 if a task is running, non-zero for otherwise).
|
||||
*/
|
||||
int rkllm_is_running(LLMHandle handle);
|
||||
|
||||
#ifdef __cplusplus
|
||||
} //extern "C"
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif
|
||||
#endif
|
||||
|
||||
Binary file not shown.
@@ -1,119 +1,271 @@
|
||||
#ifndef _LLM_H_
|
||||
#define _LLM_H_
|
||||
#ifndef _RKLLM_H_
|
||||
#define _RKLLM_H_
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
typedef void* LLMHandle; /* Handle for an instance of a language model. */
|
||||
/**
|
||||
* @typedef LLMHandle
|
||||
* @brief A handle used to manage and interact with the large language model.
|
||||
*/
|
||||
typedef void* LLMHandle;
|
||||
|
||||
/**
|
||||
* @brief Structure for possible states of an inference call.
|
||||
*
|
||||
* @enum LLMCallState
|
||||
* @brief Describes the possible states of an LLM call.
|
||||
*/
|
||||
typedef enum {
|
||||
LLM_RUN_NORMAL = 0, /* Inference status is normal and inference has not yet finished. */
|
||||
LLM_RUN_FINISH = 1, /* Inference status is normal and inference has finished. */
|
||||
LLM_RUN_ERROR = 2 /* Inference status is abnormal. */
|
||||
RKLLM_RUN_NORMAL = 0, /**< The LLM call is in a normal running state. */
|
||||
RKLLM_RUN_WAITING = 1, /**< The LLM call is waiting for complete UTF-8 encoded character. */
|
||||
RKLLM_RUN_FINISH = 2, /**< The LLM call has finished execution. */
|
||||
RKLLM_RUN_ERROR = 3, /**< An error occurred during the LLM call. */
|
||||
RKLLM_RUN_GET_LAST_HIDDEN_LAYER = 4 /**< Retrieve the last hidden layer during inference. */
|
||||
} LLMCallState;
|
||||
|
||||
/**
|
||||
* @brief Structure for setting up parameters for the language model
|
||||
*
|
||||
* @enum RKLLMInputType
|
||||
* @brief Defines the types of inputs that can be fed into the LLM.
|
||||
*/
|
||||
typedef enum {
|
||||
RKLLM_INPUT_PROMPT = 0, /**< Input is a text prompt. */
|
||||
RKLLM_INPUT_TOKEN = 1, /**< Input is a sequence of tokens. */
|
||||
RKLLM_INPUT_EMBED = 2, /**< Input is an embedding vector. */
|
||||
RKLLM_INPUT_MULTIMODAL = 3, /**< Input is multimodal (e.g., text and image). */
|
||||
} RKLLMInputType;
|
||||
|
||||
/**
|
||||
* @enum RKLLMInferMode
|
||||
* @brief Specifies the inference modes of the LLM.
|
||||
*/
|
||||
typedef enum {
|
||||
RKLLM_INFER_GENERATE = 0, /**< The LLM generates text based on input. */
|
||||
RKLLM_INFER_GET_LAST_HIDDEN_LAYER = 1, /**< The LLM retrieves the last hidden layer for further processing. */
|
||||
} RKLLMInferMode;
|
||||
|
||||
/**
|
||||
* @struct RKLLMExtendParam
|
||||
* @brief The extend parameters for configuring an LLM instance.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* model_path; /* Path where the model file is located. */
|
||||
int32_t num_npu_core; /* Number of NPU cores used for model inference. */
|
||||
int32_t max_context_len; /* Maximum size of the context. */
|
||||
int32_t max_new_tokens; /* Maximum number of tokens to generate during model inference. */
|
||||
int32_t top_k; /* The number of highest probability tokens to consider for generation. */
|
||||
float top_p; /* Nucleus sampling: cumulative probability cutoff to use for token selection. */
|
||||
float temperature; /* Hyperparameter to control the randomness of predictions by scaling the logits before applying softmax. */
|
||||
float repeat_penalty; /* Penalty applied to the logits of previously generated tokens, helps prevent repetitive or monotonic text. */
|
||||
float frequency_penalty; /* Penalty for repeating the same word or phrase, reducing the likelihood of repeated content. */
|
||||
float presence_penalty; /* Penalty or reward for introducing new tokens into the generated text. */
|
||||
int32_t mirostat; /* Enables mirostat algorithm, where 0 = off, 1 = use mirostat algorithm, 2 = use mirostat 2.0 algorithm. */
|
||||
float mirostat_tau; /* Target entropy (perplexity) for mirostat algorithm, setting the desired complexity of the generated text. */
|
||||
float mirostat_eta; /* Learning rate for the mirostat algorithm. */
|
||||
bool logprobs; /* Whether to return the log probabilities for each output token along with their token ids. */
|
||||
int32_t top_logprobs; /* The number of top tokens for which to return log probabilities, along with their token ids. */
|
||||
bool use_gpu; /* Flag to indicate whether to use GPU for inference. */
|
||||
int32_t base_domain_id; /**< base_domain_id */
|
||||
uint8_t reserved[112]; /**< reserved */
|
||||
} RKLLMExtendParam;
|
||||
|
||||
/**
|
||||
* @struct RKLLMParam
|
||||
* @brief Defines the parameters for configuring an LLM instance.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* model_path; /**< Path to the model file. */
|
||||
int32_t max_context_len; /**< Maximum number of tokens in the context window. */
|
||||
int32_t max_new_tokens; /**< Maximum number of new tokens to generate. */
|
||||
int32_t top_k; /**< Top-K sampling parameter for token generation. */
|
||||
float top_p; /**< Top-P (nucleus) sampling parameter. */
|
||||
float temperature; /**< Sampling temperature, affecting the randomness of token selection. */
|
||||
float repeat_penalty; /**< Penalty for repeating tokens in generation. */
|
||||
float frequency_penalty; /**< Penalizes frequent tokens during generation. */
|
||||
float presence_penalty; /**< Penalizes tokens based on their presence in the input. */
|
||||
int32_t mirostat; /**< Mirostat sampling strategy flag (0 to disable). */
|
||||
float mirostat_tau; /**< Tau parameter for Mirostat sampling. */
|
||||
float mirostat_eta; /**< Eta parameter for Mirostat sampling. */
|
||||
bool skip_special_token; /**< Whether to skip special tokens during generation. */
|
||||
bool is_async; /**< Whether to run inference asynchronously. */
|
||||
const char* img_start; /**< Starting position of an image in multimodal input. */
|
||||
const char* img_end; /**< Ending position of an image in multimodal input. */
|
||||
const char* img_content; /**< Pointer to the image content. */
|
||||
RKLLMExtendParam extend_param; /**< Extend parameters. */
|
||||
} RKLLMParam;
|
||||
|
||||
/**
|
||||
* @brief Structure representing a token with its associated log probability.
|
||||
*
|
||||
* @struct RKLLMLoraAdapter
|
||||
* @brief Defines parameters for a Lora adapter used in model fine-tuning.
|
||||
*/
|
||||
typedef struct {
|
||||
float logprob; /* Log probability corresponding to the token ID. */
|
||||
int id; /* Token ID. */
|
||||
} Token;
|
||||
const char* lora_adapter_path; /**< Path to the Lora adapter file. */
|
||||
const char* lora_adapter_name; /**< Name of the Lora adapter. */
|
||||
float scale; /**< Scaling factor for applying the Lora adapter. */
|
||||
} RKLLMLoraAdapter;
|
||||
|
||||
/**
|
||||
* @brief Structure to hold the results from the language model inference, including text and token details.
|
||||
*
|
||||
* @struct RKLLMEmbedInput
|
||||
* @brief Represents an embedding input to the LLM.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* text; /* Decoded text from the inference output. */
|
||||
Token* tokens; /* Array of Token structures, each containing a log probability and a token ID. */
|
||||
int num; /* Number of top tokens returned, typically those with the highest probabilities. */
|
||||
float* embed; /**< Pointer to the embedding vector (of size n_tokens * n_embed). */
|
||||
size_t n_tokens; /**< Number of tokens represented in the embedding. */
|
||||
} RKLLMEmbedInput;
|
||||
|
||||
/**
|
||||
* @struct RKLLMTokenInput
|
||||
* @brief Represents token input to the LLM.
|
||||
*/
|
||||
typedef struct {
|
||||
int32_t* input_ids; /**< Array of token IDs. */
|
||||
size_t n_tokens; /**< Number of tokens in the input. */
|
||||
} RKLLMTokenInput;
|
||||
|
||||
/**
|
||||
* @struct RKLLMMultiModelInput
|
||||
* @brief Represents multimodal input (e.g., text and image).
|
||||
*/
|
||||
typedef struct {
|
||||
char* prompt; /**< Text prompt input. */
|
||||
float* image_embed; /**< Embedding of the image (of size n_image_tokens * n_image_embed). */
|
||||
size_t n_image_tokens; /**< Number of image tokens. */
|
||||
} RKLLMMultiModelInput;
|
||||
|
||||
/**
|
||||
* @struct RKLLMInput
|
||||
* @brief Represents different types of input to the LLM via a union.
|
||||
*/
|
||||
typedef struct {
|
||||
RKLLMInputType input_type; /**< Specifies the type of input provided (e.g., prompt, token, embed, multimodal). */
|
||||
union {
|
||||
const char* prompt_input; /**< Text prompt input if input_type is RKLLM_INPUT_PROMPT. */
|
||||
RKLLMEmbedInput embed_input; /**< Embedding input if input_type is RKLLM_INPUT_EMBED. */
|
||||
RKLLMTokenInput token_input; /**< Token input if input_type is RKLLM_INPUT_TOKEN. */
|
||||
RKLLMMultiModelInput multimodal_input; /**< Multimodal input if input_type is RKLLM_INPUT_MULTIMODAL. */
|
||||
};
|
||||
} RKLLMInput;
|
||||
|
||||
/**
|
||||
* @struct RKLLMLoraParam
|
||||
* @brief Structure defining parameters for Lora adapters.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* lora_adapter_name; /**< Name of the Lora adapter. */
|
||||
} RKLLMLoraParam;
|
||||
|
||||
/**
|
||||
* @struct RKLLMPromptCacheParam
|
||||
* @brief Structure to define parameters for caching prompts.
|
||||
*/
|
||||
typedef struct {
|
||||
int save_prompt_cache; /**< Flag to indicate whether to save the prompt cache (0 = don't save, 1 = save). */
|
||||
const char* prompt_cache_path; /**< Path to the prompt cache file. */
|
||||
} RKLLMPromptCacheParam;
|
||||
|
||||
/**
|
||||
* @struct RKLLMInferParam
|
||||
* @brief Structure for defining parameters during inference.
|
||||
*/
|
||||
typedef struct {
|
||||
RKLLMInferMode mode; /**< Inference mode (e.g., generate or get last hidden layer). */
|
||||
RKLLMLoraParam* lora_params; /**< Pointer to Lora adapter parameters. */
|
||||
RKLLMPromptCacheParam* prompt_cache_params; /**< Pointer to prompt cache parameters. */
|
||||
} RKLLMInferParam;
|
||||
|
||||
/**
|
||||
* @struct RKLLMResultLastHiddenLayer
|
||||
* @brief Structure to hold the hidden states from the last layer.
|
||||
*/
|
||||
typedef struct {
|
||||
const float* hidden_states; /**< Pointer to the hidden states (of size num_tokens * embd_size). */
|
||||
int embd_size; /**< Size of the embedding vector. */
|
||||
int num_tokens; /**< Number of tokens for which hidden states are stored. */
|
||||
} RKLLMResultLastHiddenLayer;
|
||||
|
||||
/**
|
||||
* @struct RKLLMResult
|
||||
* @brief Structure to represent the result of LLM inference.
|
||||
*/
|
||||
typedef struct {
|
||||
const char* text; /**< Generated text result. */
|
||||
int32_t token_id; /**< ID of the generated token. */
|
||||
RKLLMResultLastHiddenLayer last_hidden_layer; /**< Hidden states of the last layer (if requested). */
|
||||
} RKLLMResult;
|
||||
|
||||
/**
|
||||
* @brief Callback function for handling inference results.
|
||||
*
|
||||
* @param result A pointer to an RKLLMResult struct containing the inference results.
|
||||
* @param userdata A pointer to user-defined function or null if no user function was provided.
|
||||
* @param state The state of the inference process, indicating success, failure, or completion.
|
||||
* @typedef LLMResultCallback
|
||||
* @brief Callback function to handle LLM results.
|
||||
* @param result Pointer to the LLM result.
|
||||
* @param userdata Pointer to user data for the callback.
|
||||
* @param state State of the LLM call (e.g., finished, error).
|
||||
*/
|
||||
typedef void(*LLMResultCallback)(RKLLMResult* result, void* userdata, LLMCallState state);
|
||||
|
||||
/**
|
||||
* @brief Initializes RKLLMParam with default settings.
|
||||
*
|
||||
* @return RKLLMParam An RKLLMParam struct with default values set.
|
||||
* @brief Creates a default RKLLMParam structure with preset values.
|
||||
* @return A default RKLLMParam structure.
|
||||
*/
|
||||
RKLLMParam rkllm_createDefaultParam();
|
||||
|
||||
/**
|
||||
* @brief Initializes the model with specified parameters.
|
||||
*
|
||||
* @param handle Pointer to a handle for the language model, which will be initialized by this function.
|
||||
* @param param An RKLLMParam struct containing all the parameters needed for the model.
|
||||
* @param callback A function pointer to the callback that handles the results of the inference.
|
||||
* @return int Returns 0 on success, or a negative error code on failure.
|
||||
* @brief Initializes the LLM with the given parameters.
|
||||
* @param handle Pointer to the LLM handle.
|
||||
* @param param Configuration parameters for the LLM.
|
||||
* @param callback Callback function to handle LLM results.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_init(LLMHandle* handle, RKLLMParam param, LLMResultCallback callback);
|
||||
int rkllm_init(LLMHandle* handle, RKLLMParam* param, LLMResultCallback callback);
|
||||
|
||||
/**
|
||||
* @brief Releases the model resources.
|
||||
*
|
||||
* @param handle The handle to the language model to be destroyed.
|
||||
* @return int Returns 0 on successful release, or a negative error code if an error occurs.
|
||||
* @brief Loads a Lora adapter into the LLM.
|
||||
* @param handle LLM handle.
|
||||
* @param lora_adapter Pointer to the Lora adapter structure.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_load_lora(LLMHandle handle, RKLLMLoraAdapter* lora_adapter);
|
||||
|
||||
/**
|
||||
* @brief Loads a prompt cache from a file.
|
||||
* @param handle LLM handle.
|
||||
* @param prompt_cache_path Path to the prompt cache file.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_load_prompt_cache(LLMHandle handle, const char* prompt_cache_path);
|
||||
|
||||
/**
|
||||
* @brief Releases the prompt cache from memory.
|
||||
* @param handle LLM handle.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_release_prompt_cache(LLMHandle handle);
|
||||
|
||||
/**
|
||||
* @brief Destroys the LLM instance and releases resources.
|
||||
* @param handle LLM handle.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_destroy(LLMHandle handle);
|
||||
|
||||
/**
|
||||
* @brief Runs model inference on the given prompt.
|
||||
*
|
||||
* @param handle The handle to the initialized language model.
|
||||
* @param prompt The text prompt on which to perform inference.
|
||||
* @param userdata Optional user-defined function that will be passed to the callback.
|
||||
* @return int Returns 0 on success, or a negative error code if an error occurs during inference.
|
||||
* @brief Runs an LLM inference task synchronously.
|
||||
* @param handle LLM handle.
|
||||
* @param rkllm_input Input data for the LLM.
|
||||
* @param rkllm_infer_params Parameters for the inference task.
|
||||
* @param userdata Pointer to user data for the callback.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_run(LLMHandle handle, const char* prompt, void* userdata);
|
||||
int rkllm_run(LLMHandle handle, RKLLMInput* rkllm_input, RKLLMInferParam* rkllm_infer_params, void* userdata);
|
||||
|
||||
/**
|
||||
* @brief Aborts the current inference process.
|
||||
*
|
||||
* @param handle The handle to the language model whose inference is to be aborted.
|
||||
* @return int Returns 0 if the process is successfully aborted, or a negative error code
|
||||
* if no process was running or if the abort fails.
|
||||
* @brief Runs an LLM inference task asynchronously.
|
||||
* @param handle LLM handle.
|
||||
* @param rkllm_input Input data for the LLM.
|
||||
* @param rkllm_infer_params Parameters for the inference task.
|
||||
* @param userdata Pointer to user data for the callback.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_run_async(LLMHandle handle, RKLLMInput* rkllm_input, RKLLMInferParam* rkllm_infer_params, void* userdata);
|
||||
|
||||
/**
|
||||
* @brief Aborts an ongoing LLM task.
|
||||
* @param handle LLM handle.
|
||||
* @return Status code (0 for success, non-zero for failure).
|
||||
*/
|
||||
int rkllm_abort(LLMHandle handle);
|
||||
|
||||
/**
|
||||
* @brief Checks if an LLM task is currently running.
|
||||
* @param handle LLM handle.
|
||||
* @return Status code (0 if a task is running, non-zero for otherwise).
|
||||
*/
|
||||
int rkllm_is_running(LLMHandle handle);
|
||||
|
||||
#ifdef __cplusplus
|
||||
} //extern "C"
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif
|
||||
#endif
|
||||
|
||||
@@ -1,29 +0,0 @@
|
||||
from rkllm.api import RKLLM
|
||||
|
||||
'''
|
||||
https://huggingface.co/Qwen/Qwen-1_8B-Chat
|
||||
Download the Qwen model from the above website.
|
||||
'''
|
||||
|
||||
modelpath = '/path/to/your/model'
|
||||
llm = RKLLM()
|
||||
|
||||
# Load model
|
||||
ret = llm.load_huggingface(model = modelpath)
|
||||
if ret != 0:
|
||||
print('Load model failed!')
|
||||
exit(ret)
|
||||
|
||||
# Build model
|
||||
ret = llm.build(do_quantization=True, optimization_level=1, quantized_dtype='w8a8', target_platform='rk3588')
|
||||
if ret != 0:
|
||||
print('Build model failed!')
|
||||
exit(ret)
|
||||
|
||||
# Export rknn model
|
||||
ret = llm.export_rkllm("./qwen.rkllm")
|
||||
if ret != 0:
|
||||
print('Export model failed!')
|
||||
exit(ret)
|
||||
|
||||
|
||||
86
rkllm-toolkit/examples/test.py
Executable file
86
rkllm-toolkit/examples/test.py
Executable file
@@ -0,0 +1,86 @@
|
||||
from rkllm.api import RKLLM
|
||||
from datasets import load_dataset
|
||||
from transformers import AutoTokenizer
|
||||
from tqdm import tqdm
|
||||
import torch
|
||||
from torch import nn
|
||||
import os
|
||||
# os.environ['CUDA_VISIBLE_DEVICES']='1'
|
||||
|
||||
'''
|
||||
https://huggingface.co/Qwen/Qwen-1_8B-Chat
|
||||
从上面网址中下载Qwen模型
|
||||
'''
|
||||
|
||||
modelpath = './path/to/model'
|
||||
# modelpath = "./path/to/Qwen-1.8B-F16.gguf"
|
||||
llm = RKLLM()
|
||||
|
||||
# Load model
|
||||
# Use 'export CUDA_VISIBLE_DEVICES=2' to specify GPU device
|
||||
# options ['cpu', 'cuda']
|
||||
ret = llm.load_huggingface(model=modelpath, model_lora = None, device='cpu')
|
||||
# ret = llm.load_gguf(model = modelpath)
|
||||
if ret != 0:
|
||||
print('Load model failed!')
|
||||
exit(ret)
|
||||
|
||||
# Build model
|
||||
dataset = "./data_quant.json"
|
||||
# Json file format, please note to add prompt in the input,like this:
|
||||
# [{"input":"Human: 你好!\nAssistant: ", "target": "你好!我是人工智能助手KK!"},...]
|
||||
|
||||
qparams = None
|
||||
# qparams = 'gdq.qparams' # Use extra_qparams
|
||||
ret = llm.build(do_quantization=True, optimization_level=1, quantized_dtype='w8a8',
|
||||
quantized_algorithm='normal', target_platform='rk3588', num_npu_core=3, extra_qparams=qparams, dataset=dataset)
|
||||
|
||||
if ret != 0:
|
||||
print('Build model failed!')
|
||||
exit(ret)
|
||||
|
||||
# Evaluate Accuracy
|
||||
def eval_wikitext(llm):
|
||||
seqlen = 512
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
modelpath, trust_remote_code=True)
|
||||
# Dataset download link:
|
||||
# https://huggingface.co/datasets/Salesforce/wikitext/tree/main/wikitext-2-raw-v1
|
||||
testenc = load_dataset(
|
||||
"parquet", data_files='./wikitext/wikitext-2-raw-1/test-00000-of-00001.parquet', split='train')
|
||||
testenc = tokenizer("\n\n".join(
|
||||
testenc['text']), return_tensors="pt").input_ids
|
||||
nsamples = testenc.numel() // seqlen
|
||||
nlls = []
|
||||
for i in tqdm(range(nsamples), desc="eval_wikitext: "):
|
||||
batch = testenc[:, (i * seqlen): ((i + 1) * seqlen)]
|
||||
inputs = {"input_ids": batch}
|
||||
lm_logits = llm.get_logits(inputs)
|
||||
if lm_logits is None:
|
||||
print("get logits failed!")
|
||||
return
|
||||
shift_logits = lm_logits[:, :-1, :]
|
||||
shift_labels = batch[:, 1:].to(lm_logits.device)
|
||||
loss_fct = nn.CrossEntropyLoss().to(lm_logits.device)
|
||||
loss = loss_fct(
|
||||
shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
|
||||
neg_log_likelihood = loss.float() * seqlen
|
||||
nlls.append(neg_log_likelihood)
|
||||
ppl = torch.exp(torch.stack(nlls).sum() / (nsamples * seqlen))
|
||||
print(f'wikitext-2-raw-1-test ppl: {round(ppl.item(), 2)}')
|
||||
|
||||
# eval_wikitext(llm)
|
||||
|
||||
|
||||
# Chat with model
|
||||
messages = "<|im_start|>system You are a helpful assistant.<|im_end|><|im_start|>user你好!\n<|im_end|><|im_start|>assistant"
|
||||
kwargs = {"max_length": 128, "top_k": 1, "top_p": 0.8,
|
||||
"temperature": 0.8, "do_sample": True, "repetition_penalty": 1.1}
|
||||
# print(llm.chat_model(messages, kwargs))
|
||||
|
||||
|
||||
# Export rkllm model
|
||||
ret = llm.export_rkllm("./qwen.rkllm")
|
||||
if ret != 0:
|
||||
print('Export model failed!')
|
||||
exit(ret)
|
||||
@@ -1 +1 @@
|
||||
cf82f9756844793bbcd132aae58d0b81 rkllm_toolkit-1.0.1-cp38-cp38-linux_x86_64.whl
|
||||
1e69c83aa9b718a5b75e368aa51d0ac4 rkllm_toolkit-1.1.0-cp38-cp38-linux_x86_64.whl
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
rknpu-driver/rknpu_driver_0.9.8_20241009.tar.bz2
Normal file
BIN
rknpu-driver/rknpu_driver_0.9.8_20241009.tar.bz2
Normal file
Binary file not shown.
Reference in New Issue
Block a user