mirror of
https://github.com/mudler/LocalAI.git
synced 2026-10-09 22:54:42 -04:00
fix(voice): keep one voice store per embedding dimension
The in-memory local-store rejects vectors of another size than the ones it holds. With 192-value voices registered, adding a 256-value voice failed, and identifying with a 256-value probe returned an error. The registry now keeps one store per embedding dimension. The first dimension seen uses the configured store name, so single-encoder instances are unchanged. Later dimensions use "<name>-<dim>". Identify searches only the store of the probe size and returns no match when no voice of that size exists. Forget finds the store from the stored embedding. A name may hold one voice per encoder. Assisted-by: Claude Code:claude-sonnet-5-5 golangci-lint
This commit is contained in:
4 files changed
+270
-3
No files matched your search
@@ -201,6 +201,8 @@ func newApplication(appConfig *config.ApplicationConfig) *Application {
|
||||
// Voice (speaker) recognition registry — same plumbing, separate
|
||||
// namespace so embedding spaces stay isolated (a face vector and a
|
||||
// speaker vector are not comparable and differ in dimensionality).
|
||||
// The registry also splits its store per embedding dimension, so speaker
|
||||
// encoders of different sizes can coexist.
|
||||
voiceStoreResolver := func(_ context.Context, storeName string) (pkggrpc.Backend, error) {
|
||||
return corebackend.StoreBackend(ml, appConfig, app.backendLoader, storeName, "")
|
||||
}
|
||||
|
||||
@@ -27,6 +27,12 @@ type StoreResolver func(ctx context.Context, storeName string) (grpc.Backend, er
|
||||
// pass 0 to accept whatever dimension arrives (useful when the voice
|
||||
// backend exposes recognizers of different sizes, e.g. ECAPA-TDNN at
|
||||
// 192 vs ResNet at 256).
|
||||
//
|
||||
// A vector store holds one dimension only, so with dim 0 the registry keeps
|
||||
// one store per embedding dimension. The first dimension seen uses storeName
|
||||
// itself (single-encoder setups are unchanged); every later dimension uses
|
||||
// "<storeName>-<dim>". Identify compares the probe with voices of its own
|
||||
// dimension only.
|
||||
func NewStoreRegistry(resolve StoreResolver, storeName string, dim int) Registry {
|
||||
return &storeRegistry{
|
||||
resolve: resolve,
|
||||
@@ -46,6 +52,33 @@ type storeRegistry struct {
|
||||
// every registration with its metadata. It is rebuilt on every Register
|
||||
// and lost on restart, which matches the lifetime of the in-memory store.
|
||||
idIndex sync.Map // map[string]Entry
|
||||
|
||||
// namespaces maps an embedding dimension to the store that holds it.
|
||||
nsMu sync.Mutex
|
||||
namespaces map[int]string
|
||||
}
|
||||
|
||||
// namespace returns the store name for a dimension. With create set, an
|
||||
// unknown dimension gets a new name; without it, ok is false for a dimension
|
||||
// that has no voices yet.
|
||||
func (r *storeRegistry) namespace(dim int, create bool) (name string, ok bool) {
|
||||
r.nsMu.Lock()
|
||||
defer r.nsMu.Unlock()
|
||||
if name, ok := r.namespaces[dim]; ok {
|
||||
return name, true
|
||||
}
|
||||
if !create {
|
||||
return "", false
|
||||
}
|
||||
if r.namespaces == nil {
|
||||
r.namespaces = map[int]string{}
|
||||
}
|
||||
name = r.storeName
|
||||
if len(r.namespaces) > 0 {
|
||||
name = fmt.Sprintf("%s-%d", r.storeName, dim)
|
||||
}
|
||||
r.namespaces[dim] = name
|
||||
return name, true
|
||||
}
|
||||
|
||||
func (r *storeRegistry) Register(ctx context.Context, embedding []float32, meta Metadata) (Metadata, error) {
|
||||
@@ -56,7 +89,8 @@ func (r *storeRegistry) Register(ctx context.Context, embedding []float32, meta
|
||||
return Metadata{}, fmt.Errorf("%w: expected %d, got %d", ErrDimensionMismatch, r.dim, len(embedding))
|
||||
}
|
||||
|
||||
backend, err := r.resolve(ctx, r.storeName)
|
||||
ns, _ := r.namespace(len(embedding), true)
|
||||
backend, err := r.resolve(ctx, ns)
|
||||
if err != nil {
|
||||
return Metadata{}, fmt.Errorf("voicerecognition: resolve store: %w", err)
|
||||
}
|
||||
@@ -91,7 +125,13 @@ func (r *storeRegistry) Identify(ctx context.Context, probe []float32, topK int)
|
||||
topK = 5
|
||||
}
|
||||
|
||||
backend, err := r.resolve(ctx, r.storeName)
|
||||
// No voice of this dimension was ever registered: nothing can match, and
|
||||
// querying another dimension's store would be an error.
|
||||
ns, ok := r.namespace(len(probe), false)
|
||||
if !ok {
|
||||
return []Match{}, nil
|
||||
}
|
||||
backend, err := r.resolve(ctx, ns)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("voicerecognition: resolve store: %w", err)
|
||||
}
|
||||
@@ -126,7 +166,8 @@ func (r *storeRegistry) Forget(ctx context.Context, id string) error {
|
||||
}
|
||||
embedding := raw.(Entry).Embedding
|
||||
|
||||
backend, err := r.resolve(ctx, r.storeName)
|
||||
ns, _ := r.namespace(len(embedding), true)
|
||||
backend, err := r.resolve(ctx, ns)
|
||||
if err != nil {
|
||||
return fmt.Errorf("voicerecognition: resolve store: %w", err)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,203 @@
|
||||
package voicerecognition_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"math"
|
||||
"sync"
|
||||
|
||||
"github.com/mudler/LocalAI/core/services/voicerecognition"
|
||||
"github.com/mudler/LocalAI/pkg/grpc"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
ggrpc "google.golang.org/grpc"
|
||||
)
|
||||
|
||||
// strictStore mimics local-store: one dimension per store, an error on a
|
||||
// mixed-dimension Set or Find, cosine similarity for Find.
|
||||
type strictStore struct {
|
||||
grpc.Backend
|
||||
mu sync.Mutex
|
||||
keys [][]float32
|
||||
vals [][]byte
|
||||
}
|
||||
|
||||
func (s *strictStore) dimErr(n int) error {
|
||||
if len(s.keys) > 0 && len(s.keys[0]) != n {
|
||||
return fmt.Errorf("Try to add key with length %d when existing length is %d", n, len(s.keys[0]))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *strictStore) StoresSet(_ context.Context, in *pb.StoresSetOptions, _ ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for i, k := range in.Keys {
|
||||
if err := s.dimErr(len(k.Floats)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s.keys = append(s.keys, k.Floats)
|
||||
s.vals = append(s.vals, in.Values[i].Bytes)
|
||||
}
|
||||
return &pb.Result{Success: true}, nil
|
||||
}
|
||||
|
||||
func (s *strictStore) StoresDelete(_ context.Context, in *pb.StoresDeleteOptions, _ ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for _, d := range in.Keys {
|
||||
for i, k := range s.keys {
|
||||
if fmt.Sprint(k) == fmt.Sprint(d.Floats) {
|
||||
s.keys = append(s.keys[:i], s.keys[i+1:]...)
|
||||
s.vals = append(s.vals[:i], s.vals[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return &pb.Result{Success: true}, nil
|
||||
}
|
||||
|
||||
func (s *strictStore) StoresFind(_ context.Context, in *pb.StoresFindOptions, _ ...ggrpc.CallOption) (*pb.StoresFindResult, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if err := s.dimErr(len(in.Key.Floats)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
res := &pb.StoresFindResult{}
|
||||
for i, k := range s.keys {
|
||||
var dot, na, nb float64
|
||||
for j := range k {
|
||||
dot += float64(k[j] * in.Key.Floats[j])
|
||||
na += float64(k[j] * k[j])
|
||||
nb += float64(in.Key.Floats[j] * in.Key.Floats[j])
|
||||
}
|
||||
res.Keys = append(res.Keys, &pb.StoresKey{Floats: k})
|
||||
res.Values = append(res.Values, &pb.StoresValue{Bytes: s.vals[i]})
|
||||
res.Similarities = append(res.Similarities, float32(dot/math.Sqrt(na*nb)))
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
func unit(dim, hot int) []float32 {
|
||||
v := make([]float32, dim)
|
||||
v[hot] = 1
|
||||
return v
|
||||
}
|
||||
|
||||
var _ = Describe("storeRegistry with several embedding dimensions", func() {
|
||||
var (
|
||||
mu sync.Mutex
|
||||
stores map[string]*strictStore
|
||||
reg voicerecognition.Registry
|
||||
ctx = context.Background()
|
||||
)
|
||||
BeforeEach(func() {
|
||||
stores = map[string]*strictStore{}
|
||||
reg = voicerecognition.NewStoreRegistry(func(_ context.Context, name string) (grpc.Backend, error) {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if stores[name] == nil {
|
||||
stores[name] = &strictStore{}
|
||||
}
|
||||
return stores[name], nil
|
||||
}, "voices", 0)
|
||||
})
|
||||
|
||||
It("keeps the configured store name for the first dimension and suffixes later ones", func() {
|
||||
_, err := reg.Register(ctx, unit(192, 0), voicerecognition.Metadata{Name: "ada"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
_, err = reg.Register(ctx, unit(256, 0), voicerecognition.Metadata{Name: "ada"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
_, err = reg.Register(ctx, unit(192, 1), voicerecognition.Metadata{Name: "ben"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
names := []string{}
|
||||
for n := range stores {
|
||||
names = append(names, n)
|
||||
}
|
||||
Expect(names).To(ConsistOf("voices", "voices-256"))
|
||||
Expect(stores["voices"].keys).To(HaveLen(2))
|
||||
Expect(stores["voices-256"].keys).To(HaveLen(1))
|
||||
})
|
||||
|
||||
It("identifies within the probe dimension only", func() {
|
||||
a192, _ := reg.Register(ctx, unit(192, 0), voicerecognition.Metadata{Name: "ada-192"})
|
||||
a256, _ := reg.Register(ctx, unit(256, 0), voicerecognition.Metadata{Name: "ada-256"})
|
||||
|
||||
m, err := reg.Identify(ctx, unit(192, 0), 5)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(m).To(HaveLen(1))
|
||||
Expect(m[0].ID).To(Equal(a192.ID))
|
||||
|
||||
m, err = reg.Identify(ctx, unit(256, 0), 5)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(m).To(HaveLen(1))
|
||||
Expect(m[0].ID).To(Equal(a256.ID))
|
||||
Expect(m[0].Distance).To(BeNumerically("~", 0, 1e-6))
|
||||
})
|
||||
|
||||
It("returns no match, not an error, for a dimension without voices", func() {
|
||||
_, err := reg.Register(ctx, unit(192, 0), voicerecognition.Metadata{Name: "ada"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
m, err := reg.Identify(ctx, unit(512, 0), 5)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(m).To(BeEmpty())
|
||||
Expect(stores).To(HaveLen(1)) // the probe created no store
|
||||
})
|
||||
|
||||
It("returns no match on an empty registry", func() {
|
||||
m, err := reg.Identify(ctx, unit(192, 0), 5)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(m).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("lets one name hold voices from several encoders and forgets each by ID", func() {
|
||||
a192, _ := reg.Register(ctx, unit(192, 0), voicerecognition.Metadata{Name: "ada", Model: "ecapa.gguf"})
|
||||
a256, _ := reg.Register(ctx, unit(256, 0), voicerecognition.Metadata{Name: "ada", Model: "wespeaker.gguf"})
|
||||
|
||||
got, err := reg.List(ctx)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(HaveLen(2))
|
||||
Expect(got[0].Embedding).To(HaveLen(192))
|
||||
Expect(got[1].Embedding).To(HaveLen(256))
|
||||
|
||||
Expect(reg.Forget(ctx, a256.ID)).To(Succeed())
|
||||
Expect(stores["voices-256"].keys).To(BeEmpty())
|
||||
Expect(stores["voices"].keys).To(HaveLen(1))
|
||||
m, err := reg.Identify(ctx, unit(256, 0), 5)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(m).To(BeEmpty())
|
||||
|
||||
Expect(reg.Forget(ctx, a192.ID)).To(Succeed())
|
||||
Expect(reg.Forget(ctx, a192.ID)).To(MatchError(voicerecognition.ErrNotFound))
|
||||
})
|
||||
|
||||
It("still enforces a fixed dimension when one is configured", func() {
|
||||
fixed := voicerecognition.NewStoreRegistry(func(context.Context, string) (grpc.Backend, error) { return &strictStore{}, nil }, "voices", 192)
|
||||
_, err := fixed.Register(ctx, unit(256, 0), voicerecognition.Metadata{Name: "ada"})
|
||||
Expect(err).To(MatchError(voicerecognition.ErrDimensionMismatch))
|
||||
_, err = fixed.Identify(ctx, unit(256, 0), 1)
|
||||
Expect(err).To(MatchError(voicerecognition.ErrDimensionMismatch))
|
||||
})
|
||||
|
||||
It("is safe for concurrent registration and identification across dimensions", func() {
|
||||
dims := []int{192, 256, 512}
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 30; i++ {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer GinkgoRecover()
|
||||
defer wg.Done()
|
||||
d := dims[i%len(dims)]
|
||||
_, err := reg.Register(ctx, unit(d, i%d), voicerecognition.Metadata{Name: fmt.Sprintf("v%d", i)})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
_, err = reg.Identify(ctx, unit(d, 0), 3)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
got, _ := reg.List(ctx)
|
||||
Expect(got).To(HaveLen(30))
|
||||
Expect(stores).To(HaveLen(3))
|
||||
})
|
||||
})
|
||||
@@ -185,6 +185,27 @@ recognition - the voice-recognition HTTP API is designed to swap the
|
||||
backing store without changing the wire format.
|
||||
{{% /notice %}}
|
||||
|
||||
### Voices from different encoders
|
||||
|
||||
Voices from encoders with different embedding sizes can be registered on the
|
||||
same instance, for example 192-value ECAPA-TDNN voices next to 256-value
|
||||
WeSpeaker ResNet34 voices. LocalAI keeps one in-memory vector store per
|
||||
embedding size, so registering a voice of a new size no longer fails.
|
||||
|
||||
- Identification compares the probe only with voices of the same size. Voices
|
||||
from an encoder of another size are never candidates.
|
||||
- If no voice of the probe size is registered, `/v1/voice/identify` returns an
|
||||
empty `matches` list and the realtime voice gate reports an unknown speaker.
|
||||
Neither returns an error.
|
||||
- Two encoders can still give the same size (ECAPA and CAM++ both give 192
|
||||
values). Those voices share a store and the encoder tag or identity checks
|
||||
described below apply.
|
||||
- A name is not unique. One name can hold one voice per encoder, and each
|
||||
registration has its own ID. `/v1/voice/forget` removes the voice with that
|
||||
ID only, so forget each ID to remove a person from every encoder.
|
||||
- Naming in diarization and live transcription reads the registry as a whole
|
||||
and uses the voices that match the loaded encoder.
|
||||
|
||||
## Naming speakers in diarization and live transcription
|
||||
|
||||
The parakeet-cpp backend can put the names of registered voices on
|
||||
|
||||
Reference in new issue
Block a user