108 lines
3.1 KiB
Go
108 lines
3.1 KiB
Go
// Command reembed gives every chunk in a tenant a vector from the current
|
|
// embedding model.
|
|
//
|
|
// Run it after changing EMBED_PROVIDER or EMBED_MODEL. The reason it is a
|
|
// command and not something that happens automatically is that it costs
|
|
// real time and, on a hosted provider, real money — and doing that silently on
|
|
// a config change is how a deployment surprises somebody with a bill.
|
|
//
|
|
// The reason it EXISTS is that the alternative is silent too, in the worse
|
|
// direction: vectors from two models are not comparable, so after a switch the
|
|
// old ones simply stop being searched. Retrieval keeps working, keeps citing,
|
|
// and quietly halves its own recall. Nothing errors.
|
|
//
|
|
// make reembed ORG=<slug>
|
|
package main
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"flag"
|
|
"fmt"
|
|
"os"
|
|
"time"
|
|
|
|
"github.com/jackc/pgx/v5"
|
|
|
|
"github.com/krow/krow-backend/go-api/internal/config"
|
|
"github.com/krow/krow-backend/go-api/internal/db"
|
|
"github.com/krow/krow-backend/go-api/internal/knowledge"
|
|
"github.com/krow/krow-backend/go-api/internal/runtime"
|
|
)
|
|
|
|
func main() {
|
|
var (
|
|
org = flag.String("org", "", "organization slug to re-embed (required)")
|
|
batch = flag.Int("batch", 32, "chunks per request to the embedding model")
|
|
timeout = flag.Duration("timeout", 30*time.Minute, "overall timeout")
|
|
)
|
|
flag.Parse()
|
|
|
|
if err := run(*org, *batch, *timeout); err != nil {
|
|
fmt.Fprintf(os.Stderr, "reembed: %v\n", err)
|
|
os.Exit(1)
|
|
}
|
|
}
|
|
|
|
func run(orgSlug string, batch int, timeout time.Duration) error {
|
|
if orgSlug == "" {
|
|
return errors.New("an organization is required: --org=<slug>")
|
|
}
|
|
|
|
cfg, err := config.Load()
|
|
if err != nil {
|
|
return fmt.Errorf("load configuration: %w", err)
|
|
}
|
|
|
|
embedder := runtime.NewEmbedder(*cfg)
|
|
if embedder == nil {
|
|
return errors.New("no embedding model is configured; set EMBED_PROVIDER " +
|
|
"(and its model) before re-embedding")
|
|
}
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
|
defer cancel()
|
|
|
|
database, err := db.Open(ctx, cfg.DB)
|
|
if err != nil {
|
|
return fmt.Errorf("connect: %w", err)
|
|
}
|
|
defer database.Close()
|
|
|
|
var orgID string
|
|
err = database.Pool.QueryRow(ctx,
|
|
`SELECT id::text FROM organizations WHERE slug = $1`, orgSlug).Scan(&orgID)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return fmt.Errorf("no organization with slug %q", orgSlug)
|
|
}
|
|
if err != nil {
|
|
return fmt.Errorf("resolve organization: %w", err)
|
|
}
|
|
|
|
fmt.Printf("re-embedding %s with %s\n", orgSlug, embedder.Model())
|
|
|
|
started := time.Now()
|
|
last := 0
|
|
done, err := knowledge.NewIngester(database.Pool, embedder).
|
|
Reembed(ctx, orgID, batch, func(d, total int) {
|
|
// Reported as it goes. A corpus takes long enough that a silent
|
|
// command is one somebody kills halfway, which is the worst place
|
|
// to stop.
|
|
if d-last >= batch || d == total {
|
|
fmt.Printf(" %d/%d chunks (%s elapsed)\n", d, total,
|
|
time.Since(started).Round(time.Second))
|
|
last = d
|
|
}
|
|
})
|
|
if err != nil {
|
|
return fmt.Errorf("after %d chunk(s): %w", done, err)
|
|
}
|
|
|
|
if done == 0 {
|
|
fmt.Println("nothing to do — every chunk already carries this model's vectors")
|
|
return nil
|
|
}
|
|
fmt.Printf("\n%d chunk(s) re-embedded in %s\n", done, time.Since(started).Round(time.Second))
|
|
return nil
|
|
}
|