Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- const DIM = 1536
- type embedding struct {
- name string
- size []int
- f [][DIM]float64
- f2 []float64
- }
- func newEmbeddingJSON(name string) *embedding {
- r, err := os.Open(name)
- if err != nil {
- panic(err)
- }
- defer r.Close()
- res := &embedding{name: name}
- if err := json.NewDecoder(r).Decode(&res.f); err != nil {
- return nil
- }
- res.updatef2()
- return res
- }
- func newEmbeddingBytes(name string, size []int, b []byte) *embedding {
- res := &embedding{name: name}
- res.f = make([][DIM]float64, len(size), len(size))
- res.size = size
- for i := 0; i < len(size); i++ {
- res.f[i] = *(*[DIM]float64)(unsafe.Pointer(&b[i*DIM*8]))
- }
- res.updatef2()
- return res
- }
- func (e *embedding) updatef2() {
- if e.f2 == nil {
- e.f2 = make([]float64, len(e.f), len(e.f))
- }
- for i := 0; i < len(e.f); i++ {
- e.f2[i] = 0
- for j := 0; j < DIM; j++ {
- e.f2[i] += e.f[i][j] * e.f[i][j]
- }
- }
- }
- func (e *embedding) Bytes() [][]byte {
- res := make([][]byte, len(e.f))
- for i := range e.f {
- res[i] = (*[DIM * 8]byte)(unsafe.Pointer(&e.f[i][0]))[:]
- }
- return res
- }
- func (e *embedding) hash() []byte {
- var b [8]byte
- h := sha256.New()
- h.Write([]byte(e.name))
- for _, v := range e.f {
- for i := 0; i < DIM; i++ {
- var u uint64 = *(*uint64)(unsafe.Pointer(&v[i]))
- binary.BigEndian.PutUint64(b[:], u)
- h.Write(b[:])
- }
- }
- return h.Sum(nil)
- }
- func (a *embedding) cosine(b *embedding) float64 {
- var num, den float64
- for ai := 0; ai < len(a.f); ai++ {
- for bi := 0; bi < len(b.f); bi++ {
- var dot float64
- for i := 0; i < DIM; i++ {
- dot += a.f[ai][i] * b.f[bi][i]
- }
- c := dot / (a.f2[ai] * b.f2[bi])
- if c < 0.5 {
- c = 0
- }
- if a.size != nil {
- num += c * float64(b.size[bi]) * float64(a.size[ai])
- den += float64(b.size[bi]) * float64(a.size[ai])
- } else {
- num += c * float64(b.size[bi])
- den += float64(b.size[bi])
- }
- }
- }
- return num / den
- }
- type db struct {
- f *os.File
- e []*embedding
- }
- type encodedMetadata struct {
- Name string `json:"name"`
- Size []int `json:"size"`
- }
- func newDB(path string, checkHash bool) *db {
- f, err := os.OpenFile(path, os.O_RDONLY, 0o755)
- if err != nil {
- panic(err)
- }
- m, err := mmap.Map(f, mmap.RDONLY, 0)
- if err != nil {
- panic(err)
- }
- b := []byte(m)
- hash := b[:sha256.Size]
- // fmt.Printf("deserialize: hash1: %x\n", hash)
- b = b[sha256.Size:]
- n := int64(binary.BigEndian.Uint64(b[:8]))
- // fmt.Printf("deserialize: n: %d\n", n)
- b = b[8:]
- var meta []encodedMetadata
- mb := b[:n]
- for mb[len(mb)-1] == 0 {
- mb = mb[:len(mb)-1]
- }
- if err := json.Unmarshal(mb, &meta); err != nil {
- panic(err)
- }
- b = b[n:]
- db := &db{f: f}
- for _, m := range meta {
- db.e = append(db.e, newEmbeddingBytes(m.Name, m.Size, b))
- b = b[len(m.Size)*DIM*8:]
- }
- if checkHash {
- h := sha256.New()
- for _, e := range db.e {
- eh := e.hash()
- h.Write(eh)
- }
- // fmt.Printf("deserialize: hash2: %x\n", h.Sum(nil))
- if !strings.EqualFold(fmt.Sprintf("%x", h.Sum(nil)), fmt.Sprintf("%x", hash)) {
- panic(fmt.Sprintf("hash mismatch: %x != %x", h.Sum(nil), hash))
- }
- }
- return db
- }
- func newDBFromJSON(path string) *db {
- db := &db{}
- fmt.Printf("reading embeddings...\n")
- begin := time.Now()
- fs.WalkDir(os.DirFS("."), ".", func(path string, d fs.DirEntry, err error) error {
- if err != nil {
- return err
- }
- if !strings.HasSuffix(path, ".json") {
- return nil
- }
- e := newEmbeddingJSON(path)
- if e == nil {
- return nil
- }
- db.e = append(db.e, e)
- if len(db.e)%100 == 0 {
- fmt.Printf("%d... ", len(db.e))
- }
- return nil
- })
- fmt.Printf("%s (%d)\n", time.Since(begin).Round(time.Millisecond), len(db.e))
- return db
- }
- func (db *db) Close() error {
- return db.f.Close()
- }
- func (db *db) serialize(path string) {
- f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR|os.O_TRUNC, 0o755)
- if err != nil {
- panic(err)
- }
- defer f.Close()
- h := sha256.New()
- for _, e := range db.e {
- h.Write(e.hash())
- }
- sha := h.Sum(nil)
- fmt.Printf("serialize: hash: %x\n", sha)
- if _, err := f.Write(sha); err != nil {
- panic(err)
- }
- var meta []encodedMetadata
- for _, e := range db.e {
- m := encodedMetadata{Name: e.name}
- b, err := os.ReadFile("embed_req/" + e.name)
- if err != nil {
- panic(err)
- }
- var lines []string
- if err := json.Unmarshal(b, &lines); err != nil {
- panic(err)
- }
- for _, l := range lines {
- m.Size = append(m.Size, len(l))
- }
- meta = append(meta, m)
- }
- mb, err := json.Marshal(meta)
- if err != nil {
- panic(err)
- }
- for len(mb)%8 != 0 {
- mb = append(mb, 0)
- }
- b := make([]byte, 8)
- binary.BigEndian.PutUint64(b, uint64(len(mb)))
- if _, err := f.Write(b); err != nil {
- panic(err)
- }
- if _, err := f.Write(mb); err != nil {
- panic(err)
- }
- fmt.Printf("serialize: n=%d\n", len(mb))
- for _, e := range db.e {
- for _, b := range e.Bytes() {
- if _, err := f.Write(b); err != nil {
- panic(err)
- }
- }
- }
- }
- func convert(out string) {
- db := newDBFromJSON(".")
- defer db.Close()
- db.serialize(out)
- }
- func (db *db) search(e *embedding, k int) {
- W := 16
- fmt.Printf("cosine similarities (%d emb, %d wrk, %d vec)... ", len(db.e), W, len(e.f))
- begin := time.Now()
- results := make([]struct {
- e *embedding
- s float64
- }, len(db.e), len(db.e))
- var wg sync.WaitGroup
- wg.Add(W)
- for wi := 0; wi < W; wi++ {
- go func(wi int) {
- defer wg.Done()
- for ei := wi; ei < len(db.e); ei += W {
- if db.e[ei].name == e.name {
- continue
- }
- results[ei].e = db.e[ei]
- results[ei].s = e.cosine(db.e[ei])
- }
- }(wi)
- }
- wg.Wait()
- fmt.Printf("%s\n", time.Since(begin).Round(time.Millisecond))
- fmt.Printf("sorting... ")
- begin = time.Now()
- sort.Slice(results, func(i, j int) bool {
- return results[i].s > results[j].s
- })
- fmt.Printf("%s\n", time.Since(begin).Round(time.Millisecond))
- if len(results) < k {
- k = len(results)
- }
- for i := 0; i < k; i++ {
- summarize(results[i].e.name, &results[i].s)
- }
- }
- func summarize(name string, score *float64) {
- fmt.Printf("* \033[0;34m%s\033[0m", name)
- if score != nil {
- fmt.Printf(" (%.2f)", *score)
- }
- fmt.Printf("\n")
- b, err := os.ReadFile("embed_req/" + name)
- if err != nil {
- panic(err)
- }
- var lines []string
- if err := json.Unmarshal(b, &lines); err != nil {
- panic(err)
- }
- s := strings.Join(lines, " // ")
- s = strings.ReplaceAll(s, "\n", " ")
- if len(s) > 210 {
- s = s[:210] + "…"
- }
- s = wrapText(s, 80)
- s = " " + strings.ReplaceAll(s, "\n", "\n ")
- fmt.Printf("%s\n\n", s)
- }
- func wrapText(text string, colWidth int) string {
- var buf bytes.Buffer
- words := strings.Fields(text)
- lineLen := 0
- for i, word := range words {
- if i > 0 {
- buf.WriteByte(' ')
- lineLen++
- }
- if lineLen+len(word) > colWidth {
- buf.WriteByte('\n')
- lineLen = 0
- }
- buf.WriteString(word)
- lineLen += len(word)
- }
- return buf.String()
- }
Advertisement
Add Comment
Please, Sign In to add comment