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;
}

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;
}
posted @ 2026-08-11 11:29  wssheng  阅读(0)  评论(0)    收藏  举报