Merge branch 'worktree-agent-ad5da57e47b822152' into overnight-jul31
This commit is contained in:
@@ -117,18 +117,40 @@ type cachingEmbedder struct {
|
||||
seen map[string][]float32
|
||||
}
|
||||
|
||||
var _ router.AsymmetricEmbedder = (*cachingEmbedder)(nil)
|
||||
|
||||
func (c *cachingEmbedder) Dim() int { return c.inner.Dim() }
|
||||
func (c *cachingEmbedder) Close() error { return nil } // the caller owns inner
|
||||
|
||||
func (c *cachingEmbedder) Embed(ctx context.Context, text string) ([]float32, error) {
|
||||
if v, ok := c.seen[text]; ok {
|
||||
return c.cached(ctx, "embed:"+text, func() ([]float32, error) {
|
||||
return c.inner.Embed(ctx, text)
|
||||
})
|
||||
}
|
||||
|
||||
// The two sides of an asymmetric embedder give different vectors for the same
|
||||
// string, so the cache key has to say which side asked.
|
||||
func (c *cachingEmbedder) EmbedQuery(ctx context.Context, text string) ([]float32, error) {
|
||||
return c.cached(ctx, "query:"+text, func() ([]float32, error) {
|
||||
return router.EmbedQuery(ctx, c.inner, text)
|
||||
})
|
||||
}
|
||||
|
||||
func (c *cachingEmbedder) EmbedPassage(ctx context.Context, text string) ([]float32, error) {
|
||||
return c.cached(ctx, "passage:"+text, func() ([]float32, error) {
|
||||
return router.EmbedPassage(ctx, c.inner, text)
|
||||
})
|
||||
}
|
||||
|
||||
func (c *cachingEmbedder) cached(_ context.Context, key string, embed func() ([]float32, error)) ([]float32, error) {
|
||||
if v, ok := c.seen[key]; ok {
|
||||
return v, nil
|
||||
}
|
||||
v, err := c.inner.Embed(ctx, text)
|
||||
v, err := embed()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.seen[text] = v
|
||||
c.seen[key] = v
|
||||
return v, nil
|
||||
}
|
||||
|
||||
@@ -305,7 +327,7 @@ func scoreCase(ctx context.Context, emb router.Embedder, newStore NewStore, minS
|
||||
|
||||
all := append(append([]StoredNote(nil), c.Notes...), filler...)
|
||||
for _, n := range all {
|
||||
vec, err := emb.Embed(ctx, n.Text)
|
||||
vec, err := router.EmbedPassage(ctx, emb, n.Text)
|
||||
if err != nil {
|
||||
return Outcome{}, fmt.Errorf("%s: embed note %s: %w", c.ID, n.ID, err)
|
||||
}
|
||||
@@ -317,7 +339,7 @@ func scoreCase(ctx context.Context, emb router.Embedder, newStore NewStore, minS
|
||||
|
||||
o := Outcome{Case: c}
|
||||
start := time.Now()
|
||||
qvec, err := emb.Embed(ctx, c.Query)
|
||||
qvec, err := router.EmbedQuery(ctx, emb, c.Query)
|
||||
if err != nil {
|
||||
o.Latency = time.Since(start)
|
||||
o.Err = err
|
||||
|
||||
@@ -246,8 +246,8 @@ func TestONNXRecall(t *testing.T) {
|
||||
if lib == "" {
|
||||
t.Skip("MAVEN_ONNX_LIB unset — see AGENTS.md § Embedder model for intent routing")
|
||||
}
|
||||
model := filepath.Join("../../..", "models/embedder/model.onnx")
|
||||
tok := filepath.Join("../../..", "models/embedder/tokenizer.json")
|
||||
model := filepath.Join("../../..", "models/embedder/multilingual-e5-small/model_quantized.onnx")
|
||||
tok := filepath.Join("../../..", "models/embedder/multilingual-e5-small/tokenizer.json")
|
||||
for _, p := range []string{lib, model, tok} {
|
||||
if _, err := os.Stat(p); err != nil {
|
||||
t.Skipf("missing %s: %v", p, err)
|
||||
|
||||
@@ -19,6 +19,37 @@ type Embedder interface {
|
||||
Close() error
|
||||
}
|
||||
|
||||
// AsymmetricEmbedder — an embedder that wants to know whether a text is a
|
||||
// search query or a stored passage. Recall is asymmetric: a short question
|
||||
// goes in, a longer note comes out. The e5 family is trained for exactly that
|
||||
// and needs the side written into the text ("query: " / "passage: ").
|
||||
//
|
||||
// Optional on purpose: HashEmbedder has no such notion, so callers go through
|
||||
// EmbedQuery and EmbedPassage below, which fall back to plain Embed.
|
||||
type AsymmetricEmbedder interface {
|
||||
Embedder
|
||||
EmbedQuery(ctx context.Context, text string) ([]float32, error)
|
||||
EmbedPassage(ctx context.Context, text string) ([]float32, error)
|
||||
}
|
||||
|
||||
// EmbedQuery embeds text that is being searched WITH — a question.
|
||||
func EmbedQuery(ctx context.Context, e Embedder, text string) ([]float32, error) {
|
||||
if a, ok := e.(AsymmetricEmbedder); ok {
|
||||
return a.EmbedQuery(ctx, text)
|
||||
}
|
||||
return e.Embed(ctx, text)
|
||||
}
|
||||
|
||||
// EmbedPassage embeds text that is being searched FOR — a note or a fact on
|
||||
// its way into the store. Store and lookup must use these two calls, not one
|
||||
// of them twice, or the asymmetry buys nothing.
|
||||
func EmbedPassage(ctx context.Context, e Embedder, text string) ([]float32, error) {
|
||||
if a, ok := e.(AsymmetricEmbedder); ok {
|
||||
return a.EmbedPassage(ctx, text)
|
||||
}
|
||||
return e.Embed(ctx, text)
|
||||
}
|
||||
|
||||
// HashEmbedder — a deterministic bag-of-words embedder used for tests and as a
|
||||
// non-zero default floor. NOT semantically meaningful across languages; the
|
||||
// real classifier swaps in the multilingual ONNX model wholesale.
|
||||
|
||||
@@ -37,3 +37,60 @@ func TestHashEmbedderCyrillic(t *testing.T) {
|
||||
t.Fatalf("cosine(shared)=%.3f not > cosine(disjoint)=%.3f", cosine(a, b), cosine(a, c))
|
||||
}
|
||||
}
|
||||
|
||||
// recordingEmbedder â an asymmetric embedder that only remembers which side
|
||||
// was asked for. Enough to pin the dispatch; real vectors need the model.
|
||||
type recordingEmbedder struct{ calls []string }
|
||||
|
||||
func (r *recordingEmbedder) Dim() int { return 2 }
|
||||
func (r *recordingEmbedder) Close() error { return nil }
|
||||
|
||||
func (r *recordingEmbedder) Embed(_ context.Context, _ string) ([]float32, error) {
|
||||
r.calls = append(r.calls, "embed")
|
||||
return []float32{1, 0}, nil
|
||||
}
|
||||
|
||||
func (r *recordingEmbedder) EmbedQuery(_ context.Context, _ string) ([]float32, error) {
|
||||
r.calls = append(r.calls, "query")
|
||||
return []float32{1, 0}, nil
|
||||
}
|
||||
|
||||
func (r *recordingEmbedder) EmbedPassage(_ context.Context, _ string) ([]float32, error) {
|
||||
r.calls = append(r.calls, "passage")
|
||||
return []float32{0, 1}, nil
|
||||
}
|
||||
|
||||
// TestEmbedQueryAndPassageSplit â a question and a stored note must not take
|
||||
// the same path. If both ended up on the same call the asymmetric model buys
|
||||
// nothing, which is the whole reason for the swap.
|
||||
func TestEmbedQueryAndPassageSplit(t *testing.T) {
|
||||
rec := &recordingEmbedder{}
|
||||
if _, err := EmbedQuery(context.Background(), rec, "где логи?"); err != nil {
|
||||
t.Fatalf("EmbedQuery: %v", err)
|
||||
}
|
||||
if _, err := EmbedPassage(context.Background(), rec, "логи в /var/log"); err != nil {
|
||||
t.Fatalf("EmbedPassage: %v", err)
|
||||
}
|
||||
if len(rec.calls) != 2 || rec.calls[0] != "query" || rec.calls[1] != "passage" {
|
||||
t.Errorf("calls %v, want [query passage]", rec.calls)
|
||||
}
|
||||
}
|
||||
|
||||
// TestEmbedFallsBackToPlainEmbed â HashEmbedder has no sides, so both helpers
|
||||
// must still work and give the same vector.
|
||||
func TestEmbedFallsBackToPlainEmbed(t *testing.T) {
|
||||
h := NewHashEmbedder(64)
|
||||
q, err := EmbedQuery(context.Background(), h, "text")
|
||||
if err != nil {
|
||||
t.Fatalf("EmbedQuery: %v", err)
|
||||
}
|
||||
p, err := EmbedPassage(context.Background(), h, "text")
|
||||
if err != nil {
|
||||
t.Fatalf("EmbedPassage: %v", err)
|
||||
}
|
||||
for i := range q {
|
||||
if q[i] != p[i] {
|
||||
t.Fatalf("hash embedder gave two different vectors for the same text")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -185,8 +185,8 @@ func TestONNXBaseline(t *testing.T) {
|
||||
if lib == "" {
|
||||
t.Skip("MAVEN_ONNX_LIB unset — see AGENTS.md § Embedder model for intent routing")
|
||||
}
|
||||
model := filepath.Join("../../..", "models/embedder/model.onnx")
|
||||
tok := filepath.Join("../../..", "models/embedder/tokenizer.json")
|
||||
model := filepath.Join("../../..", "models/embedder/multilingual-e5-small/model_quantized.onnx")
|
||||
tok := filepath.Join("../../..", "models/embedder/multilingual-e5-small/tokenizer.json")
|
||||
for _, p := range []string{lib, model, tok} {
|
||||
if _, err := os.Stat(p); err != nil {
|
||||
t.Skipf("missing %s: %v", p, err)
|
||||
|
||||
@@ -12,6 +12,15 @@ import (
|
||||
"golang.org/x/text/unicode/norm"
|
||||
)
|
||||
|
||||
// The deployed model is multilingual-e5-small. e5 was trained with these two
|
||||
// words glued to the front of every text, and it scores badly without them —
|
||||
// they are part of the model, not a style choice. Swapping back to a symmetric
|
||||
// paraphrase model means dropping them again.
|
||||
const (
|
||||
queryPrefix = "query: "
|
||||
passagePrefix = "passage: "
|
||||
)
|
||||
|
||||
const (
|
||||
padTokenID = 1
|
||||
unkTokenID = 3
|
||||
@@ -54,7 +63,25 @@ func NewONNXEmbedder(modelPath, tokenizerPath, libPath string) (*onnxEmbedder, e
|
||||
|
||||
func (e *onnxEmbedder) Dim() int { return embedDim }
|
||||
|
||||
// Embed treats the text as a query. The classifier compares one short
|
||||
// utterance to another short seed phrase, so both sides get the same prefix
|
||||
// and the comparison stays fair. The recall path must call EmbedQuery and
|
||||
// EmbedPassage instead.
|
||||
func (e *onnxEmbedder) Embed(ctx context.Context, text string) ([]float32, error) {
|
||||
return e.embed(ctx, queryPrefix+text)
|
||||
}
|
||||
|
||||
// EmbedQuery — the question the user just asked.
|
||||
func (e *onnxEmbedder) EmbedQuery(ctx context.Context, text string) ([]float32, error) {
|
||||
return e.embed(ctx, queryPrefix+text)
|
||||
}
|
||||
|
||||
// EmbedPassage — a note or fact being stored, or re-scored at lookup time.
|
||||
func (e *onnxEmbedder) EmbedPassage(ctx context.Context, text string) ([]float32, error) {
|
||||
return e.embed(ctx, passagePrefix+text)
|
||||
}
|
||||
|
||||
func (e *onnxEmbedder) embed(ctx context.Context, text string) ([]float32, error) {
|
||||
inputIDs, attentionMask, _ := e.tokenizer.Encode(text)
|
||||
|
||||
inputShape := ort.NewShape(1, int64(maxLength))
|
||||
@@ -313,4 +340,4 @@ func preTokenize(text string) []string {
|
||||
return out
|
||||
}
|
||||
|
||||
var _ Embedder = (*onnxEmbedder)(nil)
|
||||
var _ AsymmetricEmbedder = (*onnxEmbedder)(nil)
|
||||
|
||||
Reference in New Issue
Block a user