Guest User

Untitled

a guest
Apr 16th, 2023
243
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Go 6.89 KB | None | 0 0
  1. const DIM = 1536
  2.  
  3. type embedding struct {
  4.     name string
  5.     size []int
  6.     f    [][DIM]float64
  7.     f2   []float64
  8. }
  9.  
  10. func newEmbeddingJSON(name string) *embedding {
  11.     r, err := os.Open(name)
  12.     if err != nil {
  13.         panic(err)
  14.     }
  15.     defer r.Close()
  16.     res := &embedding{name: name}
  17.     if err := json.NewDecoder(r).Decode(&res.f); err != nil {
  18.         return nil
  19.     }
  20.     res.updatef2()
  21.     return res
  22. }
  23.  
  24. func newEmbeddingBytes(name string, size []int, b []byte) *embedding {
  25.     res := &embedding{name: name}
  26.     res.f = make([][DIM]float64, len(size), len(size))
  27.     res.size = size
  28.     for i := 0; i < len(size); i++ {
  29.         res.f[i] = *(*[DIM]float64)(unsafe.Pointer(&b[i*DIM*8]))
  30.     }
  31.     res.updatef2()
  32.     return res
  33. }
  34.  
  35. func (e *embedding) updatef2() {
  36.     if e.f2 == nil {
  37.         e.f2 = make([]float64, len(e.f), len(e.f))
  38.     }
  39.     for i := 0; i < len(e.f); i++ {
  40.         e.f2[i] = 0
  41.         for j := 0; j < DIM; j++ {
  42.             e.f2[i] += e.f[i][j] * e.f[i][j]
  43.         }
  44.     }
  45. }
  46.  
  47. func (e *embedding) Bytes() [][]byte {
  48.     res := make([][]byte, len(e.f))
  49.     for i := range e.f {
  50.         res[i] = (*[DIM * 8]byte)(unsafe.Pointer(&e.f[i][0]))[:]
  51.     }
  52.     return res
  53. }
  54.  
  55. func (e *embedding) hash() []byte {
  56.     var b [8]byte
  57.     h := sha256.New()
  58.     h.Write([]byte(e.name))
  59.     for _, v := range e.f {
  60.         for i := 0; i < DIM; i++ {
  61.             var u uint64 = *(*uint64)(unsafe.Pointer(&v[i]))
  62.             binary.BigEndian.PutUint64(b[:], u)
  63.             h.Write(b[:])
  64.         }
  65.     }
  66.     return h.Sum(nil)
  67. }
  68.  
  69. func (a *embedding) cosine(b *embedding) float64 {
  70.     var num, den float64
  71.     for ai := 0; ai < len(a.f); ai++ {
  72.         for bi := 0; bi < len(b.f); bi++ {
  73.             var dot float64
  74.             for i := 0; i < DIM; i++ {
  75.                 dot += a.f[ai][i] * b.f[bi][i]
  76.             }
  77.             c := dot / (a.f2[ai] * b.f2[bi])
  78.             if c < 0.5 {
  79.                 c = 0
  80.             }
  81.             if a.size != nil {
  82.                 num += c * float64(b.size[bi]) * float64(a.size[ai])
  83.                 den += float64(b.size[bi]) * float64(a.size[ai])
  84.             } else {
  85.                 num += c * float64(b.size[bi])
  86.                 den += float64(b.size[bi])
  87.             }
  88.         }
  89.     }
  90.     return num / den
  91. }
  92.  
  93. type db struct {
  94.     f *os.File
  95.     e []*embedding
  96. }
  97.  
  98. type encodedMetadata struct {
  99.     Name string `json:"name"`
  100.     Size []int  `json:"size"`
  101. }
  102.  
  103. func newDB(path string, checkHash bool) *db {
  104.     f, err := os.OpenFile(path, os.O_RDONLY, 0o755)
  105.     if err != nil {
  106.         panic(err)
  107.     }
  108.     m, err := mmap.Map(f, mmap.RDONLY, 0)
  109.     if err != nil {
  110.         panic(err)
  111.     }
  112.     b := []byte(m)
  113.  
  114.     hash := b[:sha256.Size]
  115.     // fmt.Printf("deserialize: hash1: %x\n", hash)
  116.     b = b[sha256.Size:]
  117.  
  118.     n := int64(binary.BigEndian.Uint64(b[:8]))
  119.     // fmt.Printf("deserialize: n: %d\n", n)
  120.     b = b[8:]
  121.  
  122.     var meta []encodedMetadata
  123.     mb := b[:n]
  124.     for mb[len(mb)-1] == 0 {
  125.         mb = mb[:len(mb)-1]
  126.     }
  127.     if err := json.Unmarshal(mb, &meta); err != nil {
  128.         panic(err)
  129.     }
  130.     b = b[n:]
  131.  
  132.     db := &db{f: f}
  133.     for _, m := range meta {
  134.         db.e = append(db.e, newEmbeddingBytes(m.Name, m.Size, b))
  135.         b = b[len(m.Size)*DIM*8:]
  136.     }
  137.     if checkHash {
  138.         h := sha256.New()
  139.         for _, e := range db.e {
  140.             eh := e.hash()
  141.             h.Write(eh)
  142.         }
  143.         // fmt.Printf("deserialize: hash2: %x\n", h.Sum(nil))
  144.         if !strings.EqualFold(fmt.Sprintf("%x", h.Sum(nil)), fmt.Sprintf("%x", hash)) {
  145.             panic(fmt.Sprintf("hash mismatch: %x != %x", h.Sum(nil), hash))
  146.         }
  147.     }
  148.     return db
  149. }
  150.  
  151. func newDBFromJSON(path string) *db {
  152.     db := &db{}
  153.     fmt.Printf("reading embeddings...\n")
  154.     begin := time.Now()
  155.     fs.WalkDir(os.DirFS("."), ".", func(path string, d fs.DirEntry, err error) error {
  156.         if err != nil {
  157.             return err
  158.         }
  159.         if !strings.HasSuffix(path, ".json") {
  160.             return nil
  161.         }
  162.         e := newEmbeddingJSON(path)
  163.         if e == nil {
  164.             return nil
  165.         }
  166.         db.e = append(db.e, e)
  167.         if len(db.e)%100 == 0 {
  168.             fmt.Printf("%d... ", len(db.e))
  169.         }
  170.         return nil
  171.     })
  172.     fmt.Printf("%s (%d)\n", time.Since(begin).Round(time.Millisecond), len(db.e))
  173.     return db
  174. }
  175.  
  176. func (db *db) Close() error {
  177.     return db.f.Close()
  178. }
  179.  
  180. func (db *db) serialize(path string) {
  181.     f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR|os.O_TRUNC, 0o755)
  182.     if err != nil {
  183.         panic(err)
  184.     }
  185.     defer f.Close()
  186.     h := sha256.New()
  187.     for _, e := range db.e {
  188.         h.Write(e.hash())
  189.     }
  190.     sha := h.Sum(nil)
  191.     fmt.Printf("serialize: hash: %x\n", sha)
  192.     if _, err := f.Write(sha); err != nil {
  193.         panic(err)
  194.     }
  195.     var meta []encodedMetadata
  196.     for _, e := range db.e {
  197.         m := encodedMetadata{Name: e.name}
  198.         b, err := os.ReadFile("embed_req/" + e.name)
  199.         if err != nil {
  200.             panic(err)
  201.         }
  202.         var lines []string
  203.         if err := json.Unmarshal(b, &lines); err != nil {
  204.             panic(err)
  205.         }
  206.         for _, l := range lines {
  207.             m.Size = append(m.Size, len(l))
  208.         }
  209.         meta = append(meta, m)
  210.     }
  211.     mb, err := json.Marshal(meta)
  212.     if err != nil {
  213.         panic(err)
  214.     }
  215.     for len(mb)%8 != 0 {
  216.         mb = append(mb, 0)
  217.     }
  218.     b := make([]byte, 8)
  219.     binary.BigEndian.PutUint64(b, uint64(len(mb)))
  220.     if _, err := f.Write(b); err != nil {
  221.         panic(err)
  222.     }
  223.     if _, err := f.Write(mb); err != nil {
  224.         panic(err)
  225.     }
  226.     fmt.Printf("serialize: n=%d\n", len(mb))
  227.     for _, e := range db.e {
  228.         for _, b := range e.Bytes() {
  229.             if _, err := f.Write(b); err != nil {
  230.                 panic(err)
  231.             }
  232.         }
  233.     }
  234. }
  235.  
  236. func convert(out string) {
  237.     db := newDBFromJSON(".")
  238.     defer db.Close()
  239.     db.serialize(out)
  240. }
  241.  
  242. func (db *db) search(e *embedding, k int) {
  243.     W := 16
  244.     fmt.Printf("cosine similarities (%d emb, %d wrk, %d vec)... ", len(db.e), W, len(e.f))
  245.     begin := time.Now()
  246.     results := make([]struct {
  247.         e *embedding
  248.         s float64
  249.     }, len(db.e), len(db.e))
  250.     var wg sync.WaitGroup
  251.     wg.Add(W)
  252.     for wi := 0; wi < W; wi++ {
  253.         go func(wi int) {
  254.             defer wg.Done()
  255.             for ei := wi; ei < len(db.e); ei += W {
  256.                 if db.e[ei].name == e.name {
  257.                     continue
  258.                 }
  259.                 results[ei].e = db.e[ei]
  260.                 results[ei].s = e.cosine(db.e[ei])
  261.             }
  262.         }(wi)
  263.     }
  264.     wg.Wait()
  265.     fmt.Printf("%s\n", time.Since(begin).Round(time.Millisecond))
  266.  
  267.     fmt.Printf("sorting... ")
  268.     begin = time.Now()
  269.     sort.Slice(results, func(i, j int) bool {
  270.         return results[i].s > results[j].s
  271.     })
  272.     fmt.Printf("%s\n", time.Since(begin).Round(time.Millisecond))
  273.  
  274.     if len(results) < k {
  275.         k = len(results)
  276.     }
  277.     for i := 0; i < k; i++ {
  278.         summarize(results[i].e.name, &results[i].s)
  279.     }
  280. }
  281.  
  282. func summarize(name string, score *float64) {
  283.     fmt.Printf("* \033[0;34m%s\033[0m", name)
  284.     if score != nil {
  285.         fmt.Printf(" (%.2f)", *score)
  286.     }
  287.     fmt.Printf("\n")
  288.     b, err := os.ReadFile("embed_req/" + name)
  289.     if err != nil {
  290.         panic(err)
  291.     }
  292.     var lines []string
  293.     if err := json.Unmarshal(b, &lines); err != nil {
  294.         panic(err)
  295.     }
  296.     s := strings.Join(lines, " // ")
  297.     s = strings.ReplaceAll(s, "\n", " ")
  298.     if len(s) > 210 {
  299.         s = s[:210] + "…"
  300.     }
  301.     s = wrapText(s, 80)
  302.     s = "  " + strings.ReplaceAll(s, "\n", "\n  ")
  303.     fmt.Printf("%s\n\n", s)
  304. }
  305.  
  306. func wrapText(text string, colWidth int) string {
  307.     var buf bytes.Buffer
  308.     words := strings.Fields(text)
  309.     lineLen := 0
  310.     for i, word := range words {
  311.         if i > 0 {
  312.             buf.WriteByte(' ')
  313.             lineLen++
  314.         }
  315.         if lineLen+len(word) > colWidth {
  316.             buf.WriteByte('\n')
  317.             lineLen = 0
  318.         }
  319.         buf.WriteString(word)
  320.         lineLen += len(word)
  321.     }
  322.     return buf.String()
  323. }
Advertisement
Add Comment
Please, Sign In to add comment