13.板端Qwen3-VL记忆实现
1.背景介绍
Qwen3-VL 的语言主干是标准 Transformer Decoder,天生依靠 KV Cache 存储历史图文 token,记住本轮对话所有图片、问答上下文,属于模型推理内置能力
但是要实现断电、重启后仍具有之前的记忆,需要自己实现
2.方案选择
两个方案:
1.轻量化本地向量记忆(边缘首选)用 SQLite 存对话文本 + 图片特征向量,FAISS 做本地检索;
新提问时先检索历史对话,把相关图文拼接进 prompt 喂给 Qwen3-VL,模拟 “长期记忆”,内存占用极低,适配 RK3588 嵌入式 Linux。
2.纯文本历史落盘每次对话把完整 prompt(图文 token 文本描述)存入本地文件,新开会话时读取拼接进输入,简单但长历史会拉长 prompt、变慢。
本文以方案1为主进行说明
方案1完整流程
怎么筛选要拼接的记录(标准边缘轻量方案)?
步骤 1:每条历史对话生成「记忆向量」
记忆单元 = {
id: 自增ID
text: 本轮完整图文对话(用户提问+模型回答+图片描述)
vec: text embedding向量(128/384维轻量模型,如all-MiniLM-L6-v2量化INT8)
time: 时间戳
img_feat: 图片视觉向量(Qwen3-VL vision encoder输出,可选)
}
每次对话结束,把本轮完整对话文本向量化存入向量库。
步骤 2:当前用户提问做向量检索(核心筛选逻辑)
用户输入新问题 query_text;
用同一个 embedding 模型生成 query_vec;
在本地 FAISS 索引里做相似度 TopK 检索,取出匹配度最高 N 条历史记忆;
端侧推荐 TopK=3~8,RK3588 控制在 5 条以内最稳;
相似度阈值过滤:只保留余弦相似度 > 0.4/0.5 的记录,低于阈值直接丢弃(完全不相关)。
步骤 3:二次过滤:时间 + 长度控制
拿到 TopK 相似记录后还要两道裁剪:
- 时间衰减过滤(可选)
久远记忆降低权重:比如超过 7 天的相似记录,即便匹配高也减少拼接数量;
- 总长度预计算
把检索到的历史文本提前统计 token 数量,加上当前提问、系统 Prompt、图片 Token,总和不能超过模型 max_context_len。
从相似度最低的记录开始依次剔除,直到总长度安全。
步骤 4:把筛选后的记忆按顺序拼入上下文
拼接格式示例(塞进 Qwen3-VL 输入前置):
历史相关对话参考:
[历史1]
用户:xxx 图里是什么
模型:xxx
[历史2]
用户:刚才那个物体尺寸多少
模型:xxx
当前对话
用户:(新提问+图片)
3.两种拼接策略,适配 RK3588 嵌入式
方案 A:检索相关历史 + 保留本轮短时 KV 记忆(推荐,性能最好)
KV Cache:负责本次会话内短时记忆(刚聊完的几轮,不用检索);
向量库检索:负责跨会话长期记忆(上次开机、昨天的对话);
拼接逻辑:
先跑向量检索拿到历史相关记忆,拼在 Prompt 最前面;
再拼接当前会话未清空的多轮上下文(KV Cache 管理的近期对话);
优势:近期对话不走向量检索,省算力;久远跨会话靠向量召回。
方案 B:清空 KV,完全靠向量记忆(极简,但速度差)
每次提问都清空 KV Cache,全靠向量检索拼接所有上下文,适合单轮问答场景,不适合连续聊天。
4.直观例子
历史 1:昨天拍了汽车,问车型历史 2:前天拍了小狗,问品种历史 3:上周问家电使用方法
当前提问:“这辆车油耗多少?”
- 向量化检索,匹配度:历史 1 > 历史 3 > 历史 2;
- 阈值过滤,历史 2 相似度过低丢弃;
- 长度计算,只保留历史 1;
- 拼接历史 1 对话到 Prompt,再传入当前图片 + 问题给 Qwen3-VL。模型只会参考汽车那条记忆,不会带上小狗、家电无关内容。
5.生成的向量是什么
向量就是一组数字数组,比如 [0.12, -0.35, 0.78, ...],这里用的轻量模型输出固定 384 个浮点数,也叫文本嵌入(Embedding)。
一段对话 / 一句话会被 AI 模型压缩成这一串数字,数字组合唯一代表这句话的语义含义,不是字面文字。
举例:
文字:电动车快充多久
向量:[0.05, 0.21, -0.43 ... 共384个数]
文字:柯基小狗品种
向量:[-0.72, 0.11, 0.35 ... 共384个数]
语义相近的文本,对应的向量数字分布会高度接近;语义完全无关,数字差距很大。
二、为什么必须把文字转成向量
1.计算机不能直接理解文字,只能计算数字
文字是符号(汉字、字母),机器无法直接对比两段话 “像不像”。
只有全部转为数字数组,才能用数学公式计算两者的相似程度。
2.通过向量距离判断语义关联(核心作用)
用余弦相似度 / L2 距离计算两个向量差值:
向量越接近 → 数值差距越小 → 语义高度相关
向量差异巨大 → 数值差距大 → 内容无关
比如用户新问题 “电动车充满要多久”,它的向量和历史 “电动车续航、快充” 向量距离很近,和 “柯基小狗” 向量距离很远,程序就能自动筛选出相关历史对话,实现记忆匹配。
三.关键代码
1.store_memery.cpp
#include "memory_store.h"
#include <iostream>
#include <algorithm>
#include <set>
#include <cstring>
#include <unordered_map>
#include <cmath>
// 调试打印开关(默认关闭,与 main.cpp 同步)
#define ENABLE_DEBUG_PRINT 0
#if ENABLE_DEBUG_PRINT
#define DEBUG_PRINT(...) fprintf(stderr, __VA_ARGS__)
#else
#define DEBUG_PRINT(...) ((void)0)
#endif
MemoryStore::MemoryStore() : db_(nullptr) {}
MemoryStore::~MemoryStore() {
release();
}
int MemoryStore::init(const char* db_path) {
std::lock_guard<std::mutex> lock(mutex_);
int ret = sqlite3_open(db_path, &db_);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_open failed: " << sqlite3_errmsg(db_) << std::endl;
return -1;
}
ret = create_table();
if (ret != 0) {
return ret;
}
return 0;
}
int MemoryStore::release() {
std::lock_guard<std::mutex> lock(mutex_);
if (db_ != nullptr) {
sqlite3_close(db_);
db_ = nullptr;
}
return 0;
}
int MemoryStore::create_table() {
const char* create_sql = R"(
CREATE TABLE IF NOT EXISTS memory (
id INTEGER PRIMARY KEY AUTOINCREMENT,
text TEXT NOT NULL,
token_count INTEGER NOT NULL,
time INTEGER NOT NULL,
is_deprecated INTEGER DEFAULT 0,
confidence REAL DEFAULT 0.5,
version INTEGER DEFAULT 1
);
)";
char* err_msg;
int ret = sqlite3_exec(db_, create_sql, nullptr, nullptr, &err_msg);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_exec failed: " << err_msg << std::endl;
sqlite3_free(err_msg);
return -1;
}
const char* alter_deprecated = "ALTER TABLE memory ADD COLUMN IF NOT EXISTS is_deprecated INTEGER DEFAULT 0;";
ret = sqlite3_exec(db_, alter_deprecated, nullptr, nullptr, &err_msg);
if (ret != SQLITE_OK) {
sqlite3_free(err_msg);
}
const char* alter_confidence = "ALTER TABLE memory ADD COLUMN IF NOT EXISTS confidence REAL DEFAULT 0.5;";
ret = sqlite3_exec(db_, alter_confidence, nullptr, nullptr, &err_msg);
if (ret != SQLITE_OK) {
sqlite3_free(err_msg);
}
const char* alter_version = "ALTER TABLE memory ADD COLUMN IF NOT EXISTS version INTEGER DEFAULT 1;";
ret = sqlite3_exec(db_, alter_version, nullptr, nullptr, &err_msg);
if (ret != SQLITE_OK) {
sqlite3_free(err_msg);
}
return 0;
}
void MemoryStore::get_ngrams(const std::string& text, std::vector<std::string>& ngrams) {
ngrams.clear();
int n = 2;
for (size_t i = 0; i <= text.size() - n; i++) {
ngrams.push_back(text.substr(i, n));
}
}
float MemoryStore::idf_weighted_containment(const std::string& query, const std::string& doc,
const std::unordered_map<std::string, int>& doc_freq, int total_docs) {
std::vector<std::string> query_ngrams, doc_ngrams;
get_ngrams(query, query_ngrams);
get_ngrams(doc, doc_ngrams);
if (query_ngrams.empty()) return 0.0f;
std::set<std::string> doc_set(doc_ngrams.begin(), doc_ngrams.end());
float total_weight = 0.0f;
float matched_weight = 0.0f;
for (const auto& gram : query_ngrams) {
auto it = doc_freq.find(gram);
int df = it != doc_freq.end() ? it->second : 1;
float idf = log((float)(total_docs + 1) / df);
total_weight += idf;
if (doc_set.count(gram)) {
matched_weight += idf;
}
}
float score = total_weight > 0 ? matched_weight / total_weight : 0.0f;
DEBUG_PRINT("[Memory-DBG] IDF-containment: query=%zu ngrams, doc=%zu ngrams, score=%.4f\n",
query_ngrams.size(), doc_ngrams.size(), score);
return score;
}
int MemoryStore::add_memory(const std::string& text, int token_count) {
std::lock_guard<std::mutex> lock(mutex_);
if (db_ == nullptr) {
std::cerr << "Database not initialized" << std::endl;
return -1;
}
const char* insert_sql = "INSERT INTO memory (text, token_count, time) VALUES (?, ?, ?);";
sqlite3_stmt* stmt;
int ret = sqlite3_prepare_v2(db_, insert_sql, -1, &stmt, nullptr);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_prepare_v2 failed: " << sqlite3_errmsg(db_) << std::endl;
return -1;
}
sqlite3_bind_text(stmt, 1, text.c_str(), -1, SQLITE_TRANSIENT);
sqlite3_bind_int(stmt, 2, token_count);
sqlite3_bind_int64(stmt, 3, std::chrono::system_clock::now().time_since_epoch().count() / 1000000);
ret = sqlite3_step(stmt);
if (ret != SQLITE_DONE) {
std::cerr << "sqlite3_step failed: " << sqlite3_errmsg(db_) << std::endl;
sqlite3_finalize(stmt);
return -1;
}
sqlite3_finalize(stmt);
return 0;
}
int MemoryStore::add_memory_with_update(const std::string& user_query, const std::string& full_text, int token_count) {
std::lock_guard<std::mutex> lock(mutex_);
if (db_ == nullptr) {
std::cerr << "Database not initialized" << std::endl;
return -1;
}
const char* query_sql = "SELECT id, text, version FROM memory WHERE is_deprecated = 0;";
sqlite3_stmt* stmt;
int ret = sqlite3_prepare_v2(db_, query_sql, -1, &stmt, nullptr);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_prepare_v2 failed: " << sqlite3_errmsg(db_) << std::endl;
return -1;
}
std::vector<std::tuple<int, std::string, int>> all_memories;
while ((ret = sqlite3_step(stmt)) == SQLITE_ROW) {
int id = sqlite3_column_int(stmt, 0);
const char* text = (const char*)sqlite3_column_text(stmt, 1);
int version = sqlite3_column_int(stmt, 2);
all_memories.emplace_back(id, text, version);
}
sqlite3_finalize(stmt);
int total_docs = all_memories.size();
std::unordered_map<std::string, int> doc_freq;
for (const auto& mem : all_memories) {
std::vector<std::string> ngrams;
get_ngrams(std::get<1>(mem), ngrams);
std::set<std::string> unique_grams(ngrams.begin(), ngrams.end());
for (const auto& gram : unique_grams) {
doc_freq[gram]++;
}
}
const float UPDATE_THRESHOLD = 0.6f;
const float MIN_LENGTH_RATIO = 0.8f;
int replaced_id = -1;
int replaced_version = 1;
float max_sim = 0.0f;
for (const auto& mem : all_memories) {
float sim = idf_weighted_containment(user_query, std::get<1>(mem), doc_freq, total_docs);
if (sim > max_sim) {
max_sim = sim;
if (sim > UPDATE_THRESHOLD) {
float old_len = std::get<1>(mem).size();
float new_len = full_text.size();
if (new_len >= old_len * MIN_LENGTH_RATIO) {
replaced_id = std::get<0>(mem);
replaced_version = std::get<2>(mem);
DEBUG_PRINT("[Memory-DBG] candidate replace: id=%d, old_len=%zu, new_len=%zu, ratio=%.2f\n",
replaced_id, (size_t)old_len, (size_t)new_len, new_len/old_len);
} else {
DEBUG_PRINT("[Memory-DBG] skip replace: new text too short (old=%zu, new=%zu, ratio=%.2f < %.2f)\n",
(size_t)old_len, (size_t)new_len, new_len/old_len, MIN_LENGTH_RATIO);
}
}
}
}
float confidence = 0.5f;
if (replaced_id != -1) {
const char* update_sql = "UPDATE memory SET is_deprecated = 1 WHERE id = ?;";
ret = sqlite3_prepare_v2(db_, update_sql, -1, &stmt, nullptr);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_prepare_v2 failed: " << sqlite3_errmsg(db_) << std::endl;
return -1;
}
sqlite3_bind_int(stmt, 1, replaced_id);
ret = sqlite3_step(stmt);
sqlite3_finalize(stmt);
DEBUG_PRINT("[Memory-DBG] deprecated memory id=%d (similarity=%.4f), version=%d\n", replaced_id, max_sim, replaced_version);
confidence = 1.0f;
}
const char* insert_sql = "INSERT INTO memory (text, token_count, time, confidence, version) VALUES (?, ?, ?, ?, ?);";
if (all_memories.size() >= MAX_MEMORY_COUNT) {
DEBUG_PRINT("[Memory] 记忆数量已达上限(%d),删除最旧记忆\n", MAX_MEMORY_COUNT);
const char* delete_sql = "DELETE FROM memory WHERE id = (SELECT MIN(id) FROM memory WHERE is_deprecated = 0);";
char* err_msg;
ret = sqlite3_exec(db_, delete_sql, nullptr, nullptr, &err_msg);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_exec failed: " << err_msg << std::endl;
sqlite3_free(err_msg);
}
}
ret = sqlite3_prepare_v2(db_, insert_sql, -1, &stmt, nullptr);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_prepare_v2 failed: " << sqlite3_errmsg(db_) << std::endl;
return -1;
}
sqlite3_bind_text(stmt, 1, full_text.c_str(), -1, SQLITE_TRANSIENT);
sqlite3_bind_int(stmt, 2, token_count);
sqlite3_bind_int64(stmt, 3, std::chrono::system_clock::now().time_since_epoch().count() / 1000000);
sqlite3_bind_double(stmt, 4, confidence);
sqlite3_bind_int(stmt, 5, replaced_version + 1);
ret = sqlite3_step(stmt);
if (ret != SQLITE_DONE) {
std::cerr << "sqlite3_step failed: " << sqlite3_errmsg(db_) << std::endl;
sqlite3_finalize(stmt);
return -1;
}
sqlite3_finalize(stmt);
return 0;
}
int MemoryStore::retrieve(const std::string& query, int max_tokens,
std::vector<RetrievedMemory>& results) {
std::lock_guard<std::mutex> lock(mutex_);
results.clear();
std::vector<std::tuple<float, float, int, int>> scored_memories;
const char* query_sql = "SELECT id, text, token_count, time, confidence, version FROM memory WHERE is_deprecated = 0;";
sqlite3_stmt* stmt;
int ret = sqlite3_prepare_v2(db_, query_sql, -1, &stmt, nullptr);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_prepare_v2 failed: " << sqlite3_errmsg(db_) << std::endl;
return -1;
}
std::vector<std::tuple<int, std::string, int, long long, float>> all_memories;
while ((ret = sqlite3_step(stmt)) == SQLITE_ROW) {
int id = sqlite3_column_int(stmt, 0);
const char* text = (const char*)sqlite3_column_text(stmt, 1);
int version = sqlite3_column_int(stmt, 5);
long long time = sqlite3_column_int64(stmt, 3);
float confidence = sqlite3_column_double(stmt, 4);
all_memories.emplace_back(id, text, version, time, confidence);
}
sqlite3_finalize(stmt);
int total_docs = all_memories.size();
std::unordered_map<std::string, int> doc_freq;
for (const auto& mem : all_memories) {
std::vector<std::string> ngrams;
get_ngrams(std::get<1>(mem), ngrams);
std::set<std::string> unique_grams(ngrams.begin(), ngrams.end());
for (const auto& gram : unique_grams) {
doc_freq[gram]++;
}
}
DEBUG_PRINT("[Memory-DBG] total docs=%d, unique ngrams=%zu\n", total_docs, doc_freq.size());
long long now_ms = std::chrono::system_clock::now().time_since_epoch().count() / 1000000;
for (const auto& mem : all_memories) {
float sim = idf_weighted_containment(query, std::get<1>(mem), doc_freq, total_docs);
long long diff_ms = now_ms - std::get<3>(mem);
float days = diff_ms / (1000.0f * 60 * 60 * 24);
float recency_weight = exp(-days / 30.0f);
float confidence = std::get<4>(mem);
float final_score = sim * 0.5 + recency_weight * 0.3 + confidence * 0.2;
DEBUG_PRINT("[Memory-DBG] checking memory id=%d, sim=%.4f, recency=%.4f, conf=%.2f, final=%.4f, text=(%zu chars)\n",
std::get<0>(mem), sim, recency_weight, confidence, final_score, std::get<1>(mem).size());
if (final_score > SIMILARITY_THRESHOLD) {
scored_memories.emplace_back(final_score, sim, std::get<0>(mem), std::get<2>(mem));
DEBUG_PRINT("[Memory-DBG] => passed threshold, added to candidates\n");
}
}
DEBUG_PRINT("[Memory-DBG] total candidates after filtering: %zu\n", scored_memories.size());
std::sort(scored_memories.begin(), scored_memories.end(),
[](const std::tuple<float, float, int, int>& a, const std::tuple<float, float, int, int>& b) {
return std::get<0>(a) > std::get<0>(b);
});
const char* get_sql = "SELECT text, token_count, time, confidence, version FROM memory WHERE id = ?;";
ret = sqlite3_prepare_v2(db_, get_sql, -1, &stmt, nullptr);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_prepare_v2 failed: " << sqlite3_errmsg(db_) << std::endl;
return -1;
}
int total_tokens = 0;
int count = 0;
for (const auto& tuple : scored_memories) {
if (count >= TOP_K) break;
int id = std::get<2>(tuple);
float score = std::get<0>(tuple);
float sim = std::get<1>(tuple);
sqlite3_reset(stmt);
sqlite3_bind_int(stmt, 1, id);
ret = sqlite3_step(stmt);
if (ret != SQLITE_ROW) continue;
const char* text = (const char*)sqlite3_column_text(stmt, 0);
int token_count = sqlite3_column_int(stmt, 1);
long long time = sqlite3_column_int64(stmt, 2);
float confidence = sqlite3_column_double(stmt, 3);
int version = sqlite3_column_int(stmt, 4);
if (total_tokens + token_count <= max_tokens) {
RetrievedMemory rm;
rm.score = score;
rm.sim = sim;
rm.item.id = id;
rm.item.text = text;
rm.item.token_count = token_count;
rm.item.time = time;
rm.item.confidence = confidence;
rm.item.version = version;
rm.item.is_deprecated = 0;
results.push_back(rm);
total_tokens += token_count;
count++;
} else {
break;
}
}
sqlite3_finalize(stmt);
DEBUG_PRINT("[Memory-DBG] final retrieved: %zu memories, total_tokens=%d\n", results.size(), total_tokens);
for (size_t i = 0; i < results.size(); i++) {
DEBUG_PRINT("[Memory-DBG] [%zu] score=%.4f, conf=%.2f, version=%d, text=(%zu chars)\n",
i, results[i].score, results[i].item.confidence, results[i].item.version, results[i].item.text.size());
}
return 0;
}
int MemoryStore::get_memory_count(int* count) {
std::lock_guard<std::mutex> lock(mutex_);
const char* query_sql = "SELECT COUNT(*) FROM memory;";
sqlite3_stmt* stmt;
int ret = sqlite3_prepare_v2(db_, query_sql, -1, &stmt, nullptr);
if (ret != SQLITE_OK) {
std::cerr << "sqlite3_prepare_v2 failed: " << sqlite3_errmsg(db_) << std::endl;
return -1;
}
ret = sqlite3_step(stmt);
if (ret == SQLITE_ROW) {
*count = sqlite3_column_int(stmt, 0);
}
sqlite3_finalize(stmt);
return 0;
}
2. main.cpp
#include <stdint.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <iostream>
#include <fstream>
#include <chrono>
#include <thread>
#include <condition_variable>
#include <opencv2/opencv.hpp>
#include "image_enc.h"
#include "rkllm.h"
#include "memory_store.h"
// 调试打印开关(默认关闭)
#define ENABLE_DEBUG_PRINT 0
#if ENABLE_DEBUG_PRINT
#define DEBUG_PRINT(...) fprintf(stderr, __VA_ARGS__)
#else
#define DEBUG_PRINT(...) ((void)0)
#endif
using namespace std;
LLMHandle llmHandle = nullptr;
RKLLMParam g_llm_param;
MemoryStore memory_store;
std::string g_answer_text;
std::string g_current_prompt;
std::mutex g_answer_mutex;
std::condition_variable g_answer_cv;
bool g_answer_ready = false;
int g_last_token_count = 0;
int callback(RKLLMResult *result, void *userdata, LLMCallState state);
void exit_handler(int signal)
{
if (llmHandle != nullptr)
{
cout << "程序即将退出" << endl;
LLMHandle _tmp = llmHandle;
llmHandle = nullptr;
rkllm_destroy(_tmp);
}
memory_store.release();
exit(signal);
}
int callback(RKLLMResult *result, void *userdata, LLMCallState state)
{
if (state == RKLLM_RUN_FINISH)
{
printf("\n");
g_last_token_count = result->perf.prefill_tokens + result->perf.generate_tokens;
{
std::lock_guard<std::mutex> lock(g_answer_mutex);
g_answer_ready = true;
}
g_answer_cv.notify_one();
}
else if (state == RKLLM_RUN_ERROR)
{
printf("run error\n");
{
std::lock_guard<std::mutex> lock(g_answer_mutex);
g_answer_ready = true;
}
g_answer_cv.notify_one();
}
else if (state == RKLLM_RUN_NORMAL)
{
if (result->text != nullptr) {
printf("%s", result->text);
std::lock_guard<std::mutex> lock(g_answer_mutex);
g_answer_text += result->text;
}
}
return 0;
}
cv::Mat expand2square(const cv::Mat& img, const cv::Scalar& background_color) {
int width = img.cols;
int height = img.rows;
if (width == height) {
return img.clone();
}
int size = std::max(width, height);
cv::Mat result(size, size, img.type(), background_color);
int x_offset = (size - width) / 2;
int y_offset = (size - height) / 2;
cv::Rect roi(x_offset, y_offset, width, height);
img.copyTo(result(roi));
return result;
}
bool is_question(const std::string& str) {
std::string q_markers[] = {"吗", "什么", "几", "哪", "?", "?", "谁", "怎么", "为什么", "有什么",
"多少", "多大", "是不是", "能不能", "有没有"};
for (auto& m : q_markers) {
if (str.find(m) != std::string::npos) {
return true;
}
}
return false;
}
bool is_greeting(const std::string& str) {
std::string greetings[] = {"你好", "您好", "嗨", "哈喽", "早上好",
"早呀", "下午好",
"晚上好", "晚安", "hi", "hello", "hey"};
for (auto& g : greetings) {
if (str.find(g) != std::string::npos) {
return true;
}
}
return false;
}
bool is_fact(const std::string& str) {
std::string fact_verbs[] = {"是", "喜欢", "精通", "就读", "外号", "同学", "叫做", "来自", "毕业", "工作",
"做", "当", "住在", "担任", "成为", "获得", "觉得", "认为", "感觉"};
std::string units[] = {"岁", "年", "月", "日", "个", "人", "种", "门", "项", "本", "名", "次", "分", "公斤", "米"};
std::string weather_words[] = {"天气", "凉爽", "热", "冷", "下雨", "晴天", "阴天", "刮风", "温度", "摄氏度"};
std::string adjectives[] = {"厉害", "优秀", "专业", "酷", "牛", "棒", "好", "坏", "漂亮", "帅",
"可爱", "开心", "高兴", "难过", "伤心", "无聊", "有趣"};
bool has_number = false;
bool has_unit = false;
bool has_fact_verb = false;
bool has_weather_word = false;
bool has_adjective = false;
bool ends_with_period = false;
for (char c : str) {
if (c >= '0' && c <= '9') {
has_number = true;
}
}
for (auto& u : units) {
if (str.find(u) != std::string::npos) {
has_unit = true;
break;
}
}
for (auto& v : fact_verbs) {
if (str.find(v) != std::string::npos) {
has_fact_verb = true;
break;
}
}
for (auto& w : weather_words) {
if (str.find(w) != std::string::npos) {
has_weather_word = true;
break;
}
}
for (auto& a : adjectives) {
if (str.find(a) != std::string::npos) {
has_adjective = true;
break;
}
}
if (!str.empty()) {
char last = str.back();
if (last == '.' || last == '!') {
ends_with_period = true;
} else if (str.size() >= 3) {
std::string last_chars = str.substr(str.size() - 3);
if (last_chars == "。" || last_chars == "!") {
ends_with_period = true;
}
}
}
int score = 0;
if (has_number && has_unit) score += 3;
if (has_fact_verb) score += 2;
if (has_weather_word) score += 2;
if (has_adjective) score += 1;
if (ends_with_period) score += 1;
DEBUG_PRINT("[DEBUG] fact score: %d\n", score);
return score >= 2;
}
bool is_explicit_memory_request(const std::string& str) {
std::string markers[] = {"请记住", "记住", "要记住", "帮我记住", "别忘了"};
for (auto& m : markers) {
if (str.find(m) != std::string::npos) {
return true;
}
}
return false;
}
bool is_share(const std::string& str) {
std::string share_markers[] = {"今天", "昨天", "今天很", "今天真", "今天有点",
"天气", "心情", "感觉", "觉得", "很开心", "很高兴"};
for (auto& m : share_markers) {
if (str.find(m) != std::string::npos) {
return true;
}
}
return false;
}
std::string call_llm_sync(const std::string& prompt, int max_wait_ms = 5000) {
RKLLMInput rkllm_input;
RKLLMInferParam rkllm_infer_params;
memset(&rkllm_input, 0, sizeof(RKLLMInput));
memset(&rkllm_infer_params, 0, sizeof(RKLLMInferParam));
rkllm_input.input_type = RKLLM_INPUT_PROMPT;
rkllm_input.prompt_input = (char*)prompt.c_str();
{
std::lock_guard<std::mutex> lock(g_answer_mutex);
g_answer_text.clear();
g_answer_ready = false;
}
printf("[Extract] 正在提炼结构化摘要...\n");
try {
rkllm_run(llmHandle, &rkllm_input, &rkllm_infer_params, NULL);
} catch (const std::exception& e) {
printf("[Extract] LLM调用失败: %s\n", e.what());
return "";
}
std::unique_lock<std::mutex> lock(g_answer_mutex);
bool success = g_answer_cv.wait_for(lock, std::chrono::milliseconds(max_wait_ms),
[]{ return g_answer_ready; });
if (!success) {
printf("[Extract] LLM调用超时\n");
return "";
}
std::string result = g_answer_text;
printf("[Extract] 提炼完成: %s\n", result.c_str());
rkllm_clear_kv_cache(llmHandle, 1, nullptr, nullptr);
return result;
}
std::string extract_structured_memory(const std::string& user_input) {
std::string extract_prompt =
"请将以下信息提炼成简洁的结构化摘要,格式为\"姓名:属性1=值1,属性2=值2\",不要多余解释:\n"
+ user_input + "\n摘要:";
std::string extracted = call_llm_sync(extract_prompt, 5000);
if (extracted.empty()) {
printf("[Extract] 提炼失败,使用原始输入\n");
return user_input;
}
return extracted;
}
void store_memory(const std::string& user_query,
const std::string& answer, int token_count) {
if (is_question(user_query)) {
DEBUG_PRINT("[Memory] 跳过纯问题: %s\n", user_query.c_str());
return;
}
if (is_greeting(user_query)) {
DEBUG_PRINT("[Memory] 跳过问候语: %s\n", user_query.c_str());
return;
}
bool should_store = is_explicit_memory_request(user_query) ||
(is_fact(user_query) && !is_share(user_query));
if (!should_store) {
DEBUG_PRINT("[Memory] 跳过非重要事实: %s\n", user_query.c_str());
return;
}
std::string memory_text = extract_structured_memory(user_query);
memory_store.add_memory_with_update(user_query, memory_text, token_count);
DEBUG_PRINT("[Memory] 已保存/更新记忆: %s (长度=%zu)\n", memory_text.c_str(), memory_text.size());
}
bool is_valid_utf8(const std::string& str) {
size_t i = 0;
while (i < str.size()) {
unsigned char c = (unsigned char)str[i];
if (c < 0x80) {
i++;
} else if (c < 0xE0) {
if (i + 1 >= str.size()) return false;
if ((str[i+1] & 0xC0) != 0x80) return false;
i += 2;
} else if (c < 0xF0) {
if (i + 2 >= str.size()) return false;
if ((str[i+1] & 0xC0) != 0x80 || (str[i+2] & 0xC0) != 0x80) return false;
i += 3;
} else if (c < 0xF8) {
if (i + 3 >= str.size()) return false;
if ((str[i+1] & 0xC0) != 0x80 || (str[i+2] & 0xC0) != 0x80 || (str[i+3] & 0xC0) != 0x80) return false;
i += 4;
} else {
return false;
}
}
return true;
}
std::string sanitize_utf8(const std::string& str) {
std::string result;
result.reserve(str.size());
size_t i = 0;
while (i < str.size()) {
unsigned char c = (unsigned char)str[i];
if (c < 0x80) {
result += c;
i++;
} else if (c < 0xE0) {
if (i + 1 < str.size() && (str[i+1] & 0xC0) == 0x80) {
result += str[i];
result += str[i+1];
i += 2;
} else {
result += ' ';
i++;
}
} else if (c < 0xF0) {
if (i + 2 < str.size() && (str[i+1] & 0xC0) == 0x80 && (str[i+2] & 0xC0) == 0x80) {
result += str[i];
result += str[i+1];
result += str[i+2];
i += 3;
} else {
result += ' ';
i++;
}
} else if (c < 0xF8) {
if (i + 3 < str.size() && (str[i+1] & 0xC0) == 0x80 && (str[i+2] & 0xC0) == 0x80 && (str[i+3] & 0xC0) == 0x80) {
result += str[i];
result += str[i+1];
result += str[i+2];
result += str[i+3];
i += 4;
} else {
result += ' ';
i++;
}
} else {
result += ' ';
i++;
}
}
return result;
}
std::string build_memory_prompt(const std::vector<RetrievedMemory>& memories) {
if (memories.empty()) return "";
std::string prompt = "以下是用户告诉你的事实,请基于这些信息简短回答:\n";
int idx = 1;
for (const auto& rm : memories) {
std::string sanitized = sanitize_utf8(rm.item.text);
prompt += std::to_string(idx++) + ". " + sanitized + "\n";
}
prompt += "请根据以上信息回答:\n";
DEBUG_PRINT("[Memory-DBG] built prompt:\n%s\n", prompt.c_str());
return prompt;
}
int main(int argc, char** argv)
{
if (argc < 7) {
std::cerr << "Usage: " << argv[0]
<< " image_path encoder_model_path llm_model_path max_new_tokens max_context_len rknn_core_num "
<< "[img_start] [img_end] [img_content]\n";
return -1;
}
const char * image_path = argv[1];
const char * encoder_model_path = argv[2];
g_llm_param = rkllm_createDefaultParam();
g_llm_param.model_path = argv[3];
g_llm_param.top_k = 1;
g_llm_param.max_new_tokens = std::atoi(argv[4]);
g_llm_param.max_context_len = std::atoi(argv[5]);
g_llm_param.skip_special_token = true;
g_llm_param.extend_param.base_domain_id = 1;
g_llm_param.img_start = "<|vision_start|>";
g_llm_param.img_end = "<|vision_end|>";
g_llm_param.img_content = "<|image_pad|>";
if (argc == 7) {
std::cerr << "[Warning] Using default img_start/img_end/img_content: "
<< g_llm_param.img_start << " , "
<< g_llm_param.img_end << " , "
<< g_llm_param.img_content
<< ". Please customize these values according to your model, "
<< "otherwise the output may be incorrect.\n";
}
if (argc > 7) g_llm_param.img_start = argv[7];
if (argc > 8) g_llm_param.img_end = argv[8];
if (argc > 9) g_llm_param.img_content = argv[9];
int ret;
std::chrono::high_resolution_clock::time_point t_start_us = std::chrono::high_resolution_clock::now();
ret = rkllm_init(&llmHandle, &g_llm_param, callback);
if (ret == 0){
printf("rkllm init success\n");
} else {
printf("rkllm init failed\n");
exit_handler(-1);
}
std::chrono::high_resolution_clock::time_point t_load_end_us = std::chrono::high_resolution_clock::now();
auto load_time = std::chrono::duration_cast<std::chrono::microseconds>(t_load_end_us - t_start_us);
printf("%s: LLM Model loaded in %8.2f ms\n", __func__, load_time.count() / 1000.0);
rknn_app_context_t rknn_app_ctx;
memset(&rknn_app_ctx, 0, sizeof(rknn_app_context_t));
t_start_us = std::chrono::high_resolution_clock::now();
const int core_num = atoi(argv[6]);
ret = init_imgenc(encoder_model_path, &rknn_app_ctx, core_num);
if (ret != 0) {
printf("init_imgenc fail! ret=%d model_path=%s\n", ret, encoder_model_path);
return -1;
}
t_load_end_us = std::chrono::high_resolution_clock::now();
load_time = std::chrono::duration_cast<std::chrono::microseconds>(t_load_end_us - t_start_us);
printf("%s: ImgEnc Model loaded in %8.2f ms\n", __func__, load_time.count() / 1000.0);
ret = memory_store.init();
if (ret == 0) {
int count = 0;
memory_store.get_memory_count(&count);
printf("%s: Memory store initialized, loaded %d memories\n", __func__, count);
} else {
printf("%s: Memory store init failed\n", __func__);
}
cv::Mat img = cv::imread(image_path);
cv::cvtColor(img, img, cv::COLOR_BGR2RGB);
cv::Scalar background_color(127.5, 127.5, 127.5);
cv::Mat square_img = expand2square(img, background_color);
size_t image_width = rknn_app_ctx.model_width;
size_t image_height = rknn_app_ctx.model_height;
cv::Mat resized_img;
cv::Size new_size(image_width, image_height);
cv::resize(square_img, resized_img, new_size, 0, 0, cv::INTER_LINEAR);
size_t n_image_tokens = rknn_app_ctx.model_image_token;
size_t image_embed_len = rknn_app_ctx.model_embed_size;
size_t n_embed_output = rknn_app_ctx.io_num.n_output;
int rkllm_image_embed_len = n_image_tokens * image_embed_len * n_embed_output;
float img_vec[rkllm_image_embed_len];
memset(img_vec, 0, rkllm_image_embed_len * sizeof(float));
t_start_us = std::chrono::high_resolution_clock::now();
ret = run_imgenc(&rknn_app_ctx, resized_img.data, img_vec);
if (ret != 0) {
printf("run_imgenc fail! ret=%d\n", ret);
}
t_load_end_us = std::chrono::high_resolution_clock::now();
load_time = std::chrono::duration_cast<std::chrono::microseconds>(t_load_end_us - t_start_us);
printf("%s: ImgEnc Model inference took %8.2f ms\n", __func__, load_time.count() / 1000.0);
RKLLMInput rkllm_input;
memset(&rkllm_input, 0, sizeof(RKLLMInput));
RKLLMInferParam rkllm_infer_params;
memset(&rkllm_infer_params, 0, sizeof(RKLLMInferParam));
rkllm_infer_params.mode = RKLLM_INFER_GENERATE;
rkllm_infer_params.keep_history = 0;
vector<string> pre_input;
pre_input.push_back("<image>What is in the image?");
pre_input.push_back("<image>这张图片中有什么?");
cout << "\n**********************可输入以下问题对应序号获取回答/或自定义输入********************\n"
<< endl;
for (int i = 0; i < (int)pre_input.size(); i++)
{
cout << "[" << i << "] " << pre_input[i] << endl;
}
cout << "\n命令: exit(退出) | clear(清空记忆) | memory(查看记忆数)\n"
<< endl;
while(true) {
std::string input_str;
printf("\n");
printf("user: ");
std::getline(std::cin, input_str);
if (input_str.empty() || input_str.find_first_not_of(" \t\r\n") == std::string::npos) {
continue;
}
if (input_str == "exit") {
break;
}
if (input_str == "clear") {
ret = rkllm_clear_kv_cache(llmHandle, 1, nullptr, nullptr);
if (ret != 0) {
printf("clear kv cache failed!\n");
}
continue;
}
if (input_str == "memory") {
int count = 0;
memory_store.get_memory_count(&count);
printf("当前记忆条数: %d\n", count);
continue;
}
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;
}
}
std::string user_query = input_str;
std::string safe_input = sanitize_utf8(input_str);
bool is_question_flag = is_question(input_str);
bool is_greeting_flag = is_greeting(input_str);
bool is_fact_flag = is_fact(input_str);
bool is_share_flag = is_share(input_str);
bool should_query_memory = is_question_flag;
std::vector<RetrievedMemory> retrieved_memories;
if (should_query_memory && input_str.find("<image>") == std::string::npos) {
memory_store.retrieve(input_str, MAX_HISTORY_TOKENS, retrieved_memories);
}
bool use_memory = false;
if (!retrieved_memories.empty() && retrieved_memories[0].sim >= 0.3f) {
use_memory = true;
}
std::string memory_prompt = use_memory ? build_memory_prompt(retrieved_memories) : "";
if (!memory_prompt.empty()) {
DEBUG_PRINT("[Memory] 检索到 %zu 条相关记忆\n", retrieved_memories.size());
}
{
std::lock_guard<std::mutex> lock(g_answer_mutex);
g_answer_text.clear();
}
g_last_token_count = 0;
memset(&rkllm_input, 0, sizeof(RKLLMInput));
std::string behavior_prefix;
if (is_greeting_flag) {
behavior_prefix = "请简短友好地回应:\n";
} else if (is_question_flag) {
if (!memory_prompt.empty()) {
behavior_prefix = "";
} else {
behavior_prefix = "请简短准确地回答:\n";
}
} else if (is_fact_flag) {
behavior_prefix = "用户分享了新信息,请给出简短自然的回复,不要复述用户的话:\n";
} else {
behavior_prefix = "请像朋友一样自然地对话,简短回应:\n";
}
g_current_prompt = behavior_prefix + memory_prompt + safe_input;
if (input_str.find("<image>") == std::string::npos) {
rkllm_input.input_type = RKLLM_INPUT_PROMPT;
rkllm_input.prompt_input = (char*)g_current_prompt.c_str();
DEBUG_PRINT("[DEBUG] input_type=PROMPT, prompt_len=%zu\n", g_current_prompt.size());
} else {
rkllm_input.input_type = RKLLM_INPUT_MULTIMODAL;
rkllm_input.multimodal_input.prompt = (char*)g_current_prompt.c_str();
rkllm_input.multimodal_input.image_embed = img_vec;
rkllm_input.multimodal_input.n_image_tokens = n_image_tokens;
rkllm_input.multimodal_input.n_image = 1;
rkllm_input.multimodal_input.image_height = image_height;
rkllm_input.multimodal_input.image_width = image_width;
DEBUG_PRINT("[DEBUG] input_type=MULTIMODAL, n_image=1, prompt_len=%zu\n", g_current_prompt.size());
}
DEBUG_PRINT("[DEBUG] full_prompt=%.*s\n", std::min((int)g_current_prompt.size(), 100), g_current_prompt.c_str());
int max_tokens = g_llm_param.max_context_len;
while (g_current_prompt.size() > max_tokens * 2) {
if (retrieved_memories.empty()) {
printf("[WARN] Prompt too long (%zu chars), even without memory\n", g_current_prompt.size());
break;
}
printf("[WARN] Prompt too long (%zu chars), removing least relevant memory...\n", g_current_prompt.size());
size_t min_idx = 0;
float min_score = retrieved_memories[0].score;
for (size_t i = 1; i < retrieved_memories.size(); ++i) {
if (retrieved_memories[i].score < min_score) {
min_score = retrieved_memories[i].score;
min_idx = i;
}
}
retrieved_memories.erase(retrieved_memories.begin() + min_idx);
memory_prompt = build_memory_prompt(retrieved_memories);
if (is_greeting_flag) {
behavior_prefix = "请简短友好地回应:\n";
} else if (is_question_flag) {
if (!memory_prompt.empty()) {
behavior_prefix = "";
} else {
behavior_prefix = "请简短准确地回答:\n";
}
} else if (is_fact_flag) {
behavior_prefix = "用户分享了新信息,请给出简短自然的回复,不要复述用户的话:\n";
} else {
behavior_prefix = "请像朋友一样自然地对话,简短回应:\n";
}
g_current_prompt = behavior_prefix + memory_prompt + safe_input;
}
rkllm_abort(llmHandle);
printf("robot: ");
try {
rkllm_run(llmHandle, &rkllm_input, &rkllm_infer_params, NULL);
} catch (const std::exception& e) {
printf("\n[ERROR] rkllm_run exception: %s\n", e.what());
rkllm_clear_kv_cache(llmHandle, 1, nullptr, nullptr);
printf("[WARN] 输入可能包含非法字符,请重新输入\n");
continue;
}
std::string answer;
{
std::lock_guard<std::mutex> lock(g_answer_mutex);
answer = g_answer_text;
}
if (!answer.empty() && g_last_token_count > 0) {
store_memory(user_query, answer, g_last_token_count);
}
rkllm_clear_kv_cache(llmHandle, 1, nullptr, nullptr);
}
ret = release_imgenc(&rknn_app_ctx);
if (ret != 0) {
printf("release_imgenc fail! ret=%d\n", ret);
}
memory_store.release();
rkllm_destroy(llmHandle);
return 0;
}
浙公网安备 33010602011771号