87 lines
1.7 KiB
Go
87 lines
1.7 KiB
Go
package utils
|
|
|
|
import (
|
|
"fmt"
|
|
"sort"
|
|
"strconv"
|
|
"strings"
|
|
)
|
|
|
|
func ParseEmbedding(value interface{}) ([]float64, error) {
|
|
switch v := value.(type) {
|
|
case []float64:
|
|
return v, nil
|
|
|
|
case []interface{}:
|
|
vec := make([]float64, len(v))
|
|
for i, x := range v {
|
|
f, ok := x.(float64)
|
|
if !ok {
|
|
return nil, fmt.Errorf("invalid element type %T in embedding", x)
|
|
}
|
|
vec[i] = f
|
|
}
|
|
return vec, nil
|
|
|
|
case string:
|
|
// pgvector returns a string like "[0.12, 0.53, ...]" sometimes
|
|
s := strings.Trim(v, "[]")
|
|
parts := strings.Split(s, ",")
|
|
vec := make([]float64, len(parts))
|
|
for i, p := range parts {
|
|
f, err := strconv.ParseFloat(strings.TrimSpace(p), 64)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
vec[i] = f
|
|
}
|
|
return vec, nil
|
|
|
|
default:
|
|
return nil, fmt.Errorf("unexpected embedding type: %T", v)
|
|
}
|
|
}
|
|
|
|
func AverageVector(vs [][]float64) []float64 {
|
|
dim := len(vs[0])
|
|
avg := make([]float64, dim)
|
|
|
|
for _, v := range vs {
|
|
for i, val := range v {
|
|
avg[i] += val
|
|
}
|
|
}
|
|
for i := range avg {
|
|
avg[i] /= float64(len(vs))
|
|
}
|
|
return avg
|
|
}
|
|
|
|
func TopWords(posts []string, stopWords map[string]struct{}, limit int) []string {
|
|
counts := map[string]int{}
|
|
for _, p := range posts {
|
|
for _, w := range strings.Fields(strings.ToLower(p)) {
|
|
if len(w) < 4 {
|
|
continue
|
|
}
|
|
if _, blocked := stopWords[w]; blocked {
|
|
continue
|
|
}
|
|
counts[w]++
|
|
}
|
|
}
|
|
type kv struct {
|
|
Word string
|
|
Count int
|
|
}
|
|
arr := make([]kv, 0, len(counts))
|
|
for w, c := range counts {
|
|
arr = append(arr, kv{w, c})
|
|
}
|
|
sort.Slice(arr, func(i, j int) bool { return arr[i].Count > arr[j].Count })
|
|
out := []string{}
|
|
for i := 0; i < len(arr) && i < limit; i++ {
|
|
out = append(out, arr[i].Word)
|
|
}
|
|
return out
|
|
}
|