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. 向量化检索,匹配度:历史 1 > 历史 3 > 历史 2;
  2. 阈值过滤,历史 2 相似度过低丢弃;
  3. 长度计算,只保留历史 1;
  4. 拼接历史 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;
}
  1. 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)
    #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_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;

    }

相关推荐
fthux9 小时前
装闭 RenoPit 源码解析(06):SSE如何实时推送AI装修分析进度
人工智能·ai·开源·github·open source·renopit
LyridRelan11 小时前
Skill Cli MCP 的三者关系以及部分Q&A
人工智能·ai·个人开发
兮动人11 小时前
AI 开始接管工作台:谁会成为下一代电脑入口?
人工智能·ai·chatgpt·codex·workbuddy
AI探索派15 小时前
AI视频笔记工具怎么选:Ai好记、通义听悟、Get笔记、听脑AI横向对比
ai·效率神器·视频总结·ai笔记·视频转笔记
天天有money16 小时前
API中转站与多账号内容运营:如何统一管理不同项目的调用
gpt·ai·chatgpt·产品运营·内容运营
yinghuoAI202616 小时前
电商卖货,本质上是“视觉的战争”
人工智能·ai·ai作画·ai作图·ai生视频
Are_you_kidding_16 小时前
codeX集成deepSeek的API key步骤教程
ai·oneapi
测试_AI_一辰17 小时前
AI Agent 评测最隐蔽的坑-记忆
人工智能·算法·ai·自动化·ai编程
csdn_aspnet17 小时前
Copilot能换成本地吗?VSCode接本地大模型本地化接入方案
ide·vscode·ai·ollama
努力搬砖的咸鱼17 小时前
AI Agent测试全景图:它到底改变了什么
人工智能·python·ai·集成测试·pytest·agent·ai编程