mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-22 18:08:04 +00:00
Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| fddefc78a1 | |||
| 3f431063e2 | |||
| 18d7ad3a36 |
@ -34,15 +34,11 @@ type StaticEmbedder struct {
|
|||||||
jieba *gojieba.Jieba
|
jieba *gojieba.Jieba
|
||||||
stopWords map[string]bool
|
stopWords map[string]bool
|
||||||
|
|
||||||
// words 是词向量表。**用 float32 存**:源文件(fastText 文本格式)本身就是 float32,
|
words map[string][]float64
|
||||||
// 用 float64 存等于把 578 万……不,是 57.8 万词 × 300 维的常驻内存凭空翻倍
|
|
||||||
// (实测生产:float64 → 1.29GB,float32 → 0.65GB)。相似度计算仍在 float64 里累加,
|
|
||||||
// 精度不受影响。改回 float64 会被 TestStaticEmbedder_VectorMemIsFloat32 拦住。
|
|
||||||
words map[string][]float32
|
|
||||||
dim int
|
dim int
|
||||||
loaded bool
|
loaded bool
|
||||||
|
|
||||||
unkVec []float32
|
unkVec []float64
|
||||||
unkNorm float64
|
unkNorm float64
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -154,7 +150,7 @@ func NewStaticEmbedder(modelPaths ...string) *StaticEmbedder {
|
|||||||
e := &StaticEmbedder{
|
e := &StaticEmbedder{
|
||||||
jieba: GetJieba(),
|
jieba: GetJieba(),
|
||||||
stopWords: sw,
|
stopWords: sw,
|
||||||
words: make(map[string][]float32),
|
words: make(map[string][]float64),
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(modelPaths) == 0 {
|
if len(modelPaths) == 0 {
|
||||||
@ -258,16 +254,15 @@ func (e *StaticEmbedder) load(spec string, primary bool) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
vec := make([]float32, dim)
|
vec := make([]float64, dim)
|
||||||
for i := 0; i < dim; i++ {
|
for i := 0; i < dim; i++ {
|
||||||
// 源文件是 float32 精度的文本向量:用 32 位解析,与源数据一致。
|
v, _ := strconv.ParseFloat(fields[i+1], 64)
|
||||||
v, _ := strconv.ParseFloat(fields[i+1], 32)
|
vec[i] = v
|
||||||
vec[i] = float32(v)
|
|
||||||
}
|
}
|
||||||
e.words[word] = vec
|
e.words[word] = vec
|
||||||
if primary {
|
if primary {
|
||||||
for i := range vecSum {
|
for i := range vecSum {
|
||||||
vecSum[i] += float64(vec[i])
|
vecSum[i] += vec[i]
|
||||||
}
|
}
|
||||||
count++
|
count++
|
||||||
}
|
}
|
||||||
@ -281,13 +276,11 @@ func (e *StaticEmbedder) load(spec string, primary bool) error {
|
|||||||
for i := range vecSum {
|
for i := range vecSum {
|
||||||
vecSum[i] /= float64(count)
|
vecSum[i] /= float64(count)
|
||||||
}
|
}
|
||||||
e.unkVec = make([]float32, dim)
|
e.unkVec = make([]float64, dim)
|
||||||
for i, v := range vecSum {
|
copy(e.unkVec, vecSum)
|
||||||
e.unkVec[i] = float32(v)
|
|
||||||
}
|
|
||||||
var normSq float64
|
var normSq float64
|
||||||
for _, v := range e.unkVec {
|
for _, v := range e.unkVec {
|
||||||
normSq += float64(v) * float64(v)
|
normSq += v * v
|
||||||
}
|
}
|
||||||
e.unkNorm = float64(math.Sqrt(normSq))
|
e.unkNorm = float64(math.Sqrt(normSq))
|
||||||
e.loaded = true
|
e.loaded = true
|
||||||
@ -373,11 +366,11 @@ func (e *StaticEmbedder) Vectorize(text string) vector.Vector {
|
|||||||
|
|
||||||
if !ok {
|
if !ok {
|
||||||
for i, v := range unkVec {
|
for i, v := range unkVec {
|
||||||
sum[i] += w * float64(v)
|
sum[i] += w * v
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
for i, v := range vec {
|
for i, v := range vec {
|
||||||
sum[i] += w * float64(v)
|
sum[i] += w * v
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
weightSum += w
|
weightSum += w
|
||||||
|
|||||||
@ -1,37 +0,0 @@
|
|||||||
package memory
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
"unsafe"
|
|
||||||
)
|
|
||||||
|
|
||||||
// 词向量必须用 float32 存。
|
|
||||||
//
|
|
||||||
// 这条判据是拿生产内存换来的:向量本体 = 词数 × 维数 × 每元素字节数。
|
|
||||||
// 生产配置加载了 200000(zh) + 378151(en) = 57.8 万词 × 300 维 ⇒
|
|
||||||
// float64 = 1.29GB、float32 = 0.65GB(差 0.65GB 常驻)。
|
|
||||||
// 源数据(fastText 文本格式)本身就是 float32 精度,用 float64 存没有任何收益。
|
|
||||||
//
|
|
||||||
// 若有人把类型改回 float64,本测试**编译失败**(`var vec []float32` 的类型断言),
|
|
||||||
// 这正是想要的效果。
|
|
||||||
func TestStaticEmbedder_VectorMemIsFloat32(t *testing.T) {
|
|
||||||
e := newSynthEmbedder(t, 300)
|
|
||||||
|
|
||||||
words := 0
|
|
||||||
bytes := 0
|
|
||||||
for _, vec := range e.words {
|
|
||||||
var typed []float32 = vec // 编译期断言:存储必须是 []float32
|
|
||||||
if len(typed) != e.dim {
|
|
||||||
t.Fatalf("维度不符: %d != %d", len(typed), e.dim)
|
|
||||||
}
|
|
||||||
words++
|
|
||||||
bytes += len(typed) * int(unsafe.Sizeof(typed[0]))
|
|
||||||
}
|
|
||||||
if words == 0 {
|
|
||||||
t.Fatal("合成模型应至少加载一个词")
|
|
||||||
}
|
|
||||||
// float32:每词 300×4 = 1200 字节;float64 会是 2400
|
|
||||||
if want := words * e.dim * 4; bytes != want {
|
|
||||||
t.Fatalf("向量本体字节数应 %d(float32),实际 %d", want, bytes)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Reference in New Issue
Block a user