From c6ebbfca349b5ad7058126c68e8be038989a0781 Mon Sep 17 00:00:00 2001 From: root Date: Wed, 5 Aug 2026 09:59:30 +0800 Subject: [PATCH] =?UTF-8?q?memory=20=E6=B5=8B=E8=AF=95=E8=87=AA=E5=8C=85?= =?UTF-8?q?=E5=90=AB=EF=BC=9A=E5=90=88=E6=88=90=20word2vec=20=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E6=9B=BF=E4=BB=A3=20/tmp=20=E5=A4=A7=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E4=BE=9D=E8=B5=96=EF=BC=8C=E6=B5=8B=E8=AF=95=E4=B8=8D?= =?UTF-8?q?=E5=86=8D=E8=81=94=E7=BD=91=E4=B8=8B=E8=BD=BD=20fastText?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/memory/bilingual_test.go | 12 ++-- internal/memory/clean_stress_test.go | 10 +-- internal/memory/real_context_test.go | 95 +++++++++++++++++-------- internal/memory/static_embedder_test.go | 10 +-- internal/memory/synth_model_test.go | 88 +++++++++++++++++++++++ 5 files changed, 171 insertions(+), 44 deletions(-) create mode 100644 internal/memory/synth_model_test.go diff --git a/internal/memory/bilingual_test.go b/internal/memory/bilingual_test.go index 5ed048e..fe6884a 100644 --- a/internal/memory/bilingual_test.go +++ b/internal/memory/bilingual_test.go @@ -15,8 +15,8 @@ type bilingualEvent struct { } func TestBilingualPruningAccuracy(t *testing.T) { - zhPath := "/tmp/cc.zh.top200k.vec" - enPath := "/tmp/cc.en.top200k.vec" + zhPath := writeSynthModel(t, 300) + enPath := writeSynthModel(t, 300) // Test with Chinese-only vs Chinese+English type modelConfig struct { @@ -103,8 +103,8 @@ func TestBilingualPruningAccuracy(t *testing.T) { } func TestBilingualCrossLingualSimilarity(t *testing.T) { - zhPath := "/tmp/cc.zh.top200k.vec" - enPath := "/tmp/cc.en.top200k.vec" + zhPath := writeSynthModel(t, 300) + enPath := writeSynthModel(t, 300) e := NewStaticEmbedder(zhPath, enPath) if !e.Loaded() { t.Skip("embedder not loaded") @@ -137,7 +137,7 @@ func TestBilingualCrossLingualSimilarity(t *testing.T) { } func TestBilingualEdgeCases(t *testing.T) { - zhPath := "/tmp/cc.zh.top200k.vec" + zhPath := writeSynthModel(t, 300) e := NewStaticEmbedder(zhPath) if !e.Loaded() { t.Skip("embedder not loaded") @@ -165,7 +165,7 @@ func TestBilingualEdgeCases(t *testing.T) { } func TestBilingualVectorizeClean(t *testing.T) { - zhPath := "/tmp/cc.zh.top200k.vec" + zhPath := writeSynthModel(t, 300) e := NewStaticEmbedder(zhPath) inputs := []string{ diff --git a/internal/memory/clean_stress_test.go b/internal/memory/clean_stress_test.go index 6fb3833..ee406b3 100644 --- a/internal/memory/clean_stress_test.go +++ b/internal/memory/clean_stress_test.go @@ -41,7 +41,7 @@ func TestCleanStressPrecision(t *testing.T) { events := genStressEvents(200) topics := []string{"大学招生", "老大私聊", "前端开发", "服务器运维", "股票基金"} - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) if !e.Loaded() { t.Skip("embedder not loaded") } @@ -113,7 +113,7 @@ func TestCleanStressPrecision(t *testing.T) { func TestCleanStressCrossTopic(t *testing.T) { events := genStressEvents(200) - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) if !e.Loaded() { t.Skip("embedder not loaded") } @@ -154,7 +154,7 @@ func TestCleanStressCrossTopic(t *testing.T) { } func TestCleanTemplateNoiseSuppression(t *testing.T) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) if !e.Loaded() { t.Skip("embedder not loaded") } @@ -181,7 +181,7 @@ func TestCleanTemplateNoiseSuppression(t *testing.T) { } func TestCleanVectorConsistency(t *testing.T) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) if !e.Loaded() { t.Skip("embedder not loaded") } @@ -217,7 +217,7 @@ func TestCleanVectorConsistency(t *testing.T) { } func BenchmarkCleanVectorize(b *testing.B) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(b, 300) if !e.Loaded() { b.Skip("embedder not loaded") } diff --git a/internal/memory/real_context_test.go b/internal/memory/real_context_test.go index 0880c1d..3f566af 100644 --- a/internal/memory/real_context_test.go +++ b/internal/memory/real_context_test.go @@ -1,8 +1,6 @@ package memory import ( - "encoding/json" - "os" "sort" "strings" "testing" @@ -65,21 +63,10 @@ func TestCleanText(t *testing.T) { } func TestRealContextPerSourceVector(t *testing.T) { - modelPath := "/tmp/cc.zh.sample.vec" - if _, err := os.Stat(modelPath); os.IsNotExist(err) { - t.Skip("real embedding file not found") - } - e := NewStaticEmbedder(modelPath) + e := newSynthEmbedder(t, 300) - data, err := os.ReadFile("/tmp/context.json") - if err != nil { - t.Fatal(err) - } - var raw []realEvent - if err := json.Unmarshal(data, &raw); err != nil { - t.Fatal(err) - } - t.Logf("loaded %d real events", len(raw)) + raw := genRealEvents() + t.Logf("loaded %d synthetic events", len(raw)) type scored struct { idx int @@ -257,19 +244,9 @@ func TestRealContextPerSourceVector(t *testing.T) { } func TestRealContextEmbedderStats(t *testing.T) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") - if !e.Loaded() { - t.Skip("embedder not loaded") - } + e := newSynthEmbedder(t, 300) - data, err := os.ReadFile("/tmp/context.json") - if err != nil { - t.Fatal(err) - } - var raw []realEvent - if err := json.Unmarshal(data, &raw); err != nil { - t.Fatal(err) - } + raw := genRealEvents() for _, ev := range raw[:5] { var text string @@ -292,6 +269,68 @@ func TestRealContextEmbedderStats(t *testing.T) { } } +// genRealEvents 生成跨领域的合成事件夹具(替代曾依赖的 /tmp/context.json)。 +func genRealEvents() []realEvent { + now := time.Now() + var out []realEvent + add := func(src, input, response string) { + out = append(out, realEvent{Timestamp: now, Source: src, Input: input, Response: response}) + } + + for _, s := range []string{ + "河南医药大学招生分数线出来了", + "扶高升学咨询群讨论专业排名", + "医药大学录取结果查询方法", + "河南高考志愿填报咨询", + "大学招生简章发布了", + } { + add("qq", s, "") + } + add("agent", "来自升学群的(医药大学咨询)消息", "已回复关于河南医药大学录取分数线的咨询") + + for _, s := range []string{ + "老大私聊消息安排了一个任务", + "回复老大关于服务器配置的问题", + "老大要求检查容器运行状态", + "老大说了关于组件封装的事情", + "给老大汇报工作进展", + } { + add("qq", s, "") + } + add("agent", "收到老大的(私聊)消息", "已回复老大关于任务安排的消息") + + for _, s := range []string{ + "用户要求图片转换SVG工具", + "图片转SVG后尺寸优化完成", + "转换图片格式为SVG", + "生成SVG工具使用说明", + } { + add("cli", "查询"+s, "查到了"+s+"的结果") + } + + for _, s := range []string{ + "南航航空航天专业介绍", + "电气专业转南航的录取咨询", + "南航院校分数线讨论", + "航空航天方向的就业前景", + } { + add("qq", s, "") + } + + for _, s := range []string{ + "今天天气怎么样", + "晚上吃什么", + "推荐一部电影", + "股票基金收益行情如何", + "查询快递送达状态", + "设置一个明早的闹钟", + } { + add("cold_storage", s, "已处理,结果记录完成") + } + + return out +} + func containsAny(s string, subs []string) bool { for _, sub := range subs { if sub != "" && strings.Contains(s, sub) { diff --git a/internal/memory/static_embedder_test.go b/internal/memory/static_embedder_test.go index 8e13b96..19bd61c 100644 --- a/internal/memory/static_embedder_test.go +++ b/internal/memory/static_embedder_test.go @@ -32,7 +32,7 @@ func cosineSim(a, b vector.Vector) float64 { } func TestStaticEmbedderLoad(t *testing.T) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) if !e.Loaded() { t.Fatal("embedder should be loaded") } @@ -42,7 +42,7 @@ func TestStaticEmbedderLoad(t *testing.T) { } func TestStaticEmbedderConsistency(t *testing.T) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) v1 := e.Vectorize("今天天气怎么样") v2 := e.Vectorize("今天天气怎么样") @@ -81,7 +81,7 @@ func TestStaticEmbedderFallback(t *testing.T) { } func TestStaticEmbedderAllInDenseSpace(t *testing.T) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) texts := []string{ "今天天气怎么样", @@ -105,7 +105,7 @@ func TestStaticEmbedderAllInDenseSpace(t *testing.T) { } func TestStaticEmbedderSemanticSimilarity(t *testing.T) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) pairs := []struct { a, b string @@ -134,7 +134,7 @@ func TestStaticEmbedderSemanticSimilarity(t *testing.T) { } func TestContextPruneWithRealEmbedding(t *testing.T) { - e := NewStaticEmbedder("/tmp/cc.zh.sample.vec") + e := newSynthEmbedder(t, 300) type event struct { input string diff --git a/internal/memory/synth_model_test.go b/internal/memory/synth_model_test.go new file mode 100644 index 0000000..aa4f543 --- /dev/null +++ b/internal/memory/synth_model_test.go @@ -0,0 +1,88 @@ +package memory + +// 共享测试工具:在 t.TempDir() 中生成小型合成 word2vec 文本模型, +// 替代曾硬编码在 /tmp 的真实 fastText 模型(依赖网络下载与全局文件)。 +// 各领域词簇落在正交维度上,语义测试断言即可稳定复现。 + +import ( + "fmt" + "os" + "path/filepath" + "sort" + "strconv" + "strings" + "testing" +) + +// synthClusters 领域词簇:簇索引即向量维度,同簇词共享同一维度。 +var synthClusters = map[int][]string{ + 0: {"天气", "下雨", "明天", "今天", "台风", "降温", "气象", "预报", "雨"}, + 1: {"股票", "基金", "投资", "定投", "收益", "行情", "理财", "风险", "策略", "市场", "stock", "涨"}, + 2: {"微积分", "导数", "数学", "作业", "公式", "求解", "计算", "题目"}, + 3: {"大学", "招生", "录取", "分数", "医学", "医药", "专业", "分数线", "志愿", "高考", "升学", "排名", "咨询", "河南", "university"}, + 4: {"老大", "私聊", "消息", "回复", "汇报", "任务", "安排", "收到", "boss"}, + 5: {"图片", "转换", "工具", "图标", "画布", "svg"}, + 6: {"南航", "航空", "航天", "电气", "院校", "民航"}, + 7: {"前端", "组件", "封装", "布局", "页面", "路由", "交互", "调试", "优化", "代码", "开发", "逻辑", "react", "javascript"}, + 8: {"服务器", "配置", "部署", "容器", "代理", "证书", "备份", "恢复", "监控", "告警", "数据库", "反向", "续期", "nginx", "docker", "server", "computer", "电脑"}, +} + +// synthNeutral 通用词:落在最后一个维度,不参与领域区分。 +var synthNeutral = []string{ + "会", "不会", "帮", "查", "看", "最近", "晚上", "随便", "推荐", "电影", + "注意", "安全", "可以", "说", "事情", "要求", "检查", "状态", "获取", + "实现", "测试", "结果", "问题", "处理", "已经", "相关", "需要", "使用", + "方法", "信息", "好的", "内容", "发送", "询问", "朋友", +} + +// writeSynthModel 生成合成 word2vec 文本模型文件并返回路径。 +func writeSynthModel(t testing.TB, dim int) string { + t.Helper() + + words := make(map[string][]float64) + var clusterDims []int + for c := range synthClusters { + clusterDims = append(clusterDims, c) + } + sort.Ints(clusterDims) + + for _, c := range clusterDims { + for _, w := range synthClusters[c] { + vec := make([]float64, dim) + vec[c] = 1.0 + words[w] = vec + } + } + neutralDim := len(synthClusters) + for _, w := range synthNeutral { + vec := make([]float64, dim) + vec[neutralDim] = 1.0 + words[w] = vec + } + + path := filepath.Join(t.TempDir(), "synth.vec") + var sb strings.Builder + sb.WriteString(fmt.Sprintf("%d %d\n", len(words), dim)) + for w, vec := range words { + sb.WriteString(w) + for _, v := range vec { + sb.WriteString(" ") + sb.WriteString(strconv.FormatFloat(v, 'f', 4, 64)) + } + sb.WriteString("\n") + } + if err := os.WriteFile(path, []byte(sb.String()), 0o644); err != nil { + t.Fatalf("writeSynthModel: %v", err) + } + return path +} + +// newSynthEmbedder 返回加载了合成模型的 StaticEmbedder。 +func newSynthEmbedder(t testing.TB, dim int) *StaticEmbedder { + t.Helper() + e := NewStaticEmbedder(writeSynthModel(t, dim)) + if !e.Loaded() { + t.Fatal("synthetic embedder should be loaded") + } + return e +}