// embedpoi — build the semantic POI index (poi_vec) from the poi table. // // For every named POI (all families except place=*) it embeds // "name — kind" with the local zembed model (llama.cpp /embeddings) in // batches and upserts into poi_vec. Resumable: existing keys are skipped. // After the run it creates the HNSW cosine index. // // go run ./embed -dsn "host=localhost ... dbname=trips sslmode=disable" package main import ( "bytes" "database/sql" "encoding/json" "flag" "fmt" "io" "log" "net/http" "os" "strconv" "strings" "time" _ "github.com/lib/pq" ) func getenvDefault(k, def string) string { if v := os.Getenv(k); v != "" { return v } return def } func must(err error) { if err != nil { log.Fatal(err) } } func exec(db *sql.DB, q string, args ...any) { if _, e := db.Exec(q, args...); e != nil { log.Fatal(e) } } // vecToSQL renders a []float32 as the body of a vector literal: 0.1,-0.2,… func vecToSQL(v []float32) string { parts := make([]string, len(v)) for i, f := range v { parts[i] = strconv.FormatFloat(float64(f), 'g', 6, 32) } return strings.Join(parts, ",") } func main() { dsnFlag := flag.String("dsn", getenvDefault("SPATIAL_DSN", "host=localhost port=5432 user=trips password=trips dbname=trips sslmode=disable"), "DSN") baseFlag := flag.String("embed-base", getenvDefault("EMBED_BASE", "http://192.168.3.7:1234/v1"), "llama.cpp /v1 base") modelFlag := flag.String("embed-model", getenvDefault("EMBED_MODEL", "zembed-1-Q4_K_M"), "embedding model") batchFlag := flag.Int("batch", 64, "embeddings per request") extractFlag := flag.String("extract", "", "restrict to one extract (default: all)") qpsFlag := flag.Int("qps", 4, "requests per second (politeness toward the shared model server)") flag.Parse() db, err := sql.Open("postgres", *dsnFlag) must(err) db.SetMaxOpenConns(4) must(db.Ping()) client := &http.Client{Timeout: 120 * time.Second} embed := func(texts []string) ([][]float32, error) { body, _ := json.Marshal(map[string]any{"model": *modelFlag, "input": texts}) req, err := http.NewRequest("POST", *baseFlag+"/embeddings", bytes.NewReader(body)) if err != nil { return nil, err } req.Header.Set("Content-Type", "application/json") res, err := client.Do(req) if err != nil { return nil, err } defer res.Body.Close() if res.StatusCode != 200 { b, _ := io.ReadAll(io.LimitReader(res.Body, 300)) return nil, fmt.Errorf("embed http %d: %s", res.StatusCode, b) } var j struct { Data []struct { Embedding []float32 `json:"embedding"` } `json:"data"` } if err := json.NewDecoder(res.Body).Decode(&j); err != nil { return nil, err } var out [][]float32 for _, d := range j.Data { out = append(out, d.Embedding) } return out, nil } // ---- collect rows (everything semantic-searchable: named POIs of the // POI families; place=* localities are not "places to go") type row struct { Extract, OsmType, Name, Kind string OsmID int64 } q := `SELECT extract, osm_type, osm_id, name, kind FROM poi WHERE kind NOT LIKE 'place=%'` var args []any if *extractFlag != "" { q += " AND extract = $1" args = append(args, *extractFlag) } q += " ORDER BY extract, osm_id" rows, err := db.Query(q, args...) must(err) var all []row for rows.Next() { var r row must(rows.Scan(&r.Extract, &r.OsmType, &r.OsmID, &r.Name, &r.Kind)) all = append(all, r) } rows.Close() log.Printf("poi rows to consider: %d", len(all)) // resume: keys already embedded have := map[string]bool{} hrows, err := db.Query(`SELECT extract, osm_type, osm_id FROM poi_vec`) if err != nil { have = nil // table doesn't exist yet — everything is to do } else { for hrows.Next() { var e, t string var id int64 hrows.Scan(&e, &t, &id) have[e+"|"+t+"|"+strconv.FormatInt(id, 10)] = true } hrows.Close() } var todo []row for _, r := range all { if have == nil || !have[r.Extract+"|"+r.OsmType+"|"+strconv.FormatInt(r.OsmID, 10)] { todo = append(todo, r) } } log.Printf("already embedded: %d, to do: %d", len(all)-len(todo), len(todo)) if len(todo) == 0 { return } var dim int total := len(todo) for i := 0; i < total; i += *batchFlag { j := i + *batchFlag if j > total { j = total } chunk := todo[i:j] texts := make([]string, len(chunk)) for k, r := range chunk { texts[k] = r.Name + " — " + r.Kind } // the shared llama.cpp host evicts zembed when another model loads // (500 "proxy error") — that clears on the next request, so retry // patiently rather than dying var vecs [][]float32 for attempt := 0; ; attempt++ { var err error vecs, err = embed(texts) if err == nil { break } if attempt >= 20 { must(err) } log.Printf("batch at %d failed (%v) — retry %d/20", i, err, attempt+1) time.Sleep(time.Duration(5+attempt) * time.Second) } if dim == 0 { dim = len(vecs[0]) // derived table: safe to rebuild if the model's dim changed exec(db, `DROP TABLE IF EXISTS poi_vec`) stmt := fmt.Sprintf(`CREATE TABLE poi_vec ( extract text NOT NULL, osm_type char NOT NULL, osm_id bigint NOT NULL, name text NOT NULL, kind text NOT NULL, vec vector(%d) NOT NULL, PRIMARY KEY (extract, osm_type, osm_id))`, dim) log.Printf("embedding dim: %d — creating poi_vec", dim) exec(db, stmt) } for k, r := range chunk { _, err := db.Exec(`INSERT INTO poi_vec (extract, osm_type, osm_id, name, kind, vec) VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT (extract, osm_type, osm_id) DO UPDATE SET vec = EXCLUDED.vec`, r.Extract, r.OsmType, r.OsmID, r.Name, r.Kind, "["+vecToSQL(vecs[k])+"]") must(err) } if i == 0 || i/(*batchFlag*25) > (i-*batchFlag)/(*batchFlag*25) { log.Printf("%d/%d embedded (%.1f%%)", j, total, 100.0*float64(j)/float64(total)) } time.Sleep(time.Second / time.Duration(*qpsFlag)) } log.Printf("creating HNSW index (vector_cosine_ops) on poi_vec…") exec(db, `CREATE INDEX IF NOT EXISTS poi_vec_hnsw ON poi_vec USING hnsw (vec vector_cosine_ops)`) log.Printf("done: %d vectors in poi_vec", len(all)) }