nandi/gleanpublic⑂ Fork 0
⑂ 588cf0a
Commits
⬇ Clone ▾
git clone https://git.rickub.com/nandi/glean.git
git clone ssh://git@rickub.com/nandi/glean.git

Host key fingerprint (ed25519): SHA256:iycHnxEyq0Q7uyVpB7JlznP0G7JrTPXLYRcAU5CSLhc — verify it before your first connect.

Refactor embedding logic to support instructions and sqlite-vec knnUnverified

Julien Robert committed 2026-04-26T02:00:15+02:00 Browse files
588cf0a parent: e0fa297
modified internal/cluster/article.go +7 -17
@@ -76,7 +76,7 @@ func (e *Engine) ComputeArticleEmbeddings(ctx context.Context) error {
7676 texts[j] = a.text
7777 }
7878
79- embeddings, err := e.embedder.Embed(ctx, texts)
79+ embeddings, err := e.embedder.Embed(ctx, texts, "Represent this news article for retrieving topically similar articles. Focus on the subjects, themes, and key entities discussed.")
8080 if err != nil {
8181 return fmt.Errorf("embed batch %d: %w", i/embedBatchSize, err)
8282 }
@@ -153,8 +153,7 @@ func (e *Engine) populateContentBoost(ctx context.Context, conn *sql.Conn, userD
153153 }
154154
155155 dim := e.embedder.Dimension()
156- sumVec := make([]float32, dim)
157- count := 0
156+ var blobs [][]byte
158157 likedSet := make(map[int64]bool)
159158 for embRows.Next() {
160159 var id int64
@@ -163,28 +162,19 @@ func (e *Engine) populateContentBoost(ctx context.Context, conn *sql.Conn, userD
163162 embRows.Close()
164163 return err
165164 }
166- v := deserializeFloat32(blob)
167- if len(v) != dim {
165+ if len(blob) != dim*4 {
168166 continue
169167 }
170- for j := range sumVec {
171- sumVec[j] += v[j]
172- }
173- count++
168+ blobs = append(blobs, blob)
174169 likedSet[id] = true
175170 }
176171 embRows.Close()
177172
178- if count == 0 {
173+ if len(blobs) == 0 {
179174 return nil
180175 }
181176
182- avgVec := make([]float32, dim)
183- for j := range avgVec {
184- avgVec[j] = sumVec[j] / float32(count)
185- }
186-
187- queryBlob, err := vec.SerializeFloat32(avgVec)
177+ queryBlob, err := avgEmbeddings(blobs, dim)
188178 if err != nil {
189179 return fmt.Errorf("serialize query vector: %w", err)
190180 }
@@ -308,7 +298,7 @@ func (e *Engine) ComputeFeedEmbeddings(ctx context.Context) error {
308298 texts[j] = f.text
309299 }
310300
311- embeddings, err := e.embedder.Embed(ctx, texts)
301+ embeddings, err := e.embedder.Embed(ctx, texts, "Represent this RSS feed description for discovering feeds with similar editorial focus and topic coverage.")
312302 if err != nil {
313303 return fmt.Errorf("embed feed batch %d: %w", i/embedBatchSize, err)
314304 }
@@ -76,7 +76,7 @@ func (e *Engine) ComputeArticleEmbeddings(ctx context.Context) error {
76 texts[j] = a.text76 texts[j] = a.text
77 }77 }
78 78
79- embeddings, err := e.embedder.Embed(ctx, texts)79+ embeddings, err := e.embedder.Embed(ctx, texts, "Represent this news article for retrieving topically similar articles. Focus on the subjects, themes, and key entities discussed.")
80 if err != nil {80 if err != nil {
81 return fmt.Errorf("embed batch %d: %w", i/embedBatchSize, err)81 return fmt.Errorf("embed batch %d: %w", i/embedBatchSize, err)
82 }82 }
@@ -153,8 +153,7 @@ func (e *Engine) populateContentBoost(ctx context.Context, conn *sql.Conn, userD
153 }153 }
154 154
155 dim := e.embedder.Dimension()155 dim := e.embedder.Dimension()
156- sumVec := make([]float32, dim)156+ var blobs [][]byte
157- count := 0
158 likedSet := make(map[int64]bool)157 likedSet := make(map[int64]bool)
159 for embRows.Next() {158 for embRows.Next() {
160 var id int64159 var id int64
@@ -163,28 +162,19 @@ func (e *Engine) populateContentBoost(ctx context.Context, conn *sql.Conn, userD
163 embRows.Close()162 embRows.Close()
164 return err163 return err
165 }164 }
166- v := deserializeFloat32(blob)165+ if len(blob) != dim*4 {
167- if len(v) != dim {
168 continue166 continue
169 }167 }
170- for j := range sumVec {168+ blobs = append(blobs, blob)
171- sumVec[j] += v[j]
172- }
173- count++
174 likedSet[id] = true169 likedSet[id] = true
175 }170 }
176 embRows.Close()171 embRows.Close()
177 172
178- if count == 0 {173+ if len(blobs) == 0 {
179 return nil174 return nil
180 }175 }
181 176
182- avgVec := make([]float32, dim)177+ queryBlob, err := avgEmbeddings(blobs, dim)
183- for j := range avgVec {
184- avgVec[j] = sumVec[j] / float32(count)
185- }
186-
187- queryBlob, err := vec.SerializeFloat32(avgVec)
188 if err != nil {178 if err != nil {
189 return fmt.Errorf("serialize query vector: %w", err)179 return fmt.Errorf("serialize query vector: %w", err)
190 }180 }
@@ -308,7 +298,7 @@ func (e *Engine) ComputeFeedEmbeddings(ctx context.Context) error {
308 texts[j] = f.text298 texts[j] = f.text
309 }299 }
310 300
311- embeddings, err := e.embedder.Embed(ctx, texts)301+ embeddings, err := e.embedder.Embed(ctx, texts, "Represent this RSS feed description for discovering feeds with similar editorial focus and topic coverage.")
312 if err != nil {302 if err != nil {
313 return fmt.Errorf("embed feed batch %d: %w", i/embedBatchSize, err)303 return fmt.Errorf("embed feed batch %d: %w", i/embedBatchSize, err)
314 }304 }
modified internal/cluster/embed.go +43 -18
@@ -1,35 +1,34 @@
11 package cluster
22
33 import (
4- "bytes"
54 "context"
6- "encoding/binary"
5+ "unsafe"
76
7+ vec "github.com/asg017/sqlite-vec-go-bindings/cgo"
88 "github.com/openai/openai-go"
99 "github.com/openai/openai-go/option"
1010 )
1111
12-// Embedder generates vector embeddings for text inputs. Implementations must be
13-// safe for concurrent use.
12+// Embedder generates vector embeddings for text inputs.
1413 type Embedder interface {
15- Embed(ctx context.Context, texts []string) ([][]float32, error)
14+ Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error)
1615 Dimension() int
1716 }
1817
19-type OpenAIEmbedder struct {
18+type EmbedderClient struct {
2019 client openai.Client
2120 model string
2221 dimension int
2322 }
2423
25-type OpenAIEmbedderConfig struct {
24+type EmbedderClientConfig struct {
2625 BaseURL string
2726 APIKey string
2827 Model string
2928 Dimension int
3029 }
3130
32-func NewOpenAIEmbedder(cfg OpenAIEmbedderConfig) *OpenAIEmbedder {
31+func NewEmbedderClient(cfg EmbedderClientConfig) *EmbedderClient {
3332 opts := []option.RequestOption{}
3433 if cfg.BaseURL != "" {
3534 opts = append(opts, option.WithBaseURL(cfg.BaseURL))
@@ -37,22 +36,29 @@ func NewOpenAIEmbedder(cfg OpenAIEmbedderConfig) *OpenAIEmbedder {
3736 if cfg.APIKey != "" {
3837 opts = append(opts, option.WithAPIKey(cfg.APIKey))
3938 }
40- return &OpenAIEmbedder{
39+ return &EmbedderClient{
4140 client: openai.NewClient(opts...),
4241 model: cfg.Model,
4342 dimension: cfg.Dimension,
4443 }
4544 }
4645
47-func (e *OpenAIEmbedder) Dimension() int {
46+func (e *EmbedderClient) Dimension() int {
4847 return e.dimension
4948 }
5049
51-func (e *OpenAIEmbedder) Embed(ctx context.Context, texts []string) ([][]float32, error) {
50+func (e *EmbedderClient) Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) {
51+ inputs := texts
52+ if instruction != "" {
53+ inputs = make([]string, len(texts))
54+ for i, t := range texts {
55+ inputs[i] = instruction + "\n" + t
56+ }
57+ }
5258 resp, err := e.client.Embeddings.New(ctx, openai.EmbeddingNewParams{
5359 Model: e.model,
5460 Input: openai.EmbeddingNewParamsInputUnion{
55- OfArrayOfStrings: texts,
61+ OfArrayOfStrings: inputs,
5662 },
5763 })
5864 if err != nil {
@@ -69,12 +75,31 @@ func (e *OpenAIEmbedder) Embed(ctx context.Context, texts []string) ([][]float32
6975 return embeddings, nil
7076 }
7177
72-func deserializeFloat32(data []byte) []float32 {
73- if len(data)%4 != 0 {
78+func avgEmbeddings(blobs [][]byte, dim int) ([]byte, error) {
79+ sum := make([]float32, dim)
80+ count := 0
81+ for _, blob := range blobs {
82+ v := bytesToFloat32s(blob, dim)
83+ if v == nil {
84+ continue
85+ }
86+ for j := range sum {
87+ sum[j] += v[j]
88+ }
89+ count++
90+ }
91+ if count == 0 {
92+ return nil, nil
93+ }
94+ for j := range sum {
95+ sum[j] /= float32(count)
96+ }
97+ return vec.SerializeFloat32(sum)
98+}
99+
100+func bytesToFloat32s(data []byte, expectedDim int) []float32 {
101+ if len(data) != expectedDim*4 {
74102 return nil
75103 }
76- result := make([]float32, len(data)/4)
77- r := bytes.NewReader(data)
78- _ = binary.Read(r, binary.LittleEndian, &result)
79- return result
104+ return unsafe.Slice((*float32)(unsafe.Pointer(&data[0])), expectedDim)
80105 }
@@ -1,35 +1,34 @@
1 package cluster1 package cluster
2 2
3 import (3 import (
4- "bytes"
5 "context"4 "context"
6- "encoding/binary"5+ "unsafe"
7 6
7+ vec "github.com/asg017/sqlite-vec-go-bindings/cgo"
8 "github.com/openai/openai-go"8 "github.com/openai/openai-go"
9 "github.com/openai/openai-go/option"9 "github.com/openai/openai-go/option"
10 )10 )
11 11
12-// Embedder generates vector embeddings for text inputs. Implementations must be12+// Embedder generates vector embeddings for text inputs.
13-// safe for concurrent use.
14 type Embedder interface {13 type Embedder interface {
15- Embed(ctx context.Context, texts []string) ([][]float32, error)14+ Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error)
16 Dimension() int15 Dimension() int
17 }16 }
18 17
19-type OpenAIEmbedder struct {18+type EmbedderClient struct {
20 client openai.Client19 client openai.Client
21 model string20 model string
22 dimension int21 dimension int
23 }22 }
24 23
25-type OpenAIEmbedderConfig struct {24+type EmbedderClientConfig struct {
26 BaseURL string25 BaseURL string
27 APIKey string26 APIKey string
28 Model string27 Model string
29 Dimension int28 Dimension int
30 }29 }
31 30
32-func NewOpenAIEmbedder(cfg OpenAIEmbedderConfig) *OpenAIEmbedder {31+func NewEmbedderClient(cfg EmbedderClientConfig) *EmbedderClient {
33 opts := []option.RequestOption{}32 opts := []option.RequestOption{}
34 if cfg.BaseURL != "" {33 if cfg.BaseURL != "" {
35 opts = append(opts, option.WithBaseURL(cfg.BaseURL))34 opts = append(opts, option.WithBaseURL(cfg.BaseURL))
@@ -37,22 +36,29 @@ func NewOpenAIEmbedder(cfg OpenAIEmbedderConfig) *OpenAIEmbedder {
37 if cfg.APIKey != "" {36 if cfg.APIKey != "" {
38 opts = append(opts, option.WithAPIKey(cfg.APIKey))37 opts = append(opts, option.WithAPIKey(cfg.APIKey))
39 }38 }
40- return &OpenAIEmbedder{39+ return &EmbedderClient{
41 client: openai.NewClient(opts...),40 client: openai.NewClient(opts...),
42 model: cfg.Model,41 model: cfg.Model,
43 dimension: cfg.Dimension,42 dimension: cfg.Dimension,
44 }43 }
45 }44 }
46 45
47-func (e *OpenAIEmbedder) Dimension() int {46+func (e *EmbedderClient) Dimension() int {
48 return e.dimension47 return e.dimension
49 }48 }
50 49
51-func (e *OpenAIEmbedder) Embed(ctx context.Context, texts []string) ([][]float32, error) {50+func (e *EmbedderClient) Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) {
51+ inputs := texts
52+ if instruction != "" {
53+ inputs = make([]string, len(texts))
54+ for i, t := range texts {
55+ inputs[i] = instruction + "\n" + t
56+ }
57+ }
52 resp, err := e.client.Embeddings.New(ctx, openai.EmbeddingNewParams{58 resp, err := e.client.Embeddings.New(ctx, openai.EmbeddingNewParams{
53 Model: e.model,59 Model: e.model,
54 Input: openai.EmbeddingNewParamsInputUnion{60 Input: openai.EmbeddingNewParamsInputUnion{
55- OfArrayOfStrings: texts,61+ OfArrayOfStrings: inputs,
56 },62 },
57 })63 })
58 if err != nil {64 if err != nil {
@@ -69,12 +75,31 @@ func (e *OpenAIEmbedder) Embed(ctx context.Context, texts []string) ([][]float32
69 return embeddings, nil75 return embeddings, nil
70 }76 }
71 77
72-func deserializeFloat32(data []byte) []float32 {78+func avgEmbeddings(blobs [][]byte, dim int) ([]byte, error) {
73- if len(data)%4 != 0 {79+ sum := make([]float32, dim)
80+ count := 0
81+ for _, blob := range blobs {
82+ v := bytesToFloat32s(blob, dim)
83+ if v == nil {
84+ continue
85+ }
86+ for j := range sum {
87+ sum[j] += v[j]
88+ }
89+ count++
90+ }
91+ if count == 0 {
92+ return nil, nil
93+ }
94+ for j := range sum {
95+ sum[j] /= float32(count)
96+ }
97+ return vec.SerializeFloat32(sum)
98+}
99+
100+func bytesToFloat32s(data []byte, expectedDim int) []float32 {
101+ if len(data) != expectedDim*4 {
74 return nil102 return nil
75 }103 }
76- result := make([]float32, len(data)/4)104+ return unsafe.Slice((*float32)(unsafe.Pointer(&data[0])), expectedDim)
77- r := bytes.NewReader(data)
78- _ = binary.Read(r, binary.LittleEndian, &result)
79- return result
80 }105 }
modified internal/cluster/jaccard.go +52 -66
@@ -5,7 +5,6 @@ import (
55 "database/sql"
66 "fmt"
77 "log/slog"
8- "math"
98 "sync"
109 )
1110
@@ -122,96 +121,83 @@ func (e *Engine) computeEmbeddingSimilarity(ctx context.Context, tx *sql.Tx) err
122121 if e.embedder == nil {
123122 return nil
124123 }
125- e.logger.Debug("computing embedding similarity")
124+ e.logger.Debug("computing embedding similarity via vec0 KNN")
126125
127- stagingRows, err := tx.QueryContext(ctx, `SELECT feed_a, feed_b FROM _feed_sim_staging`)
126+ feedRows, err := tx.QueryContext(ctx, `SELECT feed_url, embedding FROM recs.feed_embeddings`)
128127 if err != nil {
129128 return err
130129 }
131130
132- type pair struct{ a, b string }
133- var pairs []pair
134- feedSet := make(map[string]bool)
135- for stagingRows.Next() {
136- var p pair
137- if err := stagingRows.Scan(&p.a, &p.b); err != nil {
138- stagingRows.Close()
131+ type feedEmb struct {
132+ url string
133+ vec []byte
134+ }
135+ var feeds []feedEmb
136+ for feedRows.Next() {
137+ var f feedEmb
138+ if err := feedRows.Scan(&f.url, &f.vec); err != nil {
139+ feedRows.Close()
139140 return err
140141 }
141- pairs = append(pairs, p)
142- feedSet[p.a] = true
143- feedSet[p.b] = true
142+ feeds = append(feeds, f)
144143 }
145- stagingRows.Close()
144+ feedRows.Close()
146145
147- if len(feedSet) == 0 {
146+ if len(feeds) == 0 {
148147 return nil
149148 }
150149
151- ph := make([]string, 0, len(feedSet))
152- args := make([]any, 0, len(feedSet))
153- for url := range feedSet {
154- ph = append(ph, "?")
155- args = append(args, url)
156- }
157- embRows, err := tx.QueryContext(ctx,
158- fmt.Sprintf("SELECT feed_url, embedding FROM recs.feed_embeddings WHERE feed_url IN (%s)", joinPh(ph)),
159- args...,
160- )
150+ const knnLimit = 50
151+ stmt, err := tx.PrepareContext(ctx, `
152+ SELECT feed_url, distance
153+ FROM recs.feed_embeddings
154+ WHERE embedding MATCH ? AND k = ?
155+ ORDER BY distance
156+ `)
161157 if err != nil {
162158 return err
163159 }
160+ defer stmt.Close()
164161
165- embeddings := make(map[string][]float32)
166- for embRows.Next() {
167- var url string
168- var blob []byte
169- if err := embRows.Scan(&url, &blob); err != nil {
170- embRows.Close()
171- return err
172- }
173- v := deserializeFloat32(blob)
174- if len(v) > 0 {
175- embeddings[url] = v
176- }
162+ updateStmt, err := tx.PrepareContext(ctx,
163+ `UPDATE _feed_sim_staging SET jaccard = jaccard + ? WHERE feed_a = ? AND feed_b = ?`,
164+ )
165+ if err != nil {
166+ return err
177167 }
178- embRows.Close()
168+ defer updateStmt.Close()
179169
180- for _, p := range pairs {
181- vecA, okA := embeddings[p.a]
182- vecB, okB := embeddings[p.b]
183- if !okA || !okB {
184- continue
185- }
186- sim := cosineSimilarity(vecA, vecB)
187- if sim <= 0 {
188- continue
189- }
190- boost := sim * e.config.DescriptionWeight
191- if _, err := tx.ExecContext(ctx,
192- `UPDATE _feed_sim_staging SET jaccard = jaccard + ? WHERE feed_a = ? AND feed_b = ?`,
193- boost, p.a, p.b,
194- ); err != nil {
170+ for _, f := range feeds {
171+ knnRows, err := stmt.QueryContext(ctx, f.vec, knnLimit)
172+ if err != nil {
195173 return err
196174 }
175+ for knnRows.Next() {
176+ var neighborURL string
177+ var dist float64
178+ if err := knnRows.Scan(&neighborURL, &dist); err != nil {
179+ knnRows.Close()
180+ return err
181+ }
182+ if neighborURL == f.url || dist <= 0 {
183+ continue
184+ }
185+ boost := 1.0 / (1.0 + dist) * e.config.DescriptionWeight
186+ a, b := f.url, neighborURL
187+ if a > b {
188+ a, b = b, a
189+ }
190+ if _, err := updateStmt.ExecContext(ctx, boost, a, b); err != nil {
191+ knnRows.Close()
192+ return err
193+ }
194+ }
195+ knnRows.Close()
197196 }
198197
199198 return nil
200199 }
201200
202-func cosineSimilarity(a, b []float32) float64 {
203- var dot, normA, normB float64
204- for i := range a {
205- dot += float64(a[i]) * float64(b[i])
206- normA += float64(a[i]) * float64(a[i])
207- normB += float64(b[i]) * float64(b[i])
208- }
209- if normA == 0 || normB == 0 {
210- return 0
211- }
212- return dot / (math.Sqrt(normA) * math.Sqrt(normB))
213-}
214-
215201 // ComputeUserSimilarity recomputes the user_similarity table: subscription
216202 // Jaccard + time-decayed like co-occurrence + tag overlap + follow boost.
217203 func (e *Engine) ComputeUserSimilarity(ctx context.Context) error {
@@ -5,7 +5,6 @@ import (
5 "database/sql"5 "database/sql"
6 "fmt"6 "fmt"
7 "log/slog"7 "log/slog"
8- "math"
9 "sync"8 "sync"
10 )9 )
11 10
@@ -122,96 +121,83 @@ func (e *Engine) computeEmbeddingSimilarity(ctx context.Context, tx *sql.Tx) err
122 if e.embedder == nil {121 if e.embedder == nil {
123 return nil122 return nil
124 }123 }
125- e.logger.Debug("computing embedding similarity")124+ e.logger.Debug("computing embedding similarity via vec0 KNN")
126 125
127- stagingRows, err := tx.QueryContext(ctx, `SELECT feed_a, feed_b FROM _feed_sim_staging`)126+ feedRows, err := tx.QueryContext(ctx, `SELECT feed_url, embedding FROM recs.feed_embeddings`)
128 if err != nil {127 if err != nil {
129 return err128 return err
130 }129 }
131 130
132- type pair struct{ a, b string }131+ type feedEmb struct {
133- var pairs []pair132+ url string
134- feedSet := make(map[string]bool)133+ vec []byte
135- for stagingRows.Next() {134+ }
136- var p pair135+ var feeds []feedEmb
137- if err := stagingRows.Scan(&p.a, &p.b); err != nil {136+ for feedRows.Next() {
138- stagingRows.Close()137+ var f feedEmb
138+ if err := feedRows.Scan(&f.url, &f.vec); err != nil {
139+ feedRows.Close()
139 return err140 return err
140 }141 }
141- pairs = append(pairs, p)142+ feeds = append(feeds, f)
142- feedSet[p.a] = true
143- feedSet[p.b] = true
144 }143 }
145- stagingRows.Close()144+ feedRows.Close()
146 145
147- if len(feedSet) == 0 {146+ if len(feeds) == 0 {
148 return nil147 return nil
149 }148 }
150 149
151- ph := make([]string, 0, len(feedSet))150+ const knnLimit = 50
152- args := make([]any, 0, len(feedSet))151+ stmt, err := tx.PrepareContext(ctx, `
153- for url := range feedSet {152+ SELECT feed_url, distance
154- ph = append(ph, "?")153+ FROM recs.feed_embeddings
155- args = append(args, url)154+ WHERE embedding MATCH ? AND k = ?
156- }155+ ORDER BY distance
157- embRows, err := tx.QueryContext(ctx,156+ `)
158- fmt.Sprintf("SELECT feed_url, embedding FROM recs.feed_embeddings WHERE feed_url IN (%s)", joinPh(ph)),
159- args...,
160- )
161 if err != nil {157 if err != nil {
162 return err158 return err
163 }159 }
160+ defer stmt.Close()
164 161
165- embeddings := make(map[string][]float32)162+ updateStmt, err := tx.PrepareContext(ctx,
166- for embRows.Next() {163+ `UPDATE _feed_sim_staging SET jaccard = jaccard + ? WHERE feed_a = ? AND feed_b = ?`,
167- var url string164+ )
168- var blob []byte165+ if err != nil {
169- if err := embRows.Scan(&url, &blob); err != nil {166+ return err
170- embRows.Close()
171- return err
172- }
173- v := deserializeFloat32(blob)
174- if len(v) > 0 {
175- embeddings[url] = v
176- }
177 }167 }
178- embRows.Close()168+ defer updateStmt.Close()
179 169
180- for _, p := range pairs {170+ for _, f := range feeds {
181- vecA, okA := embeddings[p.a]171+ knnRows, err := stmt.QueryContext(ctx, f.vec, knnLimit)
182- vecB, okB := embeddings[p.b]172+ if err != nil {
183- if !okA || !okB {
184- continue
185- }
186- sim := cosineSimilarity(vecA, vecB)
187- if sim <= 0 {
188- continue
189- }
190- boost := sim * e.config.DescriptionWeight
191- if _, err := tx.ExecContext(ctx,
192- `UPDATE _feed_sim_staging SET jaccard = jaccard + ? WHERE feed_a = ? AND feed_b = ?`,
193- boost, p.a, p.b,
194- ); err != nil {
195 return err173 return err
196 }174 }
175+ for knnRows.Next() {
176+ var neighborURL string
177+ var dist float64
178+ if err := knnRows.Scan(&neighborURL, &dist); err != nil {
179+ knnRows.Close()
180+ return err
181+ }
182+ if neighborURL == f.url || dist <= 0 {
183+ continue
184+ }
185+ boost := 1.0 / (1.0 + dist) * e.config.DescriptionWeight
186+ a, b := f.url, neighborURL
187+ if a > b {
188+ a, b = b, a
189+ }
190+ if _, err := updateStmt.ExecContext(ctx, boost, a, b); err != nil {
191+ knnRows.Close()
192+ return err
193+ }
194+ }
195+ knnRows.Close()
197 }196 }
198 197
199 return nil198 return nil
200 }199 }
201 200
202-func cosineSimilarity(a, b []float32) float64 {
203- var dot, normA, normB float64
204- for i := range a {
205- dot += float64(a[i]) * float64(b[i])
206- normA += float64(a[i]) * float64(a[i])
207- normB += float64(b[i]) * float64(b[i])
208- }
209- if normA == 0 || normB == 0 {
210- return 0
211- }
212- return dot / (math.Sqrt(normA) * math.Sqrt(normB))
213-}
214-
215 // ComputeUserSimilarity recomputes the user_similarity table: subscription201 // ComputeUserSimilarity recomputes the user_similarity table: subscription
216 // Jaccard + time-decayed like co-occurrence + tag overlap + follow boost.202 // Jaccard + time-decayed like co-occurrence + tag overlap + follow boost.
217 func (e *Engine) ComputeUserSimilarity(ctx context.Context) error {203 func (e *Engine) ComputeUserSimilarity(ctx context.Context) error {
modified internal/cluster/jaccard_test.go +1 -1
@@ -584,7 +584,7 @@ func NewMockEmbedder(dimension int) *MockEmbedder {
584584 return &MockEmbedder{dimension: dimension}
585585 }
586586
587-func (m *MockEmbedder) Embed(_ context.Context, texts []string) ([][]float32, error) {
587+func (m *MockEmbedder) Embed(_ context.Context, texts []string, _ string) ([][]float32, error) {
588588 result := make([][]float32, len(texts))
589589 for i, text := range texts {
590590 vec := make([]float32, m.dimension)
@@ -584,7 +584,7 @@ func NewMockEmbedder(dimension int) *MockEmbedder {
584 return &MockEmbedder{dimension: dimension}584 return &MockEmbedder{dimension: dimension}
585 }585 }
586 586
587-func (m *MockEmbedder) Embed(_ context.Context, texts []string) ([][]float32, error) {587+func (m *MockEmbedder) Embed(_ context.Context, texts []string, _ string) ([][]float32, error) {
588 result := make([][]float32, len(texts))588 result := make([][]float32, len(texts))
589 for i, text := range texts {589 for i, text := range texts {
590 vec := make([]float32, m.dimension)590 vec := make([]float32, m.dimension)
modified internal/cluster/scoring.go +6 -18
@@ -4,8 +4,6 @@ import (
44 "context"
55 "database/sql"
66 "fmt"
7-
8- vec "github.com/asg017/sqlite-vec-go-bindings/cgo"
97 )
108
119 type FeedRecommendation struct {
@@ -292,6 +290,7 @@ func (e *Engine) coldStartFromEmbeddings(ctx context.Context, userDID string, li
292290 }
293291 defer conn.Close()
294292
293+ dim := e.embedder.Dimension()
295294 subRows, err := conn.QueryContext(ctx, `
296295 SELECT fe.feed_url, fe.embedding FROM articles.subscriptions s
297296 JOIN recs.feed_embeddings fe ON fe.feed_url = s.feed_url
@@ -301,10 +300,8 @@ func (e *Engine) coldStartFromEmbeddings(ctx context.Context, userDID string, li
301300 return nil, err
302301 }
303302
304- dim := e.embedder.Dimension()
305- sumVec := make([]float32, dim)
306- subCount := 0
307303 var subFeedURLs []string
304+ var blobs [][]byte
308305 for subRows.Next() {
309306 var url string
310307 var blob []byte
@@ -312,33 +309,24 @@ func (e *Engine) coldStartFromEmbeddings(ctx context.Context, userDID string, li
312309 subRows.Close()
313310 return nil, err
314311 }
315- v := deserializeFloat32(blob)
316- if len(v) != dim {
312+ if len(blob) != dim*4 {
317313 continue
318314 }
319- for j := range sumVec {
320- sumVec[j] += v[j]
321- }
322- subCount++
323315 subFeedURLs = append(subFeedURLs, url)
316+ blobs = append(blobs, blob)
324317 }
325318 subRows.Close()
326319
327- if subCount == 0 {
320+ if len(blobs) == 0 {
328321 return nil, nil
329322 }
330323
331- avgVec := make([]float32, dim)
332- for j := range avgVec {
333- avgVec[j] = sumVec[j] / float32(subCount)
334- }
335-
336324 subSet := make(map[string]bool, len(subFeedURLs))
337325 for _, u := range subFeedURLs {
338326 subSet[u] = true
339327 }
340328
341- queryBlob, err := vec.SerializeFloat32(avgVec)
329+ queryBlob, err := avgEmbeddings(blobs, dim)
342330 if err != nil {
343331 return nil, fmt.Errorf("serialize query vector: %w", err)
344332 }
@@ -4,8 +4,6 @@ import (
4 "context"4 "context"
5 "database/sql"5 "database/sql"
6 "fmt"6 "fmt"
7-
8- vec "github.com/asg017/sqlite-vec-go-bindings/cgo"
9 )7 )
10 8
11 type FeedRecommendation struct {9 type FeedRecommendation struct {
@@ -292,6 +290,7 @@ func (e *Engine) coldStartFromEmbeddings(ctx context.Context, userDID string, li
292 }290 }
293 defer conn.Close()291 defer conn.Close()
294 292
293+ dim := e.embedder.Dimension()
295 subRows, err := conn.QueryContext(ctx, `294 subRows, err := conn.QueryContext(ctx, `
296 SELECT fe.feed_url, fe.embedding FROM articles.subscriptions s295 SELECT fe.feed_url, fe.embedding FROM articles.subscriptions s
297 JOIN recs.feed_embeddings fe ON fe.feed_url = s.feed_url296 JOIN recs.feed_embeddings fe ON fe.feed_url = s.feed_url
@@ -301,10 +300,8 @@ func (e *Engine) coldStartFromEmbeddings(ctx context.Context, userDID string, li
301 return nil, err300 return nil, err
302 }301 }
303 302
304- dim := e.embedder.Dimension()
305- sumVec := make([]float32, dim)
306- subCount := 0
307 var subFeedURLs []string303 var subFeedURLs []string
304+ var blobs [][]byte
308 for subRows.Next() {305 for subRows.Next() {
309 var url string306 var url string
310 var blob []byte307 var blob []byte
@@ -312,33 +309,24 @@ func (e *Engine) coldStartFromEmbeddings(ctx context.Context, userDID string, li
312 subRows.Close()309 subRows.Close()
313 return nil, err310 return nil, err
314 }311 }
315- v := deserializeFloat32(blob)312+ if len(blob) != dim*4 {
316- if len(v) != dim {
317 continue313 continue
318 }314 }
319- for j := range sumVec {
320- sumVec[j] += v[j]
321- }
322- subCount++
323 subFeedURLs = append(subFeedURLs, url)315 subFeedURLs = append(subFeedURLs, url)
316+ blobs = append(blobs, blob)
324 }317 }
325 subRows.Close()318 subRows.Close()
326 319
327- if subCount == 0 {320+ if len(blobs) == 0 {
328 return nil, nil321 return nil, nil
329 }322 }
330 323
331- avgVec := make([]float32, dim)
332- for j := range avgVec {
333- avgVec[j] = sumVec[j] / float32(subCount)
334- }
335-
336 subSet := make(map[string]bool, len(subFeedURLs))324 subSet := make(map[string]bool, len(subFeedURLs))
337 for _, u := range subFeedURLs {325 for _, u := range subFeedURLs {
338 subSet[u] = true326 subSet[u] = true
339 }327 }
340 328
341- queryBlob, err := vec.SerializeFloat32(avgVec)329+ queryBlob, err := avgEmbeddings(blobs, dim)
342 if err != nil {330 if err != nil {
343 return nil, fmt.Errorf("serialize query vector: %w", err)331 return nil, fmt.Errorf("serialize query vector: %w", err)
344 }332 }
modified main.go +1 -1
@@ -58,7 +58,7 @@ func main() {
5858
5959 var embedder cluster.Embedder
6060 if embedURL := envOr("GLEAN_EMBED_BASE_URL", ""); embedURL != "" {
61- embedder = cluster.NewOpenAIEmbedder(cluster.OpenAIEmbedderConfig{
61+ embedder = cluster.NewEmbedderClient(cluster.EmbedderClientConfig{
6262 BaseURL: embedURL,
6363 APIKey: envOr("GLEAN_EMBED_API_KEY", ""),
6464 Model: envOr("GLEAN_EMBED_MODEL", "text-embedding-3-small"),
@@ -58,7 +58,7 @@ func main() {
58 58
59 var embedder cluster.Embedder59 var embedder cluster.Embedder
60 if embedURL := envOr("GLEAN_EMBED_BASE_URL", ""); embedURL != "" {60 if embedURL := envOr("GLEAN_EMBED_BASE_URL", ""); embedURL != "" {
61- embedder = cluster.NewOpenAIEmbedder(cluster.OpenAIEmbedderConfig{61+ embedder = cluster.NewEmbedderClient(cluster.EmbedderClientConfig{
62 BaseURL: embedURL,62 BaseURL: embedURL,
63 APIKey: envOr("GLEAN_EMBED_API_KEY", ""),63 APIKey: envOr("GLEAN_EMBED_API_KEY", ""),
64 Model: envOr("GLEAN_EMBED_MODEL", "text-embedding-3-small"),64 Model: envOr("GLEAN_EMBED_MODEL", "text-embedding-3-small"),