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;
}
-
main.cpp
#include <stdint.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include
#include
#include
#include
#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)
#endifusing 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_guardstd::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_guardstd::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_guardstd::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
& 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;}