Embed a question and a stored note differently (Vikunja #371)
Note recall is asymmetric: a short question goes in, a longer note comes out. Adds EmbedQuery/EmbedPassage helpers and the e5 prefixes, and points the note/fact write path at the passage side and the query path at the query side. Reviewers: the three call sites in voice.go. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01CGeSZxh1DCtRxmFVSYVGvJ
This commit is contained in:
+3
-3
@@ -567,7 +567,7 @@ func (h *reactiveHandler) applyAction(ctx context.Context, dec router.Decision)
|
|||||||
// fail the fact write). Facts aren't in the notes table, so this is the
|
// fail the fact write). Facts aren't in the notes table, so this is the
|
||||||
// only recall path for them — "когда я пил воду?" reads back from here.
|
// only recall path for them — "когда я пил воду?" reads back from here.
|
||||||
if h.memStore != nil {
|
if h.memStore != nil {
|
||||||
if vec, err := h.embedder.Embed(ctx, dec.Utterance); err != nil {
|
if vec, err := router.EmbedPassage(ctx, h.embedder, dec.Utterance); err != nil {
|
||||||
log.Printf("voice: embed fact for memory: %v", err)
|
log.Printf("voice: embed fact for memory: %v", err)
|
||||||
} else if err := h.memStore.Insert(ctx, "fact:"+dec.Slots.Key+":"+strconv.FormatInt(now.Unix(), 10), vec, map[string]string{
|
} else if err := h.memStore.Insert(ctx, "fact:"+dec.Slots.Key+":"+strconv.FormatInt(now.Unix(), 10), vec, map[string]string{
|
||||||
"source": "voice",
|
"source": "voice",
|
||||||
@@ -681,7 +681,7 @@ func (h *reactiveHandler) applyAction(ctx context.Context, dec router.Decision)
|
|||||||
// embed the note text with the same model the classifier uses, persist
|
// embed the note text with the same model the classifier uses, persist
|
||||||
// via CoreAPI (source=tap:voice). Semantic recall lives in `notes`, not
|
// via CoreAPI (source=tap:voice). Semantic recall lives in `notes`, not
|
||||||
// facts — no predicate reads it (spec's two-memory split).
|
// facts — no predicate reads it (spec's two-memory split).
|
||||||
vec, err := h.embedder.Embed(ctx, dec.Utterance)
|
vec, err := router.EmbedPassage(ctx, h.embedder, dec.Utterance)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Printf("voice: embed note: %v", err)
|
log.Printf("voice: embed note: %v", err)
|
||||||
return "не получилось сохранить заметку."
|
return "не получилось сохранить заметку."
|
||||||
@@ -758,7 +758,7 @@ func (h *reactiveHandler) applyAction(ctx context.Context, dec router.Decision)
|
|||||||
return fmt.Sprintf("в %s сейчас %.0f градусов, %s.", w.Location, w.Temperature, w.Condition)
|
return fmt.Sprintf("в %s сейчас %.0f градусов, %s.", w.Location, w.Temperature, w.Condition)
|
||||||
}
|
}
|
||||||
|
|
||||||
vec, err := h.embedder.Embed(ctx, dec.Utterance)
|
vec, err := router.EmbedQuery(ctx, h.embedder, dec.Utterance)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Printf("voice: embed query: %v", err)
|
log.Printf("voice: embed query: %v", err)
|
||||||
return "не получилось найти ответ."
|
return "не получилось найти ответ."
|
||||||
|
|||||||
@@ -19,6 +19,37 @@ type Embedder interface {
|
|||||||
Close() error
|
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
|
// HashEmbedder — a deterministic bag-of-words embedder used for tests and as a
|
||||||
// non-zero default floor. NOT semantically meaningful across languages; the
|
// non-zero default floor. NOT semantically meaningful across languages; the
|
||||||
// real classifier swaps in the multilingual ONNX model wholesale.
|
// 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))
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -12,6 +12,15 @@ import (
|
|||||||
"golang.org/x/text/unicode/norm"
|
"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 (
|
const (
|
||||||
padTokenID = 1
|
padTokenID = 1
|
||||||
unkTokenID = 3
|
unkTokenID = 3
|
||||||
@@ -54,7 +63,25 @@ func NewONNXEmbedder(modelPath, tokenizerPath, libPath string) (*onnxEmbedder, e
|
|||||||
|
|
||||||
func (e *onnxEmbedder) Dim() int { return embedDim }
|
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) {
|
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)
|
inputIDs, attentionMask, _ := e.tokenizer.Encode(text)
|
||||||
|
|
||||||
inputShape := ort.NewShape(1, int64(maxLength))
|
inputShape := ort.NewShape(1, int64(maxLength))
|
||||||
@@ -313,4 +340,4 @@ func preTokenize(text string) []string {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
var _ Embedder = (*onnxEmbedder)(nil)
|
var _ AsymmetricEmbedder = (*onnxEmbedder)(nil)
|
||||||
|
|||||||
Reference in New Issue
Block a user