diff --git a/backend/backend.proto b/backend/backend.proto index eda5ff6c8..676e6222c 100644 --- a/backend/backend.proto +++ b/backend/backend.proto @@ -830,6 +830,7 @@ message DiarizeRequest { // Registered voices the backend may use to name speakers. Only backends that // identify speakers themselves read this; others ignore it. repeated KnownVoice known_voices = 12; + bool include_speaker_profiles = 13; // opt-in sensitive embeddings; unsupported backends must reject } message DiarizeSegment { @@ -846,12 +847,14 @@ message DiarizeSegment { // names the encoder that produced it, so a backend with a different encoder can // refuse vectors that are not comparable. message KnownVoice { + string id = 4; // registration ID; native keys must not aggregate duplicate display names string name = 1; repeated float embedding = 2; string model = 3; } message DiarizeResponse { + string speaker_profiles_json = 5; // versioned speaker_profiles object only; absent by default repeated DiarizeSegment segments = 1; int32 num_speakers = 2; // count of distinct speaker labels in `segments` float duration = 3; // total audio duration in seconds (0 if unknown) @@ -907,6 +910,11 @@ message MemoryUsageData { map breakdown = 2; } +message SpeakerEncoder { + string identity = 1; // sha256 of loaded GGUF bytes + int32 dimension = 2; +} + message StatusResponse { enum State { UNINITIALIZED = 0; @@ -916,6 +924,7 @@ message StatusResponse { } State state = 1; MemoryUsageData memory = 2; + SpeakerEncoder speaker_encoder = 3; // trusted metadata from the loaded server encoder, never request data } message Message { diff --git a/backend/go/parakeet-cpp/Makefile b/backend/go/parakeet-cpp/Makefile index 77386dc92..fc7eac0ae 100644 --- a/backend/go/parakeet-cpp/Makefile +++ b/backend/go/parakeet-cpp/Makefile @@ -1,6 +1,6 @@ # parakeet-cpp backend Makefile. # -# Upstream pin lives below as PARAKEET_VERSION?=8c8cec0c4564610a0a4b30a8a6f2ead15d1a76fb +# Upstream pin lives below as PARAKEET_VERSION?=bee7c14dfcc23613df58176c59a40459e7b47095 # (.github/bump_deps.sh) can find and update it - matches the # whisper.cpp / ds4 / vibevoice-cpp convention. # @@ -15,7 +15,7 @@ # That's what the L0 smoke test uses. The default target below does the # proper clone-at-pin + cmake build so CI doesn't need a side-checkout. -PARAKEET_VERSION?=8c8cec0c4564610a0a4b30a8a6f2ead15d1a76fb +PARAKEET_VERSION?=bee7c14dfcc23613df58176c59a40459e7b47095 PARAKEET_REPO?=https://github.com/mudler/parakeet.cpp GOCMD?=go diff --git a/backend/go/parakeet-cpp/diarize.go b/backend/go/parakeet-cpp/diarize.go index c36cfbdcd..8cbfd33f6 100644 --- a/backend/go/parakeet-cpp/diarize.go +++ b/backend/go/parakeet-cpp/diarize.go @@ -109,6 +109,11 @@ func unsupportedDiarizeFields(req *pb.DiarizeRequest) []string { // turn); otherwise, or when no ASR companion is loaded, segments carry no // text (parakeet_capi_diarize_pcm) and no error is raised. func (p *ParakeetCpp) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error) { + if req.GetIncludeSpeakerProfiles() { + if CppDiarizeProfilesPCMJSON == nil || CppSpeakerIdentity == nil || CppSpeakerDim == nil || p.spkCtx == 0 { + return pb.DiarizeResponse{}, status.Error(codes.Unimplemented, "parakeet-cpp: speaker profiles require a loaded speaker encoder and profile-capable library") + } + } if p.diarCtx == 0 { return pb.DiarizeResponse{}, status.Error(codes.FailedPrecondition, "parakeet-cpp: model is not a diarization model") @@ -136,6 +141,11 @@ func (p *ParakeetCpp) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error } wantText := req.GetIncludeText() && p.ctxPtr != 0 && CppTranscribeAndDiarizeJSON != nil + if req.GetIncludeSpeakerProfiles() { + // Preserve the no-ASR fallback, but never silently omit text when an + // ASR companion is loaded and its timestamped PCM API is unavailable. + wantText = req.GetIncludeText() && p.ctxPtr != 0 + } var reg uintptr if len(req.GetKnownVoices()) > 0 && p.spkCtx != 0 { @@ -146,34 +156,51 @@ func (p *ParakeetCpp) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error defer p.freeSpeakerRegistry(reg) } - raw, err := p.diarizeCall(pcm, wantText, reg) + raw, err := p.diarizeCall(pcm, wantText, reg, req.GetIncludeSpeakerProfiles()) if err != nil { return pb.DiarizeResponse{}, err } + var profiles string + if req.GetIncludeSpeakerProfiles() { + var doc struct { + Profiles json.RawMessage `json:"speaker_profiles"` + } + if err := json.Unmarshal([]byte(raw), &doc); err != nil || len(doc.Profiles) == 0 || string(doc.Profiles) == "null" { + return pb.DiarizeResponse{}, status.Error(codes.Internal, "parakeet-cpp: missing speaker profiles") + } + profiles = string(doc.Profiles) + } segments, err := parseDiarizeDoc(raw, wantText) if err != nil { return pb.DiarizeResponse{}, err } + displayNames := voiceNames(req.GetKnownVoices()) + for _, segment := range segments { + if name, ok := displayNames[segment.Name]; ok { + segment.Name = name + } + } + segments = applyDurationFilters(segments, req.GetMinDurationOn(), req.GetMinDurationOff()) renumberDiarizeSegments(segments) return pb.DiarizeResponse{ - Segments: segments, - NumSpeakers: distinctDiarizeSpeakers(segments), - Duration: duration, + SpeakerProfilesJson: profiles, + Segments: segments, + NumSpeakers: distinctDiarizeSpeakers(segments), + Duration: duration, }, nil } -// diarizeCall runs the single C call Diarize needs (transcribe_and_diarize_json -// when wantText, else diarize_pcm) under engineMu, and returns the raw JSON +// diarizeCall runs diarization and optional ASR under engineMu, returning a JSON // document. p.diarCtx (and, on the include_text path, p.ctxPtr) is re-checked // under the lock before the C call: Diarize's own p.diarCtx==0/wantText checks // run before this lock is taken, so a Free() racing in between (which zeroes // those fields under the same engineMu) would otherwise reach the C side with // a freed context. last_error is ctx-shared, so it is read under the same // lock as the failing call. -func (p *ParakeetCpp) diarizeCall(pcm []float32, wantText bool, reg uintptr) (string, error) { +func (p *ParakeetCpp) diarizeCall(pcm []float32, wantText bool, reg uintptr, wantProfiles bool) (string, error) { p.engineMu.Lock() defer p.engineMu.Unlock() @@ -181,8 +208,16 @@ func (p *ParakeetCpp) diarizeCall(pcm []float32, wantText bool, reg uintptr) (st return "", grpcerrors.ModelNotLoaded("parakeet-cpp") } + if wantProfiles && (p.spkCtx == 0 || CppDiarizeProfilesPCMJSON == nil) { + return "", status.Error(codes.Unimplemented, "parakeet-cpp: speaker profile capability unavailable") + } + if wantProfiles && wantText && CppTranscribePcmBatchJSON == nil { + return "", status.Error(codes.Unimplemented, "parakeet-cpp: combined profile export requires timestamped PCM transcription") + } var cstr uintptr switch { + case wantProfiles: + cstr = CppDiarizeProfilesPCMJSON(p.diarCtx, p.spkCtx, reg, &pcm[0], int32(len(pcm)), 16000, p.speakerAccept, p.speakerMargin) case reg != 0 && wantText: if CppTranscribeAndDiarizeNamedJSON == nil { return "", status.Error(codes.Unimplemented, @@ -203,13 +238,63 @@ func (p *ParakeetCpp) diarizeCall(pcm []float32, wantText bool, reg uintptr) (st cstr = CppDiarizePCM(p.diarCtx, &pcm[0], int32(len(pcm)), 16000) } if cstr == 0 { - return "", fmt.Errorf("parakeet-cpp: diarize failed: %s", diarizeLastError(p, wantText, reg != 0)) + return "", fmt.Errorf("parakeet-cpp: diarize failed: %s", diarizeLastError(p, wantText, reg != 0 || wantProfiles)) } raw := goStringFromCPtr(cstr) CppFreeString(cstr) + if wantProfiles && wantText { + // Keep both contexts alive and last_error protected through both calls. + // Only the profile call diarizes: ASR contributes words, never slots. + cstr = CppTranscribePcmBatchJSON(p.ctxPtr, pcm, []int32{int32(len(pcm))}, 1, 16000, 0) + if cstr == 0 { + return "", fmt.Errorf("parakeet-cpp: transcribe failed: %s", CppLastError(p.ctxPtr)) + } + asr := goStringFromCPtr(cstr) + CppFreeString(cstr) + return composeProfileTranscript(raw, asr) + } return raw, nil } +// composeProfileTranscript preserves the native profiles and slot-keyed names, +// adding utterances from timestamped words on that same diarization timeline. +func composeProfileTranscript(raw, asr string) (string, error) { + var doc diarizePCMDoc + if err := json.Unmarshal([]byte(raw), &doc); err != nil { + return "", fmt.Errorf("parakeet-cpp: decode profile diarization: %w", err) + } + var transcripts []transcriptJSON + if err := json.Unmarshal([]byte(asr), &transcripts); err != nil { + return "", fmt.Errorf("parakeet-cpp: decode transcript: %w", err) + } + if len(transcripts) != 1 { + return "", fmt.Errorf("parakeet-cpp: expected one transcript, got %d", len(transcripts)) + } + t := transcripts[0] + if len(t.Words) == 0 && strings.TrimSpace(t.Text) != "" { + return "", fmt.Errorf("parakeet-cpp: transcript has no timestamped words") + } + groups, slots := splitAtSpeakerChanges([][]transcriptWord{t.Words}, assignSpeakers(t.Words, doc.Segments)) + utterances := make([]diarizeUtteranceJSON, 0, len(groups)) + for i, group := range groups { + parts := make([]string, len(group)) + for j, word := range group { + parts[j] = word.W + } + utterances = append(utterances, diarizeUtteranceJSON{ + Speaker: slots[i], Start: group[0].Start, End: group[len(group)-1].End, + Text: strings.TrimSpace(strings.Join(parts, " ")), + }) + } + var fields map[string]json.RawMessage + if err := json.Unmarshal([]byte(raw), &fields); err != nil { + return "", err + } + fields["utterances"], _ = json.Marshal(utterances) + combined, err := json.Marshal(fields) + return string(combined), err +} + // diarizeLastError reads last_error off p.diarCtx and, on the include_text // path, p.ctxPtr too — the failing call is CppTranscribeAndDiarizeJSON there, // and either side of the pairing may be the one that set it — then joins diff --git a/backend/go/parakeet-cpp/diarize_test.go b/backend/go/parakeet-cpp/diarize_test.go index 72d63766d..c97d7baf6 100644 --- a/backend/go/parakeet-cpp/diarize_test.go +++ b/backend/go/parakeet-cpp/diarize_test.go @@ -259,7 +259,7 @@ var _ = Describe("ParakeetCpp.Diarize", func() { // Simulate a Free() racing between Diarize's own diarCtx==0 check and // diarizeCall's lock, exactly as it zeroes diarCtx under engineMu. p.diarCtx = 0 - _, err := p.diarizeCall(make([]float32, 10), false, 0) + _, err := p.diarizeCall(make([]float32, 10), false, 0, false) Expect(grpcerrors.IsModelNotLoaded(err)).To(BeTrue()) Expect(called).To(BeFalse(), "no C call once diarCtx was cleared") }) @@ -311,6 +311,24 @@ var _ = Describe("ParakeetCpp.Diarize", func() { } }) + It("replays distinct IDs and translates duplicate display names offline", func() { + vectors := map[string]float32{} + CppSpeakerRegistryAddEmbedding = func(_ uintptr, key string, emb *float32, _ int32) int32 { vectors[key] = *emb; return 0 } + CppDiarizeNamedPCMJSON = func(_, _, _ uintptr, _ *float32, _, _ int32, _, _ float32) uintptr { + return pool.cstr(`{"segments":[{"speaker":0,"start":0,"end":2},{"speaker":1,"start":2,"end":4}],"names":{"0":{"name":"a","score":0.9},"1":{"name":"b","score":0.8}}}`) + } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + res, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), KnownVoices: []*pb.KnownVoice{ + {Id: "a", Name: "Ada", Embedding: []float32{1, 0}}, {Id: "b", Name: "Ada", Embedding: []float32{0, 1}}, + }}) + Expect(err).NotTo(HaveOccurred()) + Expect(vectors).To(Equal(map[string]float32{"a": 1, "b": 0})) + Expect(res.Segments[0].Name).To(Equal("Ada")) + Expect(res.Segments[1].Name).To(Equal("Ada")) + Expect(res.Segments[0].Speaker).NotTo(Equal(res.Segments[1].Speaker)) + Expect(res.SpeakerProfilesJson).To(BeEmpty()) + }) + It("puts the registered names on the segments and frees the registry", func() { p := &ParakeetCpp{diarCtx: 1, spkCtx: 2, speakerAccept: 0.7, speakerMargin: 0.05} res, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), KnownVoices: ada}) diff --git a/backend/go/parakeet-cpp/goparakeetcpp.go b/backend/go/parakeet-cpp/goparakeetcpp.go index 59b74979c..04c26bff6 100644 --- a/backend/go/parakeet-cpp/goparakeetcpp.go +++ b/backend/go/parakeet-cpp/goparakeetcpp.go @@ -109,6 +109,8 @@ var ( // CppTranscribeAndDiarizeNamedJSON are ABI v9; CppSpeakerRegistryAddEmbedding and // CppDiarizeNamedPCMJSON are ABI v10. All are nil on an older libparakeet.so, and // Load refuses speaker_model: unless the v10 ones are present. + CppSpeakerIdentity func(ctx uintptr) uintptr + CppDiarizeProfilesPCMJSON func(diar, speaker, reg uintptr, samples *float32, n, sampleRate int32, acceptThreshold, margin float32) uintptr CppSpeakerDim func(ctx uintptr) int32 CppSpeakerRegistryNew func() uintptr CppSpeakerRegistryFree func(reg uintptr) diff --git a/backend/go/parakeet-cpp/main.go b/backend/go/parakeet-cpp/main.go index 20dd1358a..8c7a3b7c3 100644 --- a/backend/go/parakeet-cpp/main.go +++ b/backend/go/parakeet-cpp/main.go @@ -132,6 +132,15 @@ func main() { purego.RegisterLibFunc(&CppDiarizeNamedPCMJSON, lib, "parakeet_capi_diarize_named_pcm_json") } + for _, lf := range []LibFuncs{ + {&CppSpeakerIdentity, "parakeet_capi_speaker_identity"}, + {&CppSpeakerDim, "parakeet_capi_speaker_dim"}, + {&CppDiarizeProfilesPCMJSON, "parakeet_capi_diarize_profiles_pcm_json"}, + } { + if sym, err := purego.Dlsym(lib, lf.Name); err == nil && sym != 0 { + purego.RegisterLibFunc(lf.FuncPtr, lib, lf.Name) + } + } fmt.Fprintf(os.Stderr, "[parakeet-cpp] ABI=%d\n", CppAbiVersion()) flag.Parse() diff --git a/backend/go/parakeet-cpp/profiles.go b/backend/go/parakeet-cpp/profiles.go new file mode 100644 index 000000000..b902b84e3 --- /dev/null +++ b/backend/go/parakeet-cpp/profiles.go @@ -0,0 +1,27 @@ +// SPDX-License-Identifier: MIT +package main + +import pb "github.com/mudler/LocalAI/pkg/grpc/proto" + +// Status exposes only metadata derived from the loaded encoder. The native +// identity is borrowed, so copy it while holding the context lifetime lock. +func (p *ParakeetCpp) Status() (pb.StatusResponse, error) { + result, err := p.Base.Status() + if err != nil { + return pb.StatusResponse{}, err + } + p.engineMu.Lock() + defer p.engineMu.Unlock() + if p.spkCtx != 0 && CppSpeakerIdentity != nil && CppSpeakerDim != nil { + identity := CppSpeakerIdentity(p.spkCtx) + dim := CppSpeakerDim(p.spkCtx) + if identity != 0 && dim > 0 { + result.SpeakerEncoder = &pb.SpeakerEncoder{Identity: goStringFromCPtr(identity), Dimension: dim} + } + } + return pb.StatusResponse{ + State: result.GetState(), + Memory: result.GetMemory(), + SpeakerEncoder: result.GetSpeakerEncoder(), + }, nil +} diff --git a/backend/go/parakeet-cpp/profiles_test.go b/backend/go/parakeet-cpp/profiles_test.go new file mode 100644 index 000000000..a44310361 --- /dev/null +++ b/backend/go/parakeet-cpp/profiles_test.go @@ -0,0 +1,188 @@ +// SPDX-License-Identifier: MIT +package main + +import ( + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +var _ = Describe("profile capability", func() { + It("rejects opt-in without support instead of silently returning plain output", func() { + restore := diarizeStubs() + defer restore() + CppDiarizePCM = func(uintptr, *float32, int32, int32) uintptr { return 0 } + p := &ParakeetCpp{diarCtx: 1} + _, err := p.Diarize(&pb.DiarizeRequest{IncludeSpeakerProfiles: true}) + Expect(status.Code(err)).To(Equal(codes.Unimplemented)) + }) +}) +var _ = Describe("profile export transport", func() { + It("exports with no registry and copies trusted metadata without freeing borrowed identity", func() { + restore := diarizeStubs() + defer restore() + oldExport, oldIdentity := CppDiarizeProfilesPCMJSON, CppSpeakerIdentity + defer func() { CppDiarizeProfilesPCMJSON, CppSpeakerIdentity = oldExport, oldIdentity }() + pool := &diarizeCstrPool{} + identity := pool.cstr("sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa") + CppSpeakerIdentity = func(ctx uintptr) uintptr { Expect(ctx).To(Equal(uintptr(2))); return identity } + CppSpeakerDim = func(uintptr) int32 { return 2 } + freed := 0 + CppFreeString = func(ptr uintptr) { Expect(ptr).NotTo(Equal(identity)); freed++ } + CppDiarizePCM = func(uintptr, *float32, int32, int32) uintptr { Fail("legacy function called"); return 0 } + CppDiarizeProfilesPCMJSON = func(d, s, r uintptr, pcm *float32, n, hz int32, a, m float32) uintptr { + Expect(r).To(BeZero()) + Expect(s).To(Equal(uintptr(2))) + return pool.cstr(`{"segments":[{"speaker":0,"start":0,"end":3}],"speaker_profiles":{"version":1,"encoder":{"identity":"sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa","dimension":2},"speakers":[]}}`) + } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + out, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(3), IncludeSpeakerProfiles: true}) + Expect(err).NotTo(HaveOccurred()) + Expect(out.SpeakerProfilesJson).To(ContainSubstring(`"version":1`)) + Expect(freed).To(Equal(1)) + meta, err := p.Status() + Expect(err).NotTo(HaveOccurred()) + Expect(meta.SpeakerEncoder.Dimension).To(Equal(int32(2))) + Expect(meta.SpeakerEncoder.Identity).To(HavePrefix("sha256:")) + Expect(freed).To(Equal(1)) + p.spkCtx = 0 + meta, err = p.Status() + Expect(err).NotTo(HaveOccurred()) + Expect(meta.SpeakerEncoder).To(BeNil()) + }) + It("maps both offline and realtime matches without changing speaker slots", func() { + voices := []*pb.KnownVoice{{Id: "one", Name: "Ada"}, {Id: "two", Name: "Ada"}} + names := map[string]speakerNameJSON{"0": {Name: "one", Score: .9}, "1": {Name: "two", Score: .8}} + translateNames(names, voiceNames(voices)) + live := liveSpeakersToProto([]sceneSpeakerJSON{{Speaker: 0}, {Speaker: 1}}, names) + Expect(live[0].Name).To(Equal("Ada")) + Expect(live[1].Name).To(Equal("Ada")) + Expect(names["0"].Score).To(Equal(float32(.9))) + Expect(names["1"].Score).To(Equal(float32(.8))) + }) +}) + +var _ = Describe("combined transcript profile export", func() { + var p *ParakeetCpp + var pool *diarizeCstrPool + var restore func() + var profileCalls, asrCalls int + var profileRaw string + BeforeEach(func() { + restore = diarizeStubs() + oldExport, oldIdentity, oldBatch := CppDiarizeProfilesPCMJSON, CppSpeakerIdentity, CppTranscribePcmBatchJSON + DeferCleanup(func() { + restore() + CppDiarizeProfilesPCMJSON, CppSpeakerIdentity, CppTranscribePcmBatchJSON = oldExport, oldIdentity, oldBatch + }) + pool = &diarizeCstrPool{} + p = &ParakeetCpp{diarCtx: 1, spkCtx: 2, ctxPtr: 3} + profileCalls, asrCalls = 0, 0 + profileRaw = `{"segments":[{"speaker":7,"start":0,"end":1},{"speaker":2,"start":1,"end":3}],"names":{"7":{"name":"one","score":0.9},"2":{"name":"two","score":0.8}},"speaker_profiles":{"version":1,"encoder":{"identity":"sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa","dimension":2},"speakers":[{"speaker":2,"embedding":[0,1]},{"speaker":7,"embedding":[1,0]}]}}` + CppSpeakerIdentity = func(uintptr) uintptr { return pool.cstr("identity") } + CppSpeakerDim = func(uintptr) int32 { return 2 } + CppFreeString = func(uintptr) {} + CppLastError = func(uintptr) string { return "inference failed" } + CppDiarizePCM = func(uintptr, *float32, int32, int32) uintptr { Fail("second diarization"); return 0 } + CppTranscribeAndDiarizeJSON = func(uintptr, uintptr, *float32, int32, int32) uintptr { Fail("second diarization"); return 0 } + CppDiarizeProfilesPCMJSON = func(d, s, r uintptr, pcm *float32, n, hz int32, a, m float32) uintptr { + profileCalls++ + return pool.cstr(profileRaw) + } + CppTranscribePcmBatchJSON = func(ctx uintptr, pcm []float32, sizes []int32, clips, hz, decoder int32) uintptr { + asrCalls++ + Expect(ctx).To(Equal(uintptr(3))) + Expect(clips).To(Equal(int32(1))) + Expect(p.engineMu.TryLock()).To(BeFalse()) + return pool.cstr(`[{"text":"Hello there","words":[{"w":"Hello","start":0.1,"end":0.9},{"w":"there","start":1.2,"end":2}]}]`) + } + }) + request := func() *pb.DiarizeRequest { + return &pb.DiarizeRequest{Dst: diarizeWav(3), IncludeText: true, IncludeSpeakerProfiles: true} + } + It("exports text and profiles with an empty registry using raw sparse slots", func() { + out, err := p.Diarize(request()) + Expect(err).NotTo(HaveOccurred()) + Expect(profileCalls).To(Equal(1)) + Expect(asrCalls).To(Equal(1)) + Expect(out.Segments).To(HaveLen(2)) + Expect(out.Segments[0].Speaker).To(Equal("7")) + Expect(out.Segments[0].Text).To(Equal("Hello")) + Expect(out.Segments[1].Speaker).To(Equal("2")) + Expect(out.Segments[1].Text).To(Equal("there")) + Expect(out.SpeakerProfilesJson).To(ContainSubstring(`"speaker":2,"embedding":[0,1]`)) + Expect(out.SpeakerProfilesJson).To(ContainSubstring(`"speaker":7,"embedding":[1,0]`)) + }) + It("keeps duplicate display names independent through replay and attribution", func() { + CppSpeakerRegistryNew = func() uintptr { return 4 } + CppSpeakerRegistryFree = func(uintptr) {} + ids := []string{} + CppSpeakerRegistryAddEmbedding = func(r uintptr, id string, v *float32, n int32) int32 { + ids = append(ids, id) + if id == "one" { + Expect(*v).To(Equal(float32(1))) + } else { + Expect(*v).To(BeZero()) + } + return 0 + } + req := request() + req.KnownVoices = []*pb.KnownVoice{{Id: "one", Name: "Ada", Embedding: []float32{1, 0}}, {Id: "two", Name: "Ada", Embedding: []float32{0, 1}}} + out, err := p.Diarize(req) + Expect(err).NotTo(HaveOccurred()) + Expect(ids).To(Equal([]string{"one", "two"})) + Expect(out.Segments[0].Name).To(Equal("Ada")) + Expect(out.Segments[1].Name).To(Equal("Ada")) + Expect(out.Segments[0].NameScore).To(Equal(float32(.9))) + Expect(out.Segments[1].NameScore).To(Equal(float32(.8))) + Expect(out.Segments[0].Speaker).To(Equal("7")) + Expect(out.Segments[1].Speaker).To(Equal("2")) + }) + It("propagates profile failure", func() { + CppDiarizeProfilesPCMJSON = func(uintptr, uintptr, uintptr, *float32, int32, int32, float32, float32) uintptr { return 0 } + _, err := p.Diarize(request()) + Expect(err).To(MatchError(ContainSubstring("inference failed"))) + }) + It("propagates ASR failure", func() { + CppTranscribePcmBatchJSON = func(uintptr, []float32, []int32, int32, int32, int32) uintptr { return 0 } + _, err := p.Diarize(request()) + Expect(err).To(MatchError(ContainSubstring("inference failed"))) + }) + It("rejects missing ASR capability when an ASR companion is loaded", func() { + CppTranscribePcmBatchJSON = nil + _, err := p.Diarize(request()) + Expect(status.Code(err)).To(Equal(codes.Unimplemented)) + }) + It("rejects transcript text without timestamped words", func() { + CppTranscribePcmBatchJSON = func(uintptr, []float32, []int32, int32, int32, int32) uintptr { + return pool.cstr(`[{"text":"missing words"}]`) + } + _, err := p.Diarize(request()) + Expect(err).To(MatchError(ContainSubstring("timestamped words"))) + }) + It("preserves no-ASR fallback with profiles and no text", func() { + p.ctxPtr = 0 + out, err := p.Diarize(request()) + Expect(err).NotTo(HaveOccurred()) + Expect(asrCalls).To(BeZero()) + Expect(profileCalls).To(Equal(1)) + Expect(out.Segments[0].Text).To(BeEmpty()) + Expect(out.SpeakerProfilesJson).NotTo(BeEmpty()) + }) + It("keeps the legacy combined call when profile export is off", func() { + CppTranscribeAndDiarizeJSON = func(uintptr, uintptr, *float32, int32, int32) uintptr { + return pool.cstr(`{"utterances":[{"speaker":5,"text":"legacy","start":0,"end":3}]}`) + } + req := request() + req.IncludeSpeakerProfiles = false + out, err := p.Diarize(req) + Expect(err).NotTo(HaveOccurred()) + Expect(profileCalls).To(BeZero()) + Expect(asrCalls).To(BeZero()) + Expect(out.SpeakerProfilesJson).To(BeEmpty()) + Expect(out.Segments[0].Text).To(Equal("legacy")) + Expect(out.Segments[0].Speaker).To(Equal("5")) + }) +}) diff --git a/backend/go/parakeet-cpp/scene.go b/backend/go/parakeet-cpp/scene.go index e104e724a..099b968e2 100644 --- a/backend/go/parakeet-cpp/scene.go +++ b/backend/go/parakeet-cpp/scene.go @@ -72,11 +72,12 @@ func (p *ParakeetCpp) sceneWanted() bool { // speaker-named stream (0 for a plain one). The stream borrows both: sceneFree // frees the stream first and then the registry, which this handle owns. type sceneStreamHandle struct { - s uintptr - diar uintptr - tag uintptr - spk uintptr - reg uintptr + names map[string]string + s uintptr + diar uintptr + tag uintptr + spk uintptr + reg uintptr } // sceneBegin opens a no-ASR scene stream (diarization and/or sound events @@ -122,7 +123,7 @@ func (p *ParakeetCpp) sceneBegin(voices []*pb.KnownVoice) sceneStreamHandle { p.freeSpeakerRegistry(reg) return sceneStreamHandle{} } - return sceneStreamHandle{s: s, diar: diar, tag: tag, spk: p.spkCtx, reg: reg} + return sceneStreamHandle{s: s, diar: diar, tag: tag, spk: p.spkCtx, reg: reg, names: voiceNames(voices)} } s := CppSceneStreamBegin(0, diar, tag, &opts) if s == 0 { @@ -193,6 +194,7 @@ func (p *ParakeetCpp) sceneFeed(h sceneStreamHandle, pcm []float32, isLast bool) if err := json.Unmarshal([]byte(raw), &doc); err != nil { return sceneFeedJSON{}, fmt.Errorf("parakeet-cpp: decode scene json: %w", err) } + translateNames(doc.Names, h.names) return doc, nil } diff --git a/backend/go/parakeet-cpp/scene_test.go b/backend/go/parakeet-cpp/scene_test.go index 76a87a9be..ca4005707 100644 --- a/backend/go/parakeet-cpp/scene_test.go +++ b/backend/go/parakeet-cpp/scene_test.go @@ -68,6 +68,24 @@ var _ = Describe("scene stream with speaker names", func() { }) AfterEach(func() { restore() }) + It("replays duplicate display names independently and translates realtime matches", func() { + vectors := map[string]float32{} + CppSpeakerRegistryAddEmbedding = func(_ uintptr, key string, emb *float32, _ int32) int32 { vectors[key] = *emb; return 0 } + CppSceneStreamBeginSpeaker = func(_, _, _, _, _ uintptr, _ *cSceneOpts) uintptr { return 55 } + CppSceneStreamFeedJSON = func(uintptr, *float32, int32, int32) uintptr { + return pool.cstr(`{"speakers":[{"speaker":0,"start":0,"end":2},{"speaker":1,"start":2,"end":4}],"names":{"0":{"name":"a","score":0.9},"1":{"name":"b","score":0.8}}}`) + } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + h := p.sceneBegin([]*pb.KnownVoice{{Id: "a", Name: "Ada", Embedding: []float32{1, 0}}, {Id: "b", Name: "Ada", Embedding: []float32{0, 1}}}) + doc, err := p.sceneFeed(h, nil, true) + Expect(err).NotTo(HaveOccurred()) + Expect(vectors).To(Equal(map[string]float32{"a": 1, "b": 0})) + out := liveSpeakersToProto(doc.Speakers, doc.Names) + Expect(out[0].Name).To(Equal("Ada")) + Expect(out[1].Name).To(Equal("Ada")) + Expect(out[0].Speaker).NotTo(Equal(out[1].Speaker)) + }) + It("begins a speaker scene stream with the threshold and margin when voices are given", func() { var gotOpts cSceneOpts var gotReg, gotSpk uintptr diff --git a/backend/go/parakeet-cpp/speaker_registry.go b/backend/go/parakeet-cpp/speaker_registry.go index 18b0df99e..2d56ebddf 100644 --- a/backend/go/parakeet-cpp/speaker_registry.go +++ b/backend/go/parakeet-cpp/speaker_registry.go @@ -87,7 +87,7 @@ func (p *ParakeetCpp) buildSpeakerRegistryLocked(voices []*pb.KnownVoice) (uintp skipped++ continue } - if rc := CppSpeakerRegistryAddEmbedding(reg, v.GetName(), &emb[0], int32(len(emb))); rc != 0 { + if rc := CppSpeakerRegistryAddEmbedding(reg, voiceKey(v), &emb[0], int32(len(emb))); rc != 0 { xlog.Warn("parakeet-cpp: skipped a registered voice the speaker registry refused", "error", CppSpeakerRegistryLastError(reg)) skipped++ continue @@ -118,3 +118,26 @@ func (p *ParakeetCpp) freeSpeakerRegistry(reg uintptr) { } CppSpeakerRegistryFree(reg) } + +// Old transport clients have no IDs; retain their name-keyed semantics. +func voiceKey(v *pb.KnownVoice) string { + if v.GetId() != "" { + return v.GetId() + } + return v.GetName() +} +func voiceNames(voices []*pb.KnownVoice) map[string]string { + names := make(map[string]string, len(voices)) + for _, v := range voices { + names[voiceKey(v)] = v.GetName() + } + return names +} +func translateNames(names map[string]speakerNameJSON, display map[string]string) { + for slot, match := range names { + if name, ok := display[match.Name]; ok { + match.Name = name + names[slot] = match + } + } +} diff --git a/backend/go/parakeet-cpp/speaker_registry_test.go b/backend/go/parakeet-cpp/speaker_registry_test.go index aa046ca1c..4297b2978 100644 --- a/backend/go/parakeet-cpp/speaker_registry_test.go +++ b/backend/go/parakeet-cpp/speaker_registry_test.go @@ -69,6 +69,16 @@ var _ = Describe("buildSpeakerRegistry", func() { return &pb.KnownVoice{Name: name, Embedding: make([]float32, n)} } + It("keeps duplicate display names under distinct registration IDs", func() { + p := &ParakeetCpp{spkCtx: 5} + _, err := p.buildSpeakerRegistry([]*pb.KnownVoice{ + {Id: "id-a", Name: "Ada", Embedding: []float32{1, 0, 0}}, + {Id: "id-b", Name: "Ada", Embedding: []float32{0, 1, 0}}, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(added).To(Equal([]string{"id-a", "id-b"})) + }) + It("adds every known voice, in order", func() { p := &ParakeetCpp{spkCtx: 5} reg, err := p.buildSpeakerRegistry([]*pb.KnownVoice{ diff --git a/backend/go/parakeet-cpp/speakers.go b/backend/go/parakeet-cpp/speakers.go index 31e9111ce..5ab9a80b5 100644 --- a/backend/go/parakeet-cpp/speakers.go +++ b/backend/go/parakeet-cpp/speakers.go @@ -42,7 +42,7 @@ func (p *ParakeetCpp) diarizeSegmentsPCM(pcm []float32) ([]diarizeSegmentJSON, e if len(pcm) == 0 { return nil, nil } - raw, err := p.diarizeCall(pcm, false, 0) + raw, err := p.diarizeCall(pcm, false, 0, false) if err != nil { return nil, err } diff --git a/core/backend/diarization.go b/core/backend/diarization.go index 77c6ad22d..a87356cd4 100644 --- a/core/backend/diarization.go +++ b/core/backend/diarization.go @@ -2,8 +2,13 @@ package backend import ( "context" + "encoding/json" "fmt" "sort" + "strings" + + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/schema" @@ -20,15 +25,16 @@ import ( // don't act on. IncludeText only matters for backends that emit // per-segment transcripts as a by-product (e.g. vibevoice.cpp). type DiarizationRequest struct { - Audio string - Language string - NumSpeakers int32 - MinSpeakers int32 - MaxSpeakers int32 - ClusteringThreshold float32 - MinDurationOn float32 - MinDurationOff float32 - IncludeText bool + Audio string + Language string + NumSpeakers int32 + MinSpeakers int32 + MaxSpeakers int32 + ClusteringThreshold float32 + MinDurationOn float32 + MinDurationOff float32 + IncludeText bool + IncludeSpeakerProfiles bool // KnownVoices are registered voices a speaker-identifying backend may use // to name the speakers. Empty for every other backend and model. KnownVoices []voicerecognition.KnownVoice @@ -38,21 +44,22 @@ type DiarizationRequest struct { func (r *DiarizationRequest) toProto(threads uint32, modelIdentity string) *proto.DiarizeRequest { known := make([]*proto.KnownVoice, 0, len(r.KnownVoices)) for _, v := range r.KnownVoices { - known = append(known, &proto.KnownVoice{Name: v.Name, Embedding: v.Embedding, Model: v.Model}) + known = append(known, &proto.KnownVoice{Id: v.ID, Name: v.Name, Embedding: v.Embedding, Model: v.Model}) } return &proto.DiarizeRequest{ - ModelIdentity: modelIdentity, - Dst: r.Audio, - Threads: threads, - Language: r.Language, - NumSpeakers: r.NumSpeakers, - MinSpeakers: r.MinSpeakers, - MaxSpeakers: r.MaxSpeakers, - ClusteringThreshold: r.ClusteringThreshold, - MinDurationOn: r.MinDurationOn, - MinDurationOff: r.MinDurationOff, - IncludeText: r.IncludeText, - KnownVoices: known, + ModelIdentity: modelIdentity, + Dst: r.Audio, + Threads: threads, + Language: r.Language, + NumSpeakers: r.NumSpeakers, + MinSpeakers: r.MinSpeakers, + MaxSpeakers: r.MaxSpeakers, + ClusteringThreshold: r.ClusteringThreshold, + MinDurationOn: r.MinDurationOn, + MinDurationOff: r.MinDurationOff, + IncludeText: r.IncludeText, + IncludeSpeakerProfiles: r.IncludeSpeakerProfiles, + KnownVoices: known, } } @@ -85,16 +92,29 @@ func ModelDiarization(ctx context.Context, req DiarizationRequest, ml *model.Mod threads = uint32(*modelConfig.Threads) } + req.KnownVoices = compatiblePortableVoices(ctx, m, req.KnownVoices) r, err := m.Diarize(ctx, req.toProto(threads, modelConfig.Model)) if err != nil { return nil, err } - return diarizationResultFromProto(r), nil + out := diarizationResultFromProto(r) + if req.IncludeSpeakerProfiles { + trusted, err := speakerEncoderFromBackend(ctx, m) + if err != nil { + return nil, err + } + profiles, err := decodeSpeakerProfiles(r.GetSpeakerProfilesJson(), trusted) + if err != nil { + return nil, err + } + out.SpeakerProfiles = profiles + } + return out, nil } // diarizationResultFromProto normalizes backend speaker labels to // "SPEAKER_NN" — the convention pyannote/RTTM tooling expects — while -// keeping the original label available via the Speaker field. Each +// keeping the original label available via the Label field. Each // distinct backend label gets its own normalized id, in first-seen order. func diarizationResultFromProto(r *proto.DiarizeResponse) *schema.DiarizationResult { if r == nil { @@ -174,3 +194,59 @@ func diarizationResultFromProto(r *proto.DiarizeResponse) *schema.DiarizationRes return out } + +// ModelSpeakerEncoder obtains trusted metadata from the configured loaded model. +// HTTP enrollment must use this, never metadata supplied by the caller. +func ModelSpeakerEncoder(ctx context.Context, ml *model.ModelLoader, modelConfig config.ModelConfig, appConfig *config.ApplicationConfig) (schema.SpeakerEncoder, error) { + m, err := loadDiarizationModel(ml, modelConfig, appConfig) + if err != nil { + return schema.SpeakerEncoder{}, err + } + return speakerEncoderFromBackend(ctx, m) +} +func speakerEncoderFromBackend(ctx context.Context, m grpcPkg.Backend) (schema.SpeakerEncoder, error) { + r, err := m.Status(ctx) + if err != nil { + return schema.SpeakerEncoder{}, err + } + e := r.GetSpeakerEncoder() + trusted := schema.SpeakerEncoder{Identity: e.GetIdentity(), Dimension: int(e.GetDimension())} + if err := (schema.SpeakerProfiles{Version: 1, Encoder: trusted}).Validate(trusted); err != nil { + return schema.SpeakerEncoder{}, status.Error(codes.Unimplemented, "backend does not expose trusted speaker encoder metadata") + } + return trusted, nil +} +func decodeSpeakerProfiles(raw string, trusted schema.SpeakerEncoder) (*schema.SpeakerProfiles, error) { + if raw == "" { + return nil, status.Error(codes.Unimplemented, "backend does not support speaker profiles") + } + var profiles schema.SpeakerProfiles + if err := json.Unmarshal([]byte(raw), &profiles); err != nil { + return nil, fmt.Errorf("decode speaker profiles: %w", err) + } + if err := profiles.Validate(trusted); err != nil { + return nil, err + } + return &profiles, nil +} + +// Portable registrations require exact loaded identity and dimension. Legacy +// candidates use the trusted dimension when available; older backends without +// metadata retain their native dimension check. No registry entry sets it. +func compatiblePortableVoices(ctx context.Context, m grpcPkg.Backend, voices []voicerecognition.KnownVoice) []voicerecognition.KnownVoice { + if len(voices) == 0 { + return voices + } + trusted, err := speakerEncoderFromBackend(ctx, m) + out := make([]voicerecognition.KnownVoice, 0, len(voices)) + for _, v := range voices { + if err == nil && len(v.Embedding) != trusted.Dimension { + continue + } + if strings.HasPrefix(v.Model, "sha256:") && (err != nil || v.Model != trusted.Identity || len(v.Embedding) != trusted.Dimension) { + continue + } + out = append(out, v) + } + return out +} diff --git a/core/backend/diarization_profiles_test.go b/core/backend/diarization_profiles_test.go new file mode 100644 index 000000000..d559385a3 --- /dev/null +++ b/core/backend/diarization_profiles_test.go @@ -0,0 +1,129 @@ +// SPDX-License-Identifier: MIT +package backend + +import ( + "context" + "encoding/json" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/voicerecognition" + grpcPkg "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "strings" +) + +var _ = Describe("speaker profile transport", func() { + It("preserves opt-in and registration IDs in offline and live transport", func() { + v := voicerecognition.KnownVoice{ID: "id", Name: "Ada", Embedding: []float32{1, 0}} + r := (&DiarizationRequest{IncludeSpeakerProfiles: true, KnownVoices: []voicerecognition.KnownVoice{v}}).toProto(2, "model") + Expect(r.IncludeSpeakerProfiles).To(BeTrue()) + Expect(r.KnownVoices[0].Id).To(Equal("id")) + var o liveOptions + WithKnownVoices([]voicerecognition.KnownVoice{v})(&o) + Expect(liveConfigProto("", o).KnownVoices[0].Id).To(Equal("id")) + Expect((&DiarizationRequest{}).toProto(2, "").IncludeSpeakerProfiles).To(BeFalse()) + }) + It("validates response data against separate trusted metadata and omits defaults", func() { + trusted := schema.SpeakerEncoder{Identity: "sha256:" + strings.Repeat("a", 64), Dimension: 2} + raw, _ := json.Marshal(schema.SpeakerProfiles{Version: 1, Encoder: trusted}) + p, err := decodeSpeakerProfiles(string(raw), trusted) + Expect(err).NotTo(HaveOccurred()) + Expect(p.Encoder).To(Equal(trusted)) + trusted.Dimension = 3 + _, err = decodeSpeakerProfiles(string(raw), trusted) + Expect(err).To(HaveOccurred()) + _, err = decodeSpeakerProfiles("", trusted) + Expect(status.Code(err)).To(Equal(codes.Unimplemented)) + raw, err = json.Marshal(diarizationResultFromProto(nil)) + Expect(err).NotTo(HaveOccurred()) + Expect(string(raw)).NotTo(ContainSubstring("speaker_profiles")) + }) +}) + +var _ = Describe("profile raw slot association", func() { + It("keeps sparse out-of-order slots separate from normalized IDs and names", func() { + out := diarizationResultFromProto(&pb.DiarizeResponse{Segments: []*pb.DiarizeSegment{ + {Speaker: "7", Text: "Hello", Name: "Ada", Start: 0, End: 1}, + {Speaker: "2", Text: "there", Name: "Ada", Start: 1, End: 2}, + }}) + Expect(out.Segments[0].Speaker).To(Equal("SPEAKER_00")) + Expect(out.Segments[0].Label).To(Equal("7")) + Expect(out.Segments[1].Speaker).To(Equal("SPEAKER_01")) + Expect(out.Segments[1].Label).To(Equal("2")) + Expect(out.Speakers[0].Label).To(Equal("7")) + Expect(out.Speakers[1].Label).To(Equal("2")) + raw, err := json.Marshal(out) + Expect(err).NotTo(HaveOccurred()) + Expect(string(raw)).To(ContainSubstring(`"label":"7"`)) + Expect(string(raw)).To(ContainSubstring(`"label":"2"`)) + }) +}) + +var _ = Describe("portable voice compatibility", func() { + It("rejects same-dimension incompatible encoders and unknown identity without changing legacy selection", func() { + identity := "sha256:" + strings.Repeat("a", 64) + m := &portableStatusBackend{identity: identity} + voices := []voicerecognition.KnownVoice{{ID: "match", Model: identity, Embedding: []float32{1, 0}}, {ID: "other", Model: "sha256:" + strings.Repeat("b", 64), Embedding: []float32{0, 1}}, {ID: "legacy", Model: "speaker.gguf", Embedding: []float32{1, 0}}} + got := compatiblePortableVoices(context.Background(), m, voices) + Expect(got).To(HaveLen(2)) + Expect(got[0].ID).To(Equal("match")) + Expect(got[1].ID).To(Equal("legacy")) + m.identity = "" + got = compatiblePortableVoices(context.Background(), m, voices) + Expect(got).To(HaveLen(1)) + Expect(got[0].ID).To(Equal("legacy")) + }) +}) + +type portableStatusBackend struct { + grpcPkg.Backend + identity string + dimension int32 +} + +func (m *portableStatusBackend) Status(context.Context) (*pb.StatusResponse, error) { + dim := m.dimension + if dim == 0 { + dim = 2 + } + return &pb.StatusResponse{SpeakerEncoder: &pb.SpeakerEncoder{Identity: m.identity, Dimension: dim}}, nil +} + +var _ = Describe("selection before portable compatibility", func() { + It("keeps legacy 192 candidates regardless of unrelated portable 256 order", func() { + identity := "sha256:" + strings.Repeat("a", 64) + makeEntry := func(id, tag string, dim int) voicerecognition.Entry { + v := make([]float32, dim) + v[0] = 1 + return voicerecognition.Entry{Metadata: voicerecognition.Metadata{ID: id, Name: id, Model: tag}, Embedding: v} + } + entries := []voicerecognition.Entry{ + makeEntry("portable", "sha256:"+strings.Repeat("b", 64), 256), + makeEntry("legacy", "", 192), + makeEntry("tagged", "speaker.gguf", 192), + makeEntry("wrong-size", "speaker.gguf", 256), + } + m := &portableStatusBackend{identity: identity, dimension: 192} + for _, pair := range [][]voicerecognition.Entry{{entries[0], entries[1]}, {entries[1], entries[0]}} { + selected := voicerecognition.SelectKnownVoices(pair, "speaker.gguf") + got := compatiblePortableVoices(context.Background(), m, selected.Voices) + Expect(got).To(HaveLen(1)) + Expect(got[0].ID).To(Equal("legacy")) + } + for i := 0; i < len(entries); i++ { + entries = append(entries[1:], entries[0]) + selected := voicerecognition.SelectKnownVoices(entries, "speaker.gguf") + got := compatiblePortableVoices(context.Background(), m, selected.Voices) + Expect(got).To(HaveLen(2)) + Expect(got[0].ID).To(Equal("tagged")) + Expect(got[1].ID).To(Equal("legacy")) + offline := (&DiarizationRequest{KnownVoices: got}).toProto(2, "model") + var live liveOptions + WithKnownVoices(got)(&live) + Expect(liveConfigProto("", live).KnownVoices).To(Equal(offline.KnownVoices)) + } + }) +}) diff --git a/core/backend/transcript_live.go b/core/backend/transcript_live.go index 49b69cd27..7fcd27f82 100644 --- a/core/backend/transcript_live.go +++ b/core/backend/transcript_live.go @@ -243,7 +243,7 @@ func WithKnownVoices(v []voicerecognition.KnownVoice) LiveOption { func liveConfigProto(language string, o liveOptions) *proto.TranscriptLiveConfig { cfg := &proto.TranscriptLiveConfig{Language: language, SampleRate: liveSampleRate} for _, v := range o.knownVoices { - cfg.KnownVoices = append(cfg.KnownVoices, &proto.KnownVoice{Name: v.Name, Embedding: v.Embedding, Model: v.Model}) + cfg.KnownVoices = append(cfg.KnownVoices, &proto.KnownVoice{Id: v.ID, Name: v.Name, Embedding: v.Embedding, Model: v.Model}) } return cfg } @@ -269,6 +269,7 @@ func ModelTranscriptionLive(ctx context.Context, language string, if err != nil { return nil, err } + lo.knownVoices = compatiblePortableVoices(ctx, transcriptionModel, lo.knownVoices) release, err := AcquireGlobalBackendSlot() if err != nil { return nil, err diff --git a/core/http/endpoints/localai/portable_http_test.go b/core/http/endpoints/localai/portable_http_test.go new file mode 100644 index 000000000..8009c5cb9 --- /dev/null +++ b/core/http/endpoints/localai/portable_http_test.go @@ -0,0 +1,374 @@ +// SPDX-License-Identifier: MIT +// +//nolint:errcheck,forbidigo // These focused HTTP harnesses use testing.T and assert response status inline. +package localai_test + +import ( + "bytes" + "context" + "encoding/json" + "github.com/mudler/LocalAI/core/application" + "github.com/mudler/LocalAI/core/trace/tracepersist" + "mime/multipart" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/auth" + "github.com/mudler/LocalAI/core/http/endpoints/localai" + "github.com/mudler/LocalAI/core/http/endpoints/openai" + "github.com/mudler/LocalAI/core/http/middleware" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/voicerecognition" + grpcpkg "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + ggrpc "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "gorm.io/gorm" +) + +type profileHTTPBackend struct { + grpcpkg.Backend + profiles schema.SpeakerProfiles + last *pb.DiarizeRequest + embeds int + unsupported bool + values [][]byte +} + +func (b *profileHTTPBackend) Status(context.Context) (*pb.StatusResponse, error) { + if b.unsupported { + return nil, status.Error(codes.Unimplemented, "no profiles") + } + return &pb.StatusResponse{SpeakerEncoder: &pb.SpeakerEncoder{Identity: b.profiles.Encoder.Identity, Dimension: int32(b.profiles.Encoder.Dimension)}}, nil +} +func (b *profileHTTPBackend) Diarize(_ context.Context, r *pb.DiarizeRequest, _ ...ggrpc.CallOption) (*pb.DiarizeResponse, error) { + b.last = r + raw, _ := json.Marshal(b.profiles) + return &pb.DiarizeResponse{Segments: []*pb.DiarizeSegment{{Speaker: "7", Start: 0, End: 3, Text: "Hello"}}, SpeakerProfilesJson: string(raw)}, nil +} +func (b *profileHTTPBackend) VoiceEmbed(context.Context, *pb.VoiceEmbedRequest, ...ggrpc.CallOption) (*pb.VoiceEmbedResponse, error) { + b.embeds++ + return &pb.VoiceEmbedResponse{Embedding: []float32{1, 0}, Model: "legacy.gguf"}, nil +} +func (b *profileHTTPBackend) StoresSet(_ context.Context, in *pb.StoresSetOptions, _ ...ggrpc.CallOption) (*pb.Result, error) { + for _, v := range in.Values { + b.values = append(b.values, append([]byte(nil), v.Bytes...)) + } + return &pb.Result{Success: true}, nil +} +func profileFixture() schema.SpeakerProfiles { + return schema.SpeakerProfiles{Version: 1, Encoder: schema.SpeakerEncoder{Identity: "sha256:" + strings.Repeat("a", 64), Dimension: 2}, Speakers: []schema.SpeakerProfile{{Speaker: 7, CleanDuration: 3, Intervals: []schema.SpeakerProfileInterval{{Start: 0, End: 3}}, Embedding: []float32{1, 0}}, {Speaker: 2, CleanDuration: 3, Intervals: []schema.SpeakerProfileInterval{{Start: 3, End: 6}}, Embedding: []float32{0, 1}}}} +} +func profileServer(b *profileHTTPBackend, denied bool) (*echo.Echo, voicerecognition.Registry) { + ml := model.NewModelLoader(&system.SystemState{}) + ml.SetModelRouter(func(_ context.Context, id string, _, _, _, _ string, _ *pb.ModelOptions, _ bool) (*model.Model, error) { + return model.NewModelWithClient(id, "test://profiles", b), nil + }) + cfg := &config.ModelConfig{Name: "test", Backend: "stub"} + cfg.SetDefaults() + cfg.Options = []string{"speaker_model:speaker.gguf"} + reg := voicerecognition.NewStoreRegistry(func(context.Context, string) (grpcpkg.Backend, error) { return b, nil }, "test", 0) + e := echo.New() + var db *gorm.DB + if denied { + db = &gorm.DB{} + } + setup := func(voice bool) echo.MiddlewareFunc { + return func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + if denied { + c.Set("auth_user", &auth.User{ID: "user", Role: "user"}) + c.Set("auth_permissions", &auth.UserPermission{Permissions: auth.PermissionMap{auth.FeatureVoiceRecognition: false}}) + } + if voice { + var r schema.VoiceRegisterRequest + if err := c.Bind(&r); err != nil { + return err + } + c.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &r) + } else { + var r schema.OpenAIRequest + if strings.HasPrefix(c.Request().Header.Get("Content-Type"), "application/json") { + if err := c.Bind(&r); err != nil { + return err + } + } + r.Model = "test" + c.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &r) + } + c.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg) + return next(c) + } + } + } + for _, route := range []string{"/v1/audio/diarization", "/audio/diarization"} { + e.POST(route, openai.DiarizationEndpoint(nil, ml, &config.ApplicationConfig{SystemState: &system.SystemState{}}, reg, db), setup(false)) + } + e.POST("/v1/voice/identify", localai.VoiceIdentifyEndpoint(nil, ml, &config.ApplicationConfig{SystemState: &system.SystemState{}}, reg), func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + var r schema.VoiceIdentifyRequest + if err := c.Bind(&r); err != nil { + return err + } + c.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &r) + c.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg) + return next(c) + } + }) + e.POST("/v1/voice/register", localai.VoiceRegisterEndpoint(nil, ml, &config.ApplicationConfig{SystemState: &system.SystemState{}}, reg), setup(true), auth.RequireFeature(db, auth.FeatureVoiceRecognition)) + return e, reg +} +func profileJSON(e *echo.Echo, route string, payload any) *httptest.ResponseRecorder { + raw, _ := json.Marshal(payload) + r := httptest.NewRequest("POST", route, bytes.NewReader(raw)) + r.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + e.ServeHTTP(w, r) + return w +} +func TestPortableProfileHTTP(t *testing.T) { + b := &profileHTTPBackend{profiles: profileFixture()} + e, reg := profileServer(b, false) + for _, on := range []bool{false, true} { + w := profileJSON(e, "/v1/audio/diarization", map[string]any{"model": "test", "file": "YXVkaW8=", "include_speaker_profiles": on, "include_text": on}) + if w.Code != 200 { + t.Fatal(w.Code, w.Body.String()) + } + if bytes.Contains(w.Body.Bytes(), []byte("speaker_profiles")) != on || b.last.IncludeSpeakerProfiles != on { + t.Fatal(w.Body.String()) + } + if on && !bytes.Contains(w.Body.Bytes(), []byte("Hello")) { + t.Fatal("missing combined text") + } + } + for _, slot := range []int{7, 2} { + w := profileJSON(e, "/v1/voice/register", map[string]any{"model": "test", "name": "Ada", "speaker_slot": slot, "speaker_profiles": b.profiles}) + if w.Code != 200 { + t.Fatal(w.Code, w.Body.String()) + } + } + entries, err := reg.List(context.Background()) + if err != nil || len(entries) != 2 || entries[0].Metadata.ID == entries[1].Metadata.ID || entries[0].Embedding[0] == entries[1].Embedding[0] || entries[0].Metadata.Model != b.profiles.Encoder.Identity { + t.Fatal(entries, err) + } + // Replay uses registration IDs and exact loaded identity, not names. + wReplay := profileJSON(e, "/v1/audio/diarization", map[string]any{"model": "test", "file": "YQ=="}) + if wReplay.Code != 200 || len(b.last.KnownVoices) != 2 || b.last.KnownVoices[0].Id == b.last.KnownVoices[1].Id { + t.Fatal(wReplay.Code, b.last) + } + wIdentify := profileJSON(e, "/v1/voice/identify", map[string]any{"model": "test", "audio": "YQ=="}) + var identified schema.VoiceIdentifyResponse + if wIdentify.Code != 200 || json.Unmarshal(wIdentify.Body.Bytes(), &identified) != nil || len(identified.Matches) != 2 { + t.Fatal(wIdentify.Code, wIdentify.Body.String()) + } + b.profiles.Encoder.Identity = "sha256:" + strings.Repeat("b", 64) + profileJSON(e, "/v1/audio/diarization", map[string]any{"model": "test", "file": "YQ=="}) + if len(b.last.KnownVoices) != 0 { + t.Fatal("cross-model replay") + } + wIdentify = profileJSON(e, "/v1/voice/identify", map[string]any{"model": "test", "audio": "YQ=="}) + json.Unmarshal(wIdentify.Body.Bytes(), &identified) + if len(identified.Matches) != 0 { + t.Fatal("cross-model identify", wIdentify.Body.String()) + } + b.profiles = profileFixture() + if b.embeds != 2 { + t.Fatal("profile enrollment invoked audio encoder") + } + w := profileJSON(e, "/v1/voice/register", map[string]any{"model": "test", "name": "Legacy", "audio": "YXVkaW8="}) + if w.Code != 200 || b.embeds != 3 { + t.Fatal(w.Code, w.Body.String()) + } + for _, kind := range []string{"wrongmodel", "zero", "wrongdim", "unavailable", "missing", "audio", "version"} { + p := profileFixture() + slot := 7 + payload := map[string]any{"model": "test", "name": "Ada", "speaker_slot": slot} + switch kind { + case "wrongmodel": + p.Encoder.Identity = "sha256:" + strings.Repeat("b", 64) + case "zero": + p.Speakers[0].Embedding = []float32{0, 0} + case "wrongdim": + p.Speakers[0].Embedding = []float32{1} + case "unavailable": + reason := "overlap" + p.Speakers[0].UnavailableReason = &reason + p.Speakers[0].Embedding = nil + case "missing": + payload["speaker_slot"] = 0 + case "audio": + payload["audio"] = "YQ==" + case "version": + p.Version = 2 + } + payload["speaker_profiles"] = p + w = profileJSON(e, "/v1/voice/register", payload) + if w.Code != 400 { + t.Fatalf("%s: %d %s", kind, w.Code, w.Body.String()) + } + } + deniedServer, _ := profileServer(b, true) + deniedResponse := profileJSON(deniedServer, "/v1/voice/register", map[string]any{"model": "test", "name": "Ada", "speaker_slot": 7, "speaker_profiles": b.profiles}) + if deniedResponse.Code != 403 { + t.Fatal(deniedResponse.Code, deniedResponse.Body.String()) + } + // Non-JSON numeric values (NaN) must fail parsing before enrollment. + raw := `{"model":"test","name":"Ada","speaker_slot":7,"speaker_profiles":{"version":1,"speakers":[{"embedding":[NaN]}]}}` + request := httptest.NewRequest("POST", "/v1/voice/register", strings.NewReader(raw)) + request.Header.Set("Content-Type", "application/json") + recorder := httptest.NewRecorder() + e.ServeHTTP(recorder, request) + if recorder.Code != 400 { + t.Fatal(recorder.Code, recorder.Body.String()) + } + b.unsupported = true + w = profileJSON(e, "/v1/voice/register", map[string]any{"model": "test", "name": "Ada", "speaker_slot": 7, "speaker_profiles": b.profiles}) + if w.Code != 501 { + t.Fatal(w.Code, w.Body.String()) + } +} +func TestPortableProfileMultipartPermissions(t *testing.T) { + for _, denied := range []bool{false, true} { + for _, route := range []string{"/v1/audio/diarization", "/audio/diarization"} { + b := &profileHTTPBackend{profiles: profileFixture()} + e, _ := profileServer(b, denied) + body := &bytes.Buffer{} + mw := multipart.NewWriter(body) + mw.WriteField("model", "test") + mw.WriteField("include_speaker_profiles", "true") + mw.WriteField("include_text", "true") + f, _ := mw.CreateFormFile("file", "sample.wav") + f.Write([]byte("audio")) + mw.Close() + r := httptest.NewRequest("POST", route, body) + r.Header.Set("Content-Type", mw.FormDataContentType()) + w := httptest.NewRecorder() + e.ServeHTTP(w, r) + want := 200 + if denied { + want = 403 + } + if w.Code != want { + t.Fatal(route, w.Code, w.Body.String()) + } + if denied && b.last != nil { + t.Fatal("denied request reached backend") + } + if denied { + w = profileJSON(e, route, map[string]any{"model": "test", "file": "YQ==", "include_speaker_profiles": true}) + if w.Code != 403 { + t.Fatal(w.Code, w.Body.String()) + } + } + } + } + b := &profileHTTPBackend{profiles: profileFixture()} + e, _ := profileServer(b, false) + for _, format := range []string{"rttm", "invalid"} { + w := profileJSON(e, "/v1/audio/diarization", map[string]any{"model": "test", "file": "YQ==", "response_format": format, "include_speaker_profiles": true}) + if w.Code != 400 || b.last != nil { + t.Fatal(w.Code, w.Body.String()) + } + } +} +func (b *profileHTTPBackend) HealthCheck(context.Context) (bool, error) { return true, nil } + +func (b *profileHTTPBackend) StoresFind(_ context.Context, in *pb.StoresFindOptions, _ ...ggrpc.CallOption) (*pb.StoresFindResult, error) { + r := &pb.StoresFindResult{} + for _, v := range b.values { + r.Keys = append(r.Keys, &pb.StoresKey{Floats: in.Key.Floats}) + r.Values = append(r.Values, &pb.StoresValue{Bytes: v}) + r.Similarities = append(r.Similarities, 1) + } + return r, nil +} +func TestPortableProfileModelAccess(t *testing.T) { + b := &profileHTTPBackend{profiles: profileFixture()} + e, _ := profileServer(b, false) + e.Use(func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + c.Set("auth_user", &auth.User{ID: "limited", Role: "user"}) + c.Set("auth_permissions", &auth.UserPermission{AllowedModels: auth.ModelAllowlist{Enabled: true, Models: []string{"different-model"}}}) + return next(c) + } + }, auth.RequireModelAccess(&gorm.DB{})) + for _, route := range []string{"/v1/audio/diarization", "/audio/diarization", "/v1/voice/register"} { + w := profileJSON(e, route, map[string]any{"model": "test", "file": "YQ==", "include_speaker_profiles": true, "name": "Ada", "speaker_profiles": b.profiles, "speaker_slot": 7}) + if w.Code != 403 || b.last != nil { + t.Fatal(route, w.Code, w.Body.String()) + } + } +} + +func TestPortableProfilesNeverPersistInAPITraces(t *testing.T) { + root := t.TempDir() + app, err := application.New(config.EnableTracing, config.WithDataPath(root), config.WithDisableLocalAIAssistant(true), config.WithDisableStats(true), config.WithSystemState(&system.SystemState{Model: system.Model{ModelsPath: root}, Backend: system.Backend{BackendsPath: root}})) + if err != nil { + t.Fatal(err) + } + defer app.Shutdown() + b := &profileHTTPBackend{profiles: profileFixture()} + // Slot zero must be distinguishable from an omitted slot. + b.profiles.Speakers[0].Speaker = 0 + e, _ := profileServer(b, false) + e.Use(middleware.TraceMiddleware(app)) + e.POST("/ordinary", func(c echo.Context) error { return c.JSON(200, map[string]string{"result": "benign"}) }) + for _, route := range []string{"/v1/audio/diarization", "/audio/diarization"} { + for _, on := range []bool{true, false} { + w := profileJSON(e, route, map[string]any{"model": "test", "file": "YXVkaW8=", "include_speaker_profiles": on}) + if w.Code != 200 { + t.Fatal(w.Code, w.Body.String()) + } + if bytes.Contains(w.Body.Bytes(), []byte(`"embedding":[1,0]`)) != on { + t.Fatal(w.Body.String()) + } + } + } + w := profileJSON(e, "/v1/voice/register", map[string]any{"model": "test", "name": "Ada", "speaker_slot": 0, "speaker_profiles": b.profiles}) + if w.Code != 200 { + t.Fatal(w.Code, w.Body.String()) + } + // Queue a nonsensitive trace last: its persistence is a barrier for earlier requests. + profileJSON(e, "/ordinary", map[string]string{"message": "benign"}) + store, err := tracepersist.New[middleware.APIExchange](filepath.Join(root, "traces", "api"), 100) + if err != nil { + t.Fatal(err) + } + deadline := time.Now().Add(5 * time.Second) + for { + records, err := store.Load() + if err != nil { + t.Fatal(err) + } + found := false + for _, r := range records { + if r.Request.Path == "/ordinary" { + found = true + } + } + if found { + if len(records) != 1 { + t.Fatalf("persisted biometric exchanges: %d records", len(records)) + } + if string(*records[0].Response.Body) != "{\"result\":\"benign\"}\n" { + t.Fatal("ordinary trace changed") + } + if len(middleware.GetTraces()) != 1 { + t.Fatal("biometric exchange captured in memory") + } + break + } + if time.Now().After(deadline) { + t.Fatal("ordinary trace not persisted") + } + time.Sleep(10 * time.Millisecond) + } +} diff --git a/core/http/endpoints/localai/portable_register_test.go b/core/http/endpoints/localai/portable_register_test.go new file mode 100644 index 000000000..0c2a0bb5a --- /dev/null +++ b/core/http/endpoints/localai/portable_register_test.go @@ -0,0 +1,38 @@ +// SPDX-License-Identifier: MIT +// +//nolint:forbidigo // This focused HTTP harness uses testing.T and asserts response status inline. +package localai_test + +import ( + "bytes" + "encoding/json" + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/localai" + "github.com/mudler/LocalAI/core/http/middleware" + "github.com/mudler/LocalAI/core/schema" + "net/http/httptest" + "testing" +) + +func TestPortableRegisterRequiresExplicitSlot(t *testing.T) { + e := echo.New() + e.POST("/v1/voice/register", localai.VoiceRegisterEndpoint(nil, nil, nil, nil), func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + var r schema.VoiceRegisterRequest + if err := json.NewDecoder(c.Request().Body).Decode(&r); err != nil { + return err + } + c.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &r) + c.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{}) + return next(c) + } + }) + r := httptest.NewRequest("POST", "/v1/voice/register", bytes.NewBufferString(`{"model":"test","name":"Ada","speaker_profiles":{"version":1}}`)) + r.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + e.ServeHTTP(w, r) + if w.Code != 400 || !bytes.Contains(w.Body.Bytes(), []byte("speaker_slot")) { + t.Fatalf("expected explicit slot validation, got %d %s", w.Code, w.Body.String()) + } +} diff --git a/core/http/endpoints/localai/voice_identify.go b/core/http/endpoints/localai/voice_identify.go index eda5aec3d..dac259b59 100644 --- a/core/http/endpoints/localai/voice_identify.go +++ b/core/http/endpoints/localai/voice_identify.go @@ -3,6 +3,7 @@ package localai import ( "cmp" "net/http" + "strings" "github.com/labstack/echo/v4" "github.com/mudler/LocalAI/core/backend" @@ -57,6 +58,28 @@ func VoiceIdentifyEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, return err } + // Portable vectors require exact loaded-weight identity. Legacy audio + // registrations retain their filename-tag compatibility behavior. + var trusted schema.SpeakerEncoder + var trustedErr error + for _, m := range matches { + if strings.HasPrefix(m.Metadata.Model, "sha256:") { + trusted, trustedErr = backend.ModelSpeakerEncoder(c.Request().Context(), ml, *cfg, appConfig) + break + } + } + filtered := matches[:0] + for _, m := range matches { + if strings.HasPrefix(m.Metadata.Model, "sha256:") { + if trustedErr != nil || m.Metadata.Model != trusted.Identity || len(embed.GetEmbedding()) != trusted.Dimension { + continue + } + } else if m.Metadata.Model != "" && voicerecognition.EncoderTag(m.Metadata.Model) != voicerecognition.EncoderTag(embed.GetModel()) { + continue + } + filtered = append(filtered, m) + } + matches = filtered response := schema.VoiceIdentifyResponse{ Matches: make([]schema.VoiceIdentifyMatch, len(matches)), } diff --git a/core/http/endpoints/localai/voice_register.go b/core/http/endpoints/localai/voice_register.go index 4a7e36806..9ae4d1785 100644 --- a/core/http/endpoints/localai/voice_register.go +++ b/core/http/endpoints/localai/voice_register.go @@ -10,11 +10,11 @@ import ( "github.com/mudler/LocalAI/core/schema" "github.com/mudler/LocalAI/core/services/voicerecognition" "github.com/mudler/LocalAI/pkg/model" - "github.com/mudler/xlog" ) // VoiceRegisterEndpoint enrolls a speaker into the 1:N identification store. // @Summary Register a speaker for 1:N identification. +// @Description Supply either audio or speaker_profiles plus an explicit numeric speaker_slot. The selected model must expose matching trusted encoder metadata for portable enrollment. Registrations are global and ephemeral, with a fresh ID for each request. // @Tags voice-recognition // @Param request body schema.VoiceRegisterRequest true "query params" // @Success 200 {object} schema.VoiceRegisterResponse "Response" @@ -33,19 +33,37 @@ func VoiceRegisterEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, return echo.NewHTTPError(http.StatusBadRequest, "name is required") } - audio, cleanup, err := decodeAudioInput(input.Audio) - if err != nil { - return err + var embedding []float32 + var encoder string + if input.SpeakerProfiles != nil { + if input.Audio != "" || input.SpeakerSlot == nil { + return echo.NewHTTPError(http.StatusBadRequest, "speaker_profiles requires speaker_slot and excludes audio") + } + trusted, err := backend.ModelSpeakerEncoder(c.Request().Context(), ml, *cfg, appConfig) + if err != nil { + return mapBackendError(err) + } + selected, err := input.SpeakerProfiles.Select(*input.SpeakerSlot, trusted) + if err != nil { + return echo.NewHTTPError(http.StatusBadRequest, err.Error()) + } + embedding, encoder = selected.Embedding, trusted.Identity + } else { + if input.SpeakerSlot != nil { + return echo.NewHTTPError(http.StatusBadRequest, "speaker_slot requires speaker_profiles") + } + audio, cleanup, err := decodeAudioInput(input.Audio) + if err != nil { + return err + } + defer cleanup() + res, err := backend.VoiceEmbed(c.Request().Context(), audio, ml, appConfig, *cfg) + if err != nil { + return mapBackendError(err) + } + embedding, encoder = res.GetEmbedding(), res.GetModel() } - defer cleanup() - - xlog.Debug("VoiceRegister", "model", cfg.Name, "name", input.Name) - res, err := backend.VoiceEmbed(c.Request().Context(), audio, ml, appConfig, *cfg) - if err != nil { - return mapBackendError(err) - } - - stored, err := registry.Register(c.Request().Context(), res.GetEmbedding(), voiceMetadata(input.Name, input.Labels, res.GetModel())) + stored, err := registry.Register(c.Request().Context(), embedding, voiceMetadata(input.Name, input.Labels, encoder)) if err != nil { return err } diff --git a/core/http/endpoints/openai/diarization.go b/core/http/endpoints/openai/diarization.go index 99dbf3ce4..bd12bfadb 100644 --- a/core/http/endpoints/openai/diarization.go +++ b/core/http/endpoints/openai/diarization.go @@ -1,7 +1,9 @@ package openai import ( + "bytes" "context" + "encoding/base64" "fmt" "io" "net/http" @@ -15,10 +17,14 @@ import ( "github.com/labstack/echo/v4" "github.com/mudler/LocalAI/core/backend" "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/auth" "github.com/mudler/LocalAI/core/http/middleware" "github.com/mudler/LocalAI/core/schema" "github.com/mudler/LocalAI/core/services/voicerecognition" model "github.com/mudler/LocalAI/pkg/model" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "gorm.io/gorm" "github.com/mudler/xlog" ) @@ -35,8 +41,9 @@ import ( // (NIST RTTM, the standard interchange format used by pyannote/dscore). // // @Summary Identify speakers in audio (who spoke when). +// @Description JSON accepts model, file (raw base64 audio), include_text, include_speaker_profiles and response_format. Profiles require voice-recognition permission and json or verbose_json; unsupported backends return 501. // @Tags audio -// @accept multipart/form-data +// @accept multipart/form-data,json // @Param model formData string true "model" // @Param file formData file true "audio file" // @Param num_speakers formData int false "exact speaker count (>0 forces; 0 = auto)" @@ -46,11 +53,12 @@ import ( // @Param min_duration_on formData number false "discard segments shorter than this (seconds)" // @Param min_duration_off formData number false "merge gaps shorter than this (seconds)" // @Param language formData string false "audio language hint (only meaningful for backends that bundle ASR)" +// @Param include_speaker_profiles formData boolean false "export portable biometric profiles (voice-recognition permission; JSON formats only)" // @Param include_text formData boolean false "include per-segment transcript when the backend supports it" // @Param response_format formData string false "json (default), verbose_json, or rttm" // @Success 200 {object} schema.DiarizationResult // @Router /v1/audio/diarization [post] -func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, registry voicerecognition.Registry) echo.HandlerFunc { +func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, registry voicerecognition.Registry, authDB ...*gorm.DB) echo.HandlerFunc { return func(c echo.Context) error { input, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) if !ok || input.Model == "" { @@ -63,8 +71,20 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap } req := backend.DiarizationRequest{ - Language: input.Language, - IncludeText: parseFormBool(c, "include_text", false), + Language: input.Language, + IncludeText: parseFormBool(c, "include_text", input.IncludeText), + IncludeSpeakerProfiles: parseFormBool(c, "include_speaker_profiles", input.IncludeSpeakerProfiles), + } + if req.IncludeSpeakerProfiles { + var db *gorm.DB + if len(authDB) > 0 { + db = authDB[0] + } + allowed := false + err := auth.RequireFeature(db, auth.FeatureVoiceRecognition)(func(c echo.Context) error { allowed = true; return nil })(c) + if err != nil || !allowed { + return err + } } req.NumSpeakers = int32(parseFormInt(c, "num_speakers", 0)) req.MinSpeakers = int32(parseFormInt(c, "min_speakers", 0)) @@ -75,6 +95,15 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap attachKnownVoices(c.Request().Context(), &req, modelConfig.Options, registry) responseFormat := schema.DiarizationResponseFormatType(strings.ToLower(c.FormValue("response_format"))) + if responseFormat == "" { + if input.ResponseFormat != nil { + f, ok := input.ResponseFormat.(string) + if !ok { + return echo.NewHTTPError(http.StatusBadRequest, "response_format must be a string") + } + responseFormat = schema.DiarizationResponseFormatType(strings.ToLower(f)) + } + } if responseFormat == "" { responseFormat = schema.DiarizationResponseFormatJson } @@ -86,15 +115,30 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format (expected: json, verbose_json, rttm)") } - file, err := uploadedFile(c, "file") - if err != nil { - return err + if req.IncludeSpeakerProfiles && responseFormat == schema.DiarizationResponseFormatRTTM { + return echo.NewHTTPError(http.StatusBadRequest, "speaker_profiles requires json or verbose_json") } - f, err := file.Open() - if err != nil { - return err + var sourceName = "audio.wav" + var reader io.ReadCloser + if strings.HasPrefix(c.Request().Header.Get(echo.HeaderContentType), echo.MIMEApplicationJSON) { + raw, err := base64.StdEncoding.DecodeString(input.File) + if err != nil || len(raw) == 0 { + return echo.NewHTTPError(http.StatusBadRequest, "file must be base64 audio") + } + reader = io.NopCloser(bytes.NewReader(raw)) + } else { + file, err := uploadedFile(c, "file") + if err != nil { + return err + } + f, err := file.Open() + if err != nil { + return err + } + reader = f + sourceName = path.Base(file.Filename) } - defer func() { _ = f.Close() }() + defer func() { _ = reader.Close() }() dir, err := os.MkdirTemp("", "diarize") if err != nil { @@ -102,13 +146,13 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap } defer func() { _ = os.RemoveAll(dir) }() - dst := filepath.Join(dir, path.Base(file.Filename)) + dst := filepath.Join(dir, sourceName) dstFile, err := os.Create(dst) if err != nil { return err } - if _, err := io.Copy(dstFile, f); err != nil { - xlog.Debug("Audio file copying error", "filename", file.Filename, "dst", dst, "error", err) + if _, err := io.Copy(dstFile, reader); err != nil { + xlog.Debug("Audio file copying error", "filename", sourceName, "dst", dst, "error", err) _ = dstFile.Close() return err } @@ -117,20 +161,28 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap result, err := backend.ModelDiarization(c.Request().Context(), req, ml, *modelConfig, appConfig) if err != nil { + if status.Code(err) == codes.Unimplemented { + return echo.NewHTTPError(http.StatusNotImplemented, status.Convert(err).Message()) + } return err } + if !req.IncludeSpeakerProfiles { + result.SpeakerProfiles = nil + } switch responseFormat { case schema.DiarizationResponseFormatRTTM: c.Response().Header().Set(echo.HeaderContentType, "text/plain; charset=utf-8") - return c.String(http.StatusOK, renderRTTM(result, file.Filename)) + return c.String(http.StatusOK, renderRTTM(result, sourceName)) case schema.DiarizationResponseFormatJson: // Default JSON: drop the heavy per-speaker summary and any - // optional per-segment text so simple consumers see a tight + // unrequested per-segment text so simple consumers see a tight // payload. verbose_json keeps everything. result.Speakers = nil for i := range result.Segments { - result.Segments[i].Text = "" + if !req.IncludeText { + result.Segments[i].Text = "" + } } return c.JSON(http.StatusOK, result) case schema.DiarizationResponseFormatJsonVerbose: diff --git a/core/http/middleware/trace.go b/core/http/middleware/trace.go index 0fd2a5cec..ec77abbe9 100644 --- a/core/http/middleware/trace.go +++ b/core/http/middleware/trace.go @@ -243,6 +243,15 @@ func TraceMiddleware(app *application.Application) echo.MiddlewareFunc { return next(c) } + // Biometric routes can carry vectors in either direction and JSON + // diarization carries base64 audio even without profile export. + // Exclude the whole exchange before reading or wrapping bodies, + // including registration if tracing is installed globally later. + switch c.Path() { + case "/v1/audio/diarization", "/audio/diarization", "/v1/voice/register": + return next(c) + } + ct, _, _ := mime.ParseMediaType(c.Request().Header.Get("Content-Type")) if ct != "application/json" { return next(c) diff --git a/core/http/react-ui/e2e/diarization-profiles.spec.js b/core/http/react-ui/e2e/diarization-profiles.spec.js new file mode 100644 index 000000000..080302971 --- /dev/null +++ b/core/http/react-ui/e2e/diarization-profiles.spec.js @@ -0,0 +1,193 @@ +// SPDX-License-Identifier: MIT +import { test, expect } from './coverage-fixtures.js' + +const profiles = { version: 1, encoder: { identity: `sha256:${'a'.repeat(64)}`, dimension: 2 }, speakers: [ + { speaker: 9, clean_duration: 0, intervals: [], unavailable_reason: 'insufficient_clean_speech' }, + { speaker: 0, clean_duration: 3, intervals: [{ start: 1, end: 2 }, { start: 4, end: 6 }], unavailable_reason: null, embedding: [1, 0] }, + { speaker: 7, clean_duration: 2, intervals: [{ start: 2, end: 4 }], unavailable_reason: null, embedding: [0, 1] }, + { speaker: 3, clean_duration: 2, intervals: [{ start: 6, end: 8 }], unavailable_reason: null, embedding: [0.6, 0.8] }, +] } +const result = { speakers: [ + { id: 'SPEAKER_00', label: '7', name: 'Known' }, + { id: 'SPEAKER_01', label: '3' }, + { id: 'SPEAKER_02', label: '9' }, + { id: 'SPEAKER_03', label: '0' }, +], segments: [ + { id: 0, speaker: 'SPEAKER_03', label: '0', start: 1, end: 2, text: 'First' }, + { id: 1, speaker: 'SPEAKER_01', label: '3', start: 6, end: 8, text: 'Second' }, + { id: 2, speaker: 'SPEAKER_03', label: '0', start: 4, end: 6, text: 'Third' }, +], speaker_profiles: profiles } +const file = { name: 'meeting.wav', mimeType: 'audio/wav', buffer: Buffer.alloc(64) } +async function setup(page, permission = true, diarizationPermission = true) { + await page.route('**/api/**', route => { + const url = route.request().url() + const data = url.endsWith('/auth/status') ? { authEnabled: true, user: { role: 'user', permissions: { audio_diarization: diarizationPermission, voice_recognition: permission } } } + : url.endsWith('/models/capabilities') ? { data: ['diarizer', 'other'].map(id => ({ id, capabilities: ['FLAG_DIARIZATION'] })) } : {} + return route.fulfill({ json: data }) + }) + await page.route('**/v1/audio/diarization', route => route.fulfill({ json: result })) + await page.goto('/app/diarization') + if (!diarizationPermission) return + await expect(page.getByRole('button', { name: 'diarizer', exact: true })).toBeVisible() + await page.getByLabel('Recording', { exact: true }).setInputFiles(file) +} +async function run(page) { + await page.getByLabel('Prepare speakers to remember').check() + await page.getByRole('button', { name: 'Diarize', exact: true }).click() + await expect(page.getByTestId('speaker-0')).toBeVisible() +} +const speaker = (page, slot) => page.getByTestId(`speaker-${slot}`) + +test('sparse raw slots, known/unavailable, explicit zero and duplicate names', async ({ page }) => { + await setup(page) + const requests = [] + await page.route('**/v1/voice/register', route => { requests.push(route.request().postDataJSON()); return route.fulfill({ json: { id: `id-${requests.length}`, name: 'Known', registered_at: '2026-01-01' } }) }) + const request = page.waitForRequest('**/v1/audio/diarization') + await run(page) + const body = (await request).postData() + for (const value of ['include_speaker_profiles', 'include_text', 'verbose_json']) expect(body).toContain(value) + await expect(speaker(page, 7)).toContainText('Known') + await expect(speaker(page, 7).getByRole('button', { name: 'Name and remember' })).toHaveCount(0) + await expect(speaker(page, 9)).toContainText('Not enough clean speech') + await expect(speaker(page, 9).getByRole('button', { name: 'Name and remember' })).toBeDisabled() + for (const slot of [0, 3]) { + await speaker(page, slot).getByRole('button', { name: 'Name and remember' }).click() + await page.getByLabel('Name', { exact: true }).fill('Known') + await page.getByRole('button', { name: 'Remember', exact: true }).click() + await expect(page.getByRole('dialog')).toHaveCount(0) + expect(requests.at(-1)).toEqual({ model: 'diarizer', name: 'Known', speaker_slot: slot, speaker_profiles: profiles }) + } + await expect(page.getByTestId('segments').getByText('Known', { exact: true })).toHaveCount(3) + const stored = await page.evaluate(() => JSON.parse(localStorage.getItem('localai_voice_enrollments'))) + expect(stored.map(x => x.id)).toEqual(['id-2', 'id-1']) + expect(JSON.stringify(stored)).not.toMatch(/embedding|speaker_profiles|sampleUrl/) + await page.getByRole('link', { name: 'Manage remembered voices' }).click() + await page.getByRole('tab', { name: 'Enrollment' }).click() + await expect(page.getByText('Known', { exact: true })).toHaveCount(2) +}) + +test('save failure preserves input, no premature relabel, duplicate submission disabled', async ({ page }) => { + await setup(page); await run(page) + let release + await page.route('**/v1/voice/register', async route => { await new Promise(r => { release = r }); await route.fulfill({ status: 500, json: { error: 'Try again' } }) }) + await speaker(page, 0).getByRole('button', { name: 'Name and remember' }).click() + await page.getByLabel('Name', { exact: true }).fill('Ada') + await page.getByRole('button', { name: 'Remember', exact: true }).click() + await expect(page.getByRole('button', { name: 'Saving…' })).toBeDisabled() + await expect(speaker(page, 0)).not.toContainText('Ada') + release() + await expect(page.getByRole('dialog')).toContainText('Try again') + await expect(page.getByLabel('Name', { exact: true })).toHaveValue('Ada') + await page.route('**/v1/voice/register', route => route.fulfill({ json: { id: 'ada', name: 'Ada' } })) + await page.getByRole('button', { name: 'Remember', exact: true }).click() + await expect(speaker(page, 0)).toContainText('Ada') +}) + +test('recognition permission does not block normal diarization', async ({ page }) => { + await setup(page, false) + await expect(page.getByLabel('Prepare speakers to remember')).toHaveCount(0) + const req = page.waitForRequest('**/v1/audio/diarization') + await page.getByRole('button', { name: 'Diarize', exact: true }).click() + expect((await req).postData()).not.toContain('include_speaker_profiles') + await expect(page.getByTestId('segments')).toContainText('First') + await expect(page.getByRole('button', { name: 'Name and remember' })).toHaveCount(0) +}) + +test('preview uses clean intervals and revokes original object URL on replacement', async ({ page }) => { + await page.addInitScript(() => { + window.revoked = [] + const revoke = URL.revokeObjectURL.bind(URL) + URL.revokeObjectURL = url => { window.revoked.push(url); revoke(url) } + HTMLMediaElement.prototype.play = function () { window.playedAt = this.currentTime; return Promise.resolve() } + HTMLMediaElement.prototype.pause = function () { window.paused = true } + }) + await setup(page); await run(page) + await speaker(page, 0).getByRole('button', { name: 'Preview 1' }).click() + expect(await page.evaluate(() => window.playedAt)).toBe(1) + const audio = page.locator('audio') + const url = await audio.getAttribute('src') + await audio.evaluate(el => { el.currentTime = 2.1; el.dispatchEvent(new Event('timeupdate')) }) + expect(await page.evaluate(() => window.paused)).toBe(true) + await speaker(page, 0).getByRole('button', { name: 'Preview 2' }).click() + expect(await page.evaluate(() => window.playedAt)).toBe(4) + await page.getByLabel('Recording', { exact: true }).setInputFiles({ ...file, name: 'new.wav' }) + await expect(speaker(page, 0)).toHaveCount(0) + await expect.poll(() => page.evaluate(url => window.revoked.includes(url), url)).toBe(true) +}) + +for (const change of ['recording', 'model']) test(`late inference and save cannot relabel changed ${change}`, async ({ page }) => { + await setup(page) + let release + await page.route('**/v1/audio/diarization', async route => { await new Promise(r => { release = r }); await route.fulfill({ json: result }) }) + await page.getByLabel('Prepare speakers to remember').check() + const req = page.waitForRequest('**/v1/audio/diarization') + await page.getByRole('button', { name: 'Diarize', exact: true }).click(); await req + const replace = async () => { + if (change === 'recording') await page.getByLabel('Recording', { exact: true }).setInputFiles({ ...file, name: 'new.wav' }) + else { await page.getByRole('button', { name: 'diarizer', exact: true }).click(); await page.getByRole('option', { name: 'other', exact: true }).click() } + } + const inferenceResponse = page.waitForResponse('**/v1/audio/diarization') + await replace(); release(); await inferenceResponse + await expect(speaker(page, 0)).toHaveCount(0) + await page.route('**/v1/audio/diarization', route => route.fulfill({ json: result })) + await run(page) + await page.route('**/v1/voice/register', async route => { await new Promise(r => { release = r }); await route.fulfill({ json: { id: 'late', name: 'Late' } }) }) + await speaker(page, 0).getByRole('button', { name: 'Name and remember' }).click() + await page.getByLabel('Name', { exact: true }).fill('Late') + const save = page.waitForRequest('**/v1/voice/register') + await page.getByRole('button', { name: 'Remember', exact: true }).click(); await save + // Input state may change while a save is pending; simulate without dismissing it. + if (change === 'recording') await page.getByLabel('Recording', { exact: true }).setInputFiles({ ...file, name: 'third.wav' }) + else { + await page.getByRole('button', { name: 'other', exact: true }).evaluate(el => el.click()) + await page.getByRole('option', { name: /^diarizer/ }).evaluate(el => el.click()) + } + const saveResponse = page.waitForResponse('**/v1/voice/register') + release(); await saveResponse + await expect(page.getByRole('dialog')).toHaveCount(0) + await expect(speaker(page, 0)).toHaveCount(0) +}) + +test('unsupported export gives actionable error, never silently falls back', async ({ page }) => { + await setup(page) + await page.route('**/v1/audio/diarization', route => route.fulfill({ status: 501, json: { error: 'unsupported backend' } })) + await page.getByLabel('Prepare speakers to remember').check() + await page.getByRole('button', { name: 'Diarize', exact: true }).click() + await expect(page.getByRole('alert')).toContainText('Choose a profile-capable model') + await expect(page.getByLabel('Prepare speakers to remember')).toBeChecked() +}) + + +test('missing requested profiles is an error and normal defaults remain opt-in', async ({ page }) => { + await setup(page) + await expect(page.getByLabel('Prepare speakers to remember')).not.toBeChecked() + await page.route('**/v1/audio/diarization', route => route.fulfill({ json: { ...result, speaker_profiles: undefined } })) + await page.getByLabel('Prepare speakers to remember').check() + await page.getByRole('button', { name: 'Diarize', exact: true }).click() + await expect(page.getByRole('alert')).toContainText('did not return') + await expect(page.getByTestId('segments')).toHaveCount(0) +}) + +test('diarization permission gates direct page and Studio tab', async ({ page }) => { + await setup(page, true, false) + await expect(page).toHaveURL(/\/app$/) + await page.goto('/app/studio') + await expect(page.locator('[data-tab="diarization"]')).toHaveCount(0) +}) + +test('Studio exposes diarization and navigation revokes recording URL', async ({ page }) => { + await page.addInitScript(() => { + window.revoked = [] + const revoke = URL.revokeObjectURL.bind(URL) + URL.revokeObjectURL = url => { window.revoked.push(url); revoke(url) } + }) + await setup(page) + await page.goto('/app/studio/diarization') + await expect(page.getByRole('heading', { name: 'Speaker diarization' })).toBeVisible() + await page.getByLabel('Recording', { exact: true }).setInputFiles(file) + await run(page) + const url = await page.locator('audio').getAttribute('src') + await page.screenshot({ path: 'test-results/diarization-profiles.png', fullPage: true }) + await page.getByRole('link', { name: 'Manage remembered voices' }).click() + await expect.poll(() => page.evaluate(url => window.revoked.includes(url), url)).toBe(true) +}) diff --git a/core/http/react-ui/public/locales/en/media.json b/core/http/react-ui/public/locales/en/media.json index cf63f7389..7b2d826c0 100644 --- a/core/http/react-ui/public/locales/en/media.json +++ b/core/http/react-ui/public/locales/en/media.json @@ -7,7 +7,8 @@ "sound": "Sound", "transform": "Transform", "threed": "3D", - "overview": "Overview" + "overview": "Overview", + "diarization": "Diarization" }, "overview": { "eyebrow": "{{ready}} of {{total}} modalities ready", @@ -26,7 +27,8 @@ "threed": "Mesh generation and animation", "tts": "Text to speech using your voice library", "sound": "Music and sound effects from a prompt", - "transform": "Separation, enhancement and voice conversion" + "transform": "Separation, enhancement and voice conversion", + "diarization": "Find speaker turns and remember voices from a recording." } }, "groups": { @@ -477,5 +479,30 @@ "heading": "Request", "copyCurl": "Copy as curl", "copied": "Copied" + }, + "diarization": { + "title": "Speaker diarization", + "subtitle": "Upload a recording to see who spoke when. Preview clean speech before remembering a voice.", + "model": "Model", + "recording": "Recording", + "optIn": "Prepare speakers to remember", + "warning": "Remembered voices are shared globally on this server and are lost when it restarts. Nothing is remembered automatically.", + "manage": "Manage remembered voices", + "running": "Diarizing…", + "run": "Diarize", + "unsupported": "Choose a profile-capable model with a speaker encoder, or turn off “Prepare speakers to remember” to run normal diarization.", + "previewError": "Could not play this recording. Check that your browser supports its audio format.", + "speakers": "Speakers", + "segments": "Segments", + "duration": "Clean speech: {{seconds}} seconds", + "preview": "Preview {{number}}", + "insufficient": "Not enough clean speech or no usable profile. Try a recording with longer, non-overlapping speech.", + "nameAndRemember": "Name and remember", + "name": "Name", + "saving": "Saving…", + "remember": "Remember", + "cancel": "Cancel", + "stop": "Stop preview", + "missingProfiles": "The backend did not return the requested speaker profiles." } } diff --git a/core/http/react-ui/src/components/Modal.jsx b/core/http/react-ui/src/components/Modal.jsx index e13824ef6..b54ac1847 100644 --- a/core/http/react-ui/src/components/Modal.jsx +++ b/core/http/react-ui/src/components/Modal.jsx @@ -1,7 +1,7 @@ import { useEffect, useRef } from 'react' import '../pages/auth.css' -export default function Modal({ onClose, children, maxWidth = '600px' }) { +export default function Modal({ onClose, children, maxWidth = '600px', ariaLabel }) { const dialogRef = useRef(null) const lastFocusRef = useRef(null) const onCloseRef = useRef(onClose) @@ -54,6 +54,7 @@ export default function Modal({ onClose, children, maxWidth = '600px' }) {
diff --git a/core/http/react-ui/src/pages/Diarization.jsx b/core/http/react-ui/src/pages/Diarization.jsx new file mode 100644 index 000000000..a2efd9ea7 --- /dev/null +++ b/core/http/react-ui/src/pages/Diarization.jsx @@ -0,0 +1,173 @@ +// SPDX-License-Identifier: MIT +import { useEffect, useRef, useState } from 'react' +import { Link, useParams } from 'react-router-dom' +import { useTranslation } from 'react-i18next' +import PageHeader from '../components/PageHeader' +import ModelSelector from '../components/ModelSelector' +import Modal from '../components/Modal' +import { useAuth } from '../context/AuthContext' +import useObjectUrl from '../hooks/useObjectUrl' +import { CAP_DIARIZATION } from '../utils/capabilities' +import { diarizationApi, voiceApi } from '../utils/api' +import { rememberEnrollment } from '../utils/voiceEnrollments' + +export default function Diarization() { + const { t } = useTranslation('media') + const text = (key, values) => t(`diarization.${key}`, values) + const { model: initialModel } = useParams() + const { hasFeature } = useAuth() + const canRemember = hasFeature('voice_recognition') + const [model, setModel] = useState(initialModel || '') + const [file, setFile] = useState(null) + const [optIn, setOptIn] = useState(false) + const [result, setResult] = useState(null) + const [busy, setBusy] = useState(false) + const [error, setError] = useState('') + const [selected, setSelected] = useState(null) + const [name, setName] = useState('') + const [saving, setSaving] = useState(false) + const [saveError, setSaveError] = useState('') + const generation = useRef(0) + const saveLock = useRef(false) + const audio = useRef(null) + const end = useRef(null) + const timer = useRef(null) + const playback = useRef(0) + const url = useObjectUrl(file) + + function stop() { + playback.current++ + clearTimeout(timer.current) + audio.current?.pause() + end.current = null + } + function invalidate() { + generation.current++ + stop() + setResult(null); setSelected(null); setBusy(false); setSaving(false) + setError(''); setSaveError(''); saveLock.current = false + } + useEffect(() => { + const player = audio.current + return () => { playback.current++; clearTimeout(timer.current); player?.pause() } + }, [url]) + useEffect(() => () => { generation.current++ }, []) + useEffect(() => { + if (!canRemember) { setOptIn(false); setSelected(null) } + }, [canRemember]) + + async function submit(event) { + event.preventDefault() + if (!file || !model || busy) return + invalidate() + const token = generation.current + const requested = canRemember && optIn + setBusy(true) + try { + const data = await diarizationApi.run({ file, model, profiles: requested }) + if (token !== generation.current) return + if (requested && !data.speaker_profiles) throw new Error(text('missingProfiles')) + // Keep the actual inference model with this export, not a later selection. + setResult({ ...data, inferenceModel: model, speaker_profiles: requested ? data.speaker_profiles : undefined }) + } catch (err) { + if (token === generation.current) setError(`${err.message}${requested ? ` ${text('unsupported')}` : ''}`) + } finally { if (token === generation.current) setBusy(false) } + } + + async function preview(interval) { + stop() + const player = audio.current + if (!player) return + const token = generation.current + const playToken = playback.current + try { + player.currentTime = interval.start + end.current = interval.end + await player.play() + if (token !== generation.current || playToken !== playback.current) return + timer.current = setTimeout(stop, Math.max(0, interval.end - player.currentTime) * 1000) + } catch { if (token === generation.current) setError(text('previewError')) } + } + + async function save(event) { + event.preventDefault() + if (!canRemember || !name.trim() || selected === null || saveLock.current || !result) return + const token = generation.current + const slot = selected + saveLock.current = true; setSaving(true); setSaveError('') + try { + const registered = await voiceApi.register({ + model: result.inferenceModel, name: name.trim(), speaker_slot: slot, + speaker_profiles: result.speaker_profiles, + }) + // A completed registration is real even if its recording is no longer open. + rememberEnrollment(registered) + if (token !== generation.current) return + const relabel = rows => rows?.map(row => String(row.label) === String(slot) ? { ...row, name: registered.name } : row) + setResult(current => ({ ...current, speakers: relabel(current.speakers), segments: relabel(current.segments) })) + setSelected(null) + } catch (err) { if (token === generation.current) setSaveError(err.message) } + finally { if (token === generation.current) { saveLock.current = false; setSaving(false) } } + } + + // Raw labels are the only stable join key. Normalized IDs and array order can differ. + const summaries = result?.speakers || Array.from(new Map((result?.segments || []).map(s => [String(s.label), { ...s, id: s.speaker }])).values()) + return ( +
+ +
+
+ {text('model')} + { if (value !== model) { invalidate(); setModel(value) } }} /> +
+
+ + { invalidate(); setFile(e.target.files?.[0] || null) }} /> +
+ {canRemember &&
+ +

{text('warning')} {text('manage')}

+
} + +
+ {error &&

{error}

} + {url &&
+ ) +} diff --git a/core/http/react-ui/src/pages/Studio.jsx b/core/http/react-ui/src/pages/Studio.jsx index e2f96521e..26f8a7342 100644 --- a/core/http/react-ui/src/pages/Studio.jsx +++ b/core/http/react-ui/src/pages/Studio.jsx @@ -7,6 +7,7 @@ import ThreeDGen from './ThreeDGen' import TTS from './TTS' import Sound from './Sound' import AudioTransform from './AudioTransform' +import Diarization from './Diarization' import StudioOverview from './StudioOverview' import { useAuth } from '../context/AuthContext' import { useModels } from '../hooks/useModels' @@ -14,7 +15,7 @@ import { useOperations } from '../hooks/useOperations' import { readAllMediaHistory } from '../hooks/useMediaHistory' import { use3DHistory } from '../hooks/use3DHistory' import { - CAP_IMAGE, CAP_VIDEO, CAP_3D, CAP_3D_ANIMATION, CAP_TTS, CAP_SOUND_GENERATION, CAP_AUDIO_TRANSFORM, + CAP_DIARIZATION, CAP_IMAGE, CAP_VIDEO, CAP_3D, CAP_3D_ANIMATION, CAP_TTS, CAP_SOUND_GENERATION, CAP_AUDIO_TRANSFORM, } from '../utils/capabilities' // One table for the six generators: the capability that makes a modality @@ -22,6 +23,7 @@ import { // under. Studio owns this so the tab strip and the overview cannot disagree // about what exists. const MODALITIES = [ + { key: 'diarization', capability: CAP_DIARIZATION, icon: 'fas fa-users', group: 'voice', feature: 'audio_diarization' }, { key: 'images', capability: CAP_IMAGE, icon: 'fas fa-image', group: 'create', history: 'image' }, { key: 'video', capability: CAP_VIDEO, icon: 'fas fa-video', group: 'create', history: 'video' }, { key: 'threed', capability: CAP_3D, icon: 'fas fa-cube', group: 'create', feature: '3d' }, @@ -33,6 +35,7 @@ const MODALITIES = [ const OVERVIEW_TAB = { key: 'overview', icon: 'fas fa-compass' } const TAB_COMPONENTS = { + diarization: Diarization, images: ImageGen, video: VideoGen, threed: ThreeDGen, diff --git a/core/http/react-ui/src/pages/VoiceRecognition.jsx b/core/http/react-ui/src/pages/VoiceRecognition.jsx index 2dd295435..90e9eaedb 100644 --- a/core/http/react-ui/src/pages/VoiceRecognition.jsx +++ b/core/http/react-ui/src/pages/VoiceRecognition.jsx @@ -1,3 +1,4 @@ +import { loadEnrollments, saveEnrollments } from '../utils/voiceEnrollments' import { useEffect, useMemo, useState } from 'react' import { useOutletContext, useParams } from 'react-router-dom' import ModelSelector from '../components/ModelSelector' @@ -20,21 +21,6 @@ const TABS = [ { id: 'embed', icon: 'fas fa-code', label: 'Embedding' }, ] -const ENROLL_KEY = 'localai_voice_enrollments' - -function loadEnrollments() { - try { - const raw = localStorage.getItem(ENROLL_KEY) - if (!raw) return [] - const p = JSON.parse(raw) - return Array.isArray(p) ? p : [] - } catch (_) { return [] } -} - -function saveEnrollments(list) { - try { localStorage.setItem(ENROLL_KEY, JSON.stringify(list.slice(0, 50))) } catch (_) { /* quota */ } -} - function parseLabels(text) { const out = {} if (!text) return out diff --git a/core/http/react-ui/src/router.jsx b/core/http/react-ui/src/router.jsx index c871e3b73..f37fcd04b 100644 --- a/core/http/react-ui/src/router.jsx +++ b/core/http/react-ui/src/router.jsx @@ -85,6 +85,7 @@ const VideoGen = page('video', () => import('./pages/VideoGen')) const ThreeDGen = page('3d', () => import('./pages/ThreeDGen')) const TTS = page('tts', () => import('./pages/TTS')) const Sound = page('sound', () => import('./pages/Sound')) +const Diarization = page('diarization', () => import('./pages/Diarization')) const AudioTransform = page('transform', () => import('./pages/AudioTransform')) const Talk = page('talk', () => import('./pages/Talk')) // Referenced only from JSX below — same blind spot as Activity further down. @@ -165,6 +166,8 @@ const appChildren = [ { path: 'tts/:model', element: }, { path: 'sound', element: }, { path: 'sound/:model', element: }, + { path: 'diarization', element: }, + { path: 'diarization/:model', element: }, { path: 'transform', element: }, { path: 'transform/:model', element: }, { path: 'studio', element: }, diff --git a/core/http/react-ui/src/utils/api.js b/core/http/react-ui/src/utils/api.js index 8729876cb..eee0b2533 100644 --- a/core/http/react-ui/src/utils/api.js +++ b/core/http/react-ui/src/utils/api.js @@ -683,3 +683,18 @@ export function fileToBase64(file) { reader.readAsDataURL(file) }) } + +// Multipart requests must let the browser set the boundary. +export const diarizationApi = { + run: async ({ file, model, profiles = false }) => { + const body = new FormData() + body.append('file', file) + body.append('model', model) + if (profiles) { + body.append('include_speaker_profiles', 'true') + body.append('include_text', 'true') + body.append('response_format', 'verbose_json') + } + return handleResponse(await fetch(apiUrl('/v1/audio/diarization'), { method: 'POST', body })) + }, +} diff --git a/core/http/react-ui/src/utils/voiceEnrollments.js b/core/http/react-ui/src/utils/voiceEnrollments.js new file mode 100644 index 000000000..85ad65698 --- /dev/null +++ b/core/http/react-ui/src/utils/voiceEnrollments.js @@ -0,0 +1,20 @@ +// SPDX-License-Identifier: MIT +const ENROLL_KEY = 'localai_voice_enrollments' + +export function loadEnrollments() { + try { + const raw = localStorage.getItem(ENROLL_KEY) + if (!raw) return [] + const p = JSON.parse(raw) + return Array.isArray(p) ? p : [] + } catch (_) { return [] } +} + +export function saveEnrollments(list) { + try { localStorage.setItem(ENROLL_KEY, JSON.stringify(list.slice(0, 50))) } catch (_) { /* quota */ } +} + +// Only registration metadata belongs in the browser's management list. +export function rememberEnrollment({ id, name, registered_at }) { + saveEnrollments([{ id, name, registeredAt: registered_at, labels: {} }, ...loadEnrollments()]) +} diff --git a/core/http/routes/openai.go b/core/http/routes/openai.go index 604adacf7..eb752693d 100644 --- a/core/http/routes/openai.go +++ b/core/http/routes/openai.go @@ -194,7 +194,7 @@ func RegisterOpenAIRoutes(app *echo.Echo, app.POST("/v1/audio/transcriptions", audioHandler, audioMiddleware...) app.POST("/audio/transcriptions", audioHandler, audioMiddleware...) - diarizationHandler := openai.DiarizationEndpoint(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig(), application.VoiceRegistry()) + diarizationHandler := openai.DiarizationEndpoint(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig(), application.VoiceRegistry(), application.AuthDB()) diarizationMiddleware := []echo.MiddlewareFunc{ traceMiddleware, re.BuildFilteredFirstAvailableDefaultModel(config.BuildUsecaseFilterFn(config.FLAG_DIARIZATION)), diff --git a/core/schema/diarization.go b/core/schema/diarization.go index 8f8561ab3..5ad88c20e 100644 --- a/core/schema/diarization.go +++ b/core/schema/diarization.go @@ -34,12 +34,13 @@ type DiarizationSpeaker struct { // Speakers and segment text are omitted when empty so the default `json` // response stays minimal; verbose_json keeps both populated. type DiarizationResult struct { - Task string `json:"task"` - Duration float64 `json:"duration,omitempty"` - Language string `json:"language,omitempty"` - NumSpeakers int `json:"num_speakers"` - Segments []DiarizationSegment `json:"segments"` - Speakers []DiarizationSpeaker `json:"speakers,omitempty"` + SpeakerProfiles *SpeakerProfiles `json:"speaker_profiles,omitempty"` + Task string `json:"task"` + Duration float64 `json:"duration,omitempty"` + Language string `json:"language,omitempty"` + NumSpeakers int `json:"num_speakers"` + Segments []DiarizationSegment `json:"segments"` + Speakers []DiarizationSpeaker `json:"speakers,omitempty"` } // DiarizationResponseFormatType mirrors transcription's response_format diff --git a/core/schema/localai.go b/core/schema/localai.go index dc99a1dbe..44a859cba 100644 --- a/core/schema/localai.go +++ b/core/schema/localai.go @@ -478,6 +478,8 @@ type VoiceEmbedResponse struct { // VoiceRegisterRequest enrolls a speaker into the 1:N identification store. type VoiceRegisterRequest struct { + SpeakerProfiles *SpeakerProfiles `json:"speaker_profiles,omitempty"` + SpeakerSlot *int `json:"speaker_slot,omitempty"` BasicModelRequest Audio string `json:"audio"` Name string `json:"name"` diff --git a/core/schema/openai.go b/core/schema/openai.go index 6f3717256..4646c94c5 100644 --- a/core/schema/openai.go +++ b/core/schema/openai.go @@ -187,6 +187,8 @@ type JsonSchema struct { } type OpenAIRequest struct { + IncludeSpeakerProfiles bool `json:"include_speaker_profiles,omitempty"` + IncludeText bool `json:"include_text,omitempty"` PredictionOptions Context context.Context `json:"-"` diff --git a/core/schema/speaker_profiles.go b/core/schema/speaker_profiles.go new file mode 100644 index 000000000..20d77aa41 --- /dev/null +++ b/core/schema/speaker_profiles.go @@ -0,0 +1,132 @@ +// SPDX-License-Identifier: MIT + +package schema + +import ( + "fmt" + "math" + "regexp" + "strings" +) + +// SpeakerEncoder identifies the exact encoder weights, not a model filename. +type SpeakerEncoder struct { + Identity string `json:"identity"` + Dimension int `json:"dimension"` +} + +// SpeakerProfileInterval locates retained clean audio in the original recording, +// in seconds. It does not describe separated or synthesized audio. +type SpeakerProfileInterval struct { + Start float64 `json:"start"` + End float64 `json:"end"` +} + +// SpeakerProfile contains one sensitive voice vector per discovered speaker. +type SpeakerProfile struct { + Speaker int `json:"speaker"` + CleanDuration float64 `json:"clean_duration"` + Intervals []SpeakerProfileInterval `json:"intervals"` + UnavailableReason *string `json:"unavailable_reason"` + Embedding []float32 `json:"embedding,omitempty"` +} + +// SpeakerProfiles is the versioned speaker_profiles object exported by parakeet. +// These unsigned profiles are not proof of identity or consent to enrollment. +type SpeakerProfiles struct { + Version int `json:"version"` + Encoder SpeakerEncoder `json:"encoder"` + Speakers []SpeakerProfile `json:"speakers"` +} + +var speakerEncoderIdentity = regexp.MustCompile(`^sha256:[0-9a-f]{64}$`) + +// Validate checks portable data against metadata from the server's loaded encoder. +// trusted must never come from the request itself. This does not authorize export +// or enrollment, verify provenance, or check intervals against recording length. +func (p SpeakerProfiles) Validate(trusted SpeakerEncoder) error { + if p.Version != 1 { + return fmt.Errorf("unsupported speaker profile version: %d", p.Version) + } + if !speakerEncoderIdentity.MatchString(trusted.Identity) || trusted.Dimension <= 0 { + return fmt.Errorf("invalid trusted speaker encoder metadata") + } + if p.Encoder != trusted { + return fmt.Errorf("speaker profile encoder does not match loaded encoder") + } + seen := make(map[int]bool, len(p.Speakers)) + for _, s := range p.Speakers { + if s.Speaker < 0 || seen[s.Speaker] { + return fmt.Errorf("invalid or duplicate speaker slot: %d", s.Speaker) + } + seen[s.Speaker] = true + if err := s.validate(trusted.Dimension); err != nil { + return fmt.Errorf("speaker %d: %w", s.Speaker, err) + } + } + return nil +} + +// Select validates the complete export and returns a usable speaker for explicit +// enrollment. Callers must supply trusted loaded-encoder metadata, not p.Encoder. +func (p SpeakerProfiles) Select(speaker int, trusted SpeakerEncoder) (SpeakerProfile, error) { + if err := p.Validate(trusted); err != nil { + return SpeakerProfile{}, err + } + for _, s := range p.Speakers { + if s.Speaker == speaker { + if s.UnavailableReason != nil { + return SpeakerProfile{}, fmt.Errorf("speaker %d is unavailable", speaker) + } + return s, nil + } + } + return SpeakerProfile{}, fmt.Errorf("speaker %d not found", speaker) +} + +func (s SpeakerProfile) validate(dimension int) error { + // Native JSON rounds timestamps; allow a millisecond of serialization drift. + const tolerance = 0.001 + if !finiteProfileNumber(s.CleanDuration) || s.CleanDuration < 0 || s.CleanDuration > 30+tolerance { + return fmt.Errorf("invalid clean duration") + } + var duration, previousEnd float64 + for _, interval := range s.Intervals { + if !finiteProfileNumber(interval.Start) || !finiteProfileNumber(interval.End) || interval.Start < previousEnd || interval.End <= interval.Start { + return fmt.Errorf("invalid clean interval") + } + duration += interval.End - interval.Start + previousEnd = interval.End + } + if !finiteProfileNumber(duration) || math.Abs(duration-s.CleanDuration) > tolerance { + return fmt.Errorf("clean duration does not match intervals") + } + if s.UnavailableReason != nil { + if strings.TrimSpace(*s.UnavailableReason) == "" || len(s.Embedding) != 0 { + return fmt.Errorf("invalid unavailable profile") + } + return nil + } + if s.CleanDuration < 2-tolerance { + return fmt.Errorf("insufficient clean speech") + } + if len(s.Embedding) != dimension { + return fmt.Errorf("speaker embedding dimension mismatch") + } + var norm float64 + for _, value := range s.Embedding { + v := float64(value) + if !finiteProfileNumber(v) { + return fmt.Errorf("speaker embedding must be finite") + } + norm += v * v + } + if norm == 0 { + return fmt.Errorf("speaker embedding must be nonzero") + } + return nil +} + +func finiteProfileNumber(v float64) bool { + return !math.IsNaN(v) && !math.IsInf(v, 0) +} diff --git a/core/schema/speaker_profiles_test.go b/core/schema/speaker_profiles_test.go new file mode 100644 index 000000000..4334af847 --- /dev/null +++ b/core/schema/speaker_profiles_test.go @@ -0,0 +1,113 @@ +// SPDX-License-Identifier: MIT + +package schema + +import ( + "encoding/json" + "math" + "strings" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Portable speaker profiles", func() { + var p SpeakerProfiles + var trusted SpeakerEncoder + BeforeEach(func() { + trusted = SpeakerEncoder{Identity: "sha256:" + strings.Repeat("a", 64), Dimension: 2} + p = SpeakerProfiles{Version: 1, Encoder: trusted, Speakers: []SpeakerProfile{{Speaker: 3, CleanDuration: 3, Intervals: []SpeakerProfileInterval{{Start: 0, End: 3}}, Embedding: []float32{0.6, 0.8}}}} + }) + It("selects the requested slot against independently supplied metadata", func() { + selected, err := p.Select(3, trusted) + Expect(err).NotTo(HaveOccurred()) + Expect(selected).To(Equal(p.Speakers[0])) + _, err = p.Select(0, trusted) + Expect(err).To(HaveOccurred()) + p.Encoder.Identity = "sha256:" + strings.Repeat("b", 64) + _, err = p.Select(3, trusted) + Expect(err).To(HaveOccurred()) + p.Encoder = trusted + p.Encoder.Dimension = 3 + Expect(p.Validate(trusted)).To(HaveOccurred()) + }) + It("round trips the backend JSON including unavailable profiles without embeddings", func() { + raw := `{"version":1,"encoder":{"identity":"` + trusted.Identity + `","dimension":2},"speakers":[{"speaker":3,"clean_duration":3,"intervals":[{"start":0,"end":3}],"unavailable_reason":null,"embedding":[0.6,0.8]},{"speaker":4,"clean_duration":1,"intervals":[{"start":4,"end":5}],"unavailable_reason":"insufficient_clean_speech"}]}` + Expect(json.Unmarshal([]byte(raw), &p)).To(Succeed()) + Expect(p.Validate(trusted)).To(Succeed()) + encoded, err := json.Marshal(p) + Expect(err).NotTo(HaveOccurred()) + Expect(encoded).To(MatchJSON(raw)) + _, err = p.Select(4, trusted) + Expect(err).To(HaveOccurred()) + _, err = p.Select(3, trusted) + Expect(err).NotTo(HaveOccurred()) + }) + It("allows empty discovery but cannot select from it", func() { + p.Speakers = []SpeakerProfile{} + Expect(p.Validate(trusted)).To(Succeed()) + _, err := p.Select(0, trusted) + Expect(err).To(HaveOccurred()) + }) + It("rejects unsupported versions and invalid trusted metadata", func() { + p.Version = 2 + Expect(p.Validate(trusted)).To(HaveOccurred()) + p.Version = 1 + for _, encoder := range []SpeakerEncoder{{}, {Identity: "sha256:trusted", Dimension: 2}, {Identity: trusted.Identity, Dimension: 0}, {Identity: strings.ToUpper(trusted.Identity), Dimension: 2}} { + p.Encoder = encoder + Expect(p.Validate(encoder)).To(HaveOccurred()) + } + }) + DescribeTable("rejects invalid vectors", func(vector []float32) { + p.Speakers[0].Embedding = vector + _, err := p.Select(3, trusted) + Expect(err).To(HaveOccurred()) + }, + Entry("missing", []float32(nil)), Entry("zero", []float32{0, 0}), + Entry("wrong dimension", []float32{1}), Entry("NaN", []float32{float32(math.NaN()), 1}), + Entry("positive infinity", []float32{float32(math.Inf(1)), 1}), Entry("negative infinity", []float32{1, float32(math.Inf(-1))}), + ) + It("accepts finite nonzero vectors without imposing a second normalization policy", func() { + p.Speakers[0].Embedding = []float32{math.MaxFloat32, math.SmallestNonzeroFloat32} + Expect(p.Validate(trusted)).To(Succeed()) + }) + It("rejects duplicate or negative speaker slots", func() { + p.Speakers = append(p.Speakers, p.Speakers[0]) + Expect(p.Validate(trusted)).To(HaveOccurred()) + p.Speakers = p.Speakers[:1] + p.Speakers[0].Speaker = -1 + Expect(p.Validate(trusted)).To(HaveOccurred()) + }) + It("rejects inconsistent unavailable status", func() { + reason := "embedding_failed" + p.Speakers[0].UnavailableReason = &reason + Expect(p.Validate(trusted)).To(HaveOccurred()) + p.Speakers[0].Embedding = nil + Expect(p.Validate(trusted)).To(Succeed()) + reason = " " + Expect(p.Validate(trusted)).To(HaveOccurred()) + }) + DescribeTable("rejects unreasonable duration or intervals", func(duration float64, intervals []SpeakerProfileInterval) { + p.Speakers[0].CleanDuration = duration + p.Speakers[0].Intervals = intervals + Expect(p.Validate(trusted)).To(HaveOccurred()) + }, + Entry("negative", -1.0, []SpeakerProfileInterval(nil)), + Entry("nonfinite duration", math.NaN(), []SpeakerProfileInterval(nil)), + Entry("too long", 31.0, []SpeakerProfileInterval{{0, 31}}), + Entry("too short for enrollment", 1.0, []SpeakerProfileInterval{{0, 1}}), + Entry("mismatch", 3.0, []SpeakerProfileInterval{{0, 2}}), + Entry("missing spans", 3.0, []SpeakerProfileInterval(nil)), + Entry("negative start", 3.0, []SpeakerProfileInterval{{-1, 2}}), + Entry("reversed", 3.0, []SpeakerProfileInterval{{3, 0}}), + Entry("empty span", 3.0, []SpeakerProfileInterval{{0, 0}, {0, 3}}), + Entry("overlap", 3.0, []SpeakerProfileInterval{{0, 2}, {1, 2}}), + Entry("out of order", 3.0, []SpeakerProfileInterval{{2, 4}, {0, 1}}), + Entry("infinite end", 3.0, []SpeakerProfileInterval{{0, math.Inf(1)}}), + Entry("NaN start", 3.0, []SpeakerProfileInterval{{math.NaN(), 3}}), + ) + It("accepts disjoint original spans and serialization rounding", func() { + p.Speakers[0].Intervals = []SpeakerProfileInterval{{1, 2}, {5, 7.0001}} + Expect(p.Validate(trusted)).To(Succeed()) + }) +}) diff --git a/core/services/voicerecognition/known_voice_ids_test.go b/core/services/voicerecognition/known_voice_ids_test.go new file mode 100644 index 000000000..d6e684151 --- /dev/null +++ b/core/services/voicerecognition/known_voice_ids_test.go @@ -0,0 +1,21 @@ +// SPDX-License-Identifier: MIT +package voicerecognition + +import ( + ginkgo "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = ginkgo.Describe("known voice registration IDs", func() { + ginkgo.It("retains independent registrations sharing a display name", func() { + selected := SelectKnownVoices([]Entry{ + {Metadata: Metadata{ID: "a", Name: "Ada", Model: "speaker.gguf"}, Embedding: []float32{1, 0}}, + {Metadata: Metadata{ID: "b", Name: "Ada", Model: "speaker.gguf"}, Embedding: []float32{0, 1}}, + }, "speaker.gguf") + Expect(selected.Voices).To(HaveLen(2)) + Expect(selected.Voices[0].ID).To(Equal("a")) + Expect(selected.Voices[1].ID).To(Equal("b")) + Expect(selected.Voices[0].Embedding).To(Equal([]float32{1, 0})) + Expect(selected.Voices[1].Embedding).To(Equal([]float32{0, 1})) + }) +}) diff --git a/core/services/voicerecognition/known_voices.go b/core/services/voicerecognition/known_voices.go index 823f4b150..0f4cf8fba 100644 --- a/core/services/voicerecognition/known_voices.go +++ b/core/services/voicerecognition/known_voices.go @@ -4,12 +4,14 @@ import ( "context" "path" "path/filepath" + "sort" "strings" ) // KnownVoice is one registered voice as it is sent to a backend that matches // speakers itself. type KnownVoice struct { + ID string Name string Embedding []float32 Model string @@ -41,16 +43,14 @@ type KnownVoiceSelection struct { Untagged int // voices with no encoder tag that were included } -// SelectKnownVoices picks, from everything registered, the voices that can be -// compared with embeddings from the speaker model at speakerModelPath: those -// with the same encoder tag, then the untagged voices whose size matches the -// tagged ones (all untagged voices when none matched). Voices from another -// encoder are counted and skipped, as are voices without a name or embedding. -// The input is not modified. +// SelectKnownVoices selects filename-matching, portable and untagged candidates. +// Registry tags and vector lengths are not trusted encoder metadata: dimension +// filtering belongs to the loaded backend. Tagged candidates precede untagged +// ones, each ordered by registration ID so registry iteration order cannot +// change replay order. The input is not modified. func SelectKnownVoices(entries []Entry, speakerModelPath string) KnownVoiceSelection { tag := EncoderTag(speakerModelPath) var sel KnownVoiceSelection - matchedDim := 0 var untagged []Entry for _, e := range entries { if e.Metadata.Name == "" || len(e.Embedding) == 0 { @@ -59,23 +59,19 @@ func SelectKnownVoices(entries []Entry, speakerModelPath string) KnownVoiceSelec switch { case e.Metadata.Model == "": untagged = append(untagged, e) - case EncoderTag(e.Metadata.Model) == tag: - if matchedDim == 0 { - matchedDim = len(e.Embedding) - } - sel.Voices = append(sel.Voices, KnownVoice{Name: e.Metadata.Name, Embedding: e.Embedding, Model: e.Metadata.Model}) + // Hash-tagged portable registrations are checked against the loaded + // encoder by the backend, never against a filename or dimension alone. + case strings.HasPrefix(e.Metadata.Model, "sha256:"), EncoderTag(e.Metadata.Model) == tag: + sel.Voices = append(sel.Voices, KnownVoice{ID: e.Metadata.ID, Name: e.Metadata.Name, Embedding: e.Embedding, Model: e.Metadata.Model}) default: sel.OtherEncoder++ } } - // With no tagged match every untagged voice is included whatever its size; the - // backend skips the voices whose size differs from its speaker model's. + sort.Slice(sel.Voices, func(i, j int) bool { return sel.Voices[i].ID < sel.Voices[j].ID }) + sort.Slice(untagged, func(i, j int) bool { return untagged[i].Metadata.ID < untagged[j].Metadata.ID }) for _, e := range untagged { - if matchedDim != 0 && len(e.Embedding) != matchedDim { - continue - } sel.Untagged++ - sel.Voices = append(sel.Voices, KnownVoice{Name: e.Metadata.Name, Embedding: e.Embedding}) + sel.Voices = append(sel.Voices, KnownVoice{ID: e.Metadata.ID, Name: e.Metadata.Name, Embedding: e.Embedding}) } return sel } diff --git a/core/services/voicerecognition/known_voices_test.go b/core/services/voicerecognition/known_voices_test.go index 8d2491ea2..fbe5510cf 100644 --- a/core/services/voicerecognition/known_voices_test.go +++ b/core/services/voicerecognition/known_voices_test.go @@ -56,15 +56,16 @@ var _ = Describe("SelectKnownVoices", func() { Expect(sel.Voices).To(BeEmpty()) Expect(sel.OtherEncoder).To(Equal(1)) }) - It("includes an untagged voice only when its size matches the tagged voices", func() { + It("defers untagged dimensions to the loaded backend, not registry tags", func() { sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{ entry("ada", wespeaker, 1, 0), entry("old_same", "", 0, 1), entry("old_other", "", 0, 1, 0), }, wespeaker) - Expect(sel.Voices).To(HaveLen(2)) - Expect(sel.Voices[1].Name).To(Equal("old_same")) - Expect(sel.Untagged).To(Equal(1)) + Expect(sel.Voices).To(HaveLen(3)) + Expect(sel.Voices[1].Name).To(Equal("old_other")) + Expect(sel.Voices[2].Name).To(Equal("old_same")) + Expect(sel.Untagged).To(Equal(2)) }) It("puts tagged voices first even when an untagged one was registered earlier", func() { sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{ diff --git a/docs/content/features/audio-diarization.md b/docs/content/features/audio-diarization.md index e052f0b65..53e5e487f 100644 --- a/docs/content/features/audio-diarization.md +++ b/docs/content/features/audio-diarization.md @@ -180,7 +180,19 @@ curl http://localhost:8080/v1/audio/diarization \ ## Backend setup - parakeet-cpp (Nemotron-3-Diarization) -Nemotron-3-Diarization is Sortformer, served standalone or paired with a Parakeet ASR model. Install `parakeet-cpp-nemotron-3-diarization` from the gallery for diarization only, or `parakeet-cpp-nemotron-3-diarization-asr` for the same model paired with `parakeet-cpp-tdt_ctc-110m` through the `asr_model` option: +Choose an existing gallery entry for the output you need: + +| Output | Gallery entry | Request options | +|---|---|---| +| Speaker turns only | `parakeet-cpp-nemotron-3-diarization` | Default options | +| Speaker turns and transcript | `parakeet-cpp-nemotron-3-diarization-asr` | `include_text=true`, `response_format=verbose_json` | +| Speaker turns, transcript, and identification | `parakeet-cpp-nemotron-3-diarization-asr-speakers` | Same transcript options; explicitly enroll voices for names | + +The complete `-asr-speakers` entry downloads Nemotron-3-Diarization, Parakeet TDT+CTC 110M ASR, and the WeSpeaker ResNet34 speaker encoder. +It configures both `asr_model` and `speaker_model`; no custom gallery configuration is needed. +See [Remember speakers in the Web UI](#remember-speakers-in-the-web-ui) for installation and enrollment. + +For manual configuration, this example pairs Sortformer with ASR: ```yaml name: parakeet-diarize @@ -215,3 +227,241 @@ Sortformer clusters on voice-like characteristics, not on "is this a human". A l ## See also - [Sound Classification]({{% relref "audio-classification" %}}) - tag non-speech sound events (alarms, glass breaking, baby cry) in a clip. + +### Backend profile transport + +The parakeet backend supports opt-in speaker profile export through the internal +`DiarizeRequest.include_speaker_profiles` field. This native transport underpins +HTTP profile export and explicit enrollment through `POST /v1/voice/register`, +as described in [Portable speaker enrollment](#portable-speaker-enrollment) below. +It requires a configured `speaker_model` and a library +with `parakeet_capi_diarize_profiles_pcm_json`; an empty recognition registry +is supported. Export does not register anyone. With `include_text` and a loaded +ASR companion, one profile-capable diarization supplies all speaker slots, +profiles, names, and intervals. Timestamped ASR words are assigned to those +same slots; the backend does not run a second diarization. Either inference +failure fails the request. If no ASR companion is loaded, the existing fallback +applies: the response includes profiles and diarization segments without text. +A loaded ASR companion without the timestamped PCM API returns an explicit error. + +`DiarizeResponse.speaker_profiles_json` carries the native version-1 +`speaker_profiles` object, including original clean preview intervals and one +embedding per usable speaker. Normal requests retain their existing output. +Profile `speaker` values are raw native slot IDs. Match their decimal string to +segment `label` or speaker-summary `label`, not to normalized `SPEAKER_NN`, +array position, or display name. Slots can be sparse, and profile order can +differ from transcript order. Profiles retain their original clean intervals +even when transcript segments use word boundaries or duration filters. +These vectors are sensitive biometric data: callers must authorize export and +explicit enrollment separately. + +The internal backend Status response supplies `speaker_encoder`, derived from +the loaded encoder's SHA-256 identity and dimension. Enrollment code must use +`backend.ModelSpeakerEncoder` with server-selected model configuration and +validate profiles against that result, never against caller-provided metadata. +Unavailable metadata or unsupported export fails closed. Renaming a GGUF does +not change its identity; modifying or quantizing its bytes does. + +Recognition replay carries registration IDs separately from display names. +Distinct IDs with the same display name remain independent native entries, +and both offline and realtime matches are translated back to display names. +Legacy transport clients without IDs retain name-keyed behavior. The native +registry's aggregation defaults are unchanged. LocalAI's recognition registry +remains global and in-memory; this adds neither persistence nor automatic +registration and is unrelated to persistent TTS voice cloning. + +## Portable speaker enrollment + +Profile-capable parakeet models can export one biometric embedding per discovered +speaker, including when the recognition registry is empty. Export is opt-in: + +```bash +curl http://localhost:8080/v1/audio/diarization \ + -F model=parakeet-diarization -F file=@conversation.wav \ + -F include_speaker_profiles=true -F include_text=true \ + -F response_format=verbose_json +``` + +The `/audio/diarization` alias has the same protection. With user authentication, +export additionally requires the **voice-recognition** permission. Existing model +access controls still apply. Without opt-in, `speaker_profiles` is omitted. +Both `json` and `verbose_json` support profiles; `rttm` with profiles returns 400. +`include_text=true` retains supported transcripts in either JSON format. +Unsupported profile backends return 501 rather than silently omitting profiles. + +Alternatively send `Content-Type: application/json`: + +```json +{ + "model": "parakeet-diarization", + "file": "", + "include_speaker_profiles": true, + "include_text": true, + "response_format": "verbose_json" +} +``` + +The `speaker_profiles` response object contains `version: 1`, +`encoder: {"identity": "sha256:<64 lowercase hex digits>", "dimension": N}`, +and `speakers`. Each speaker contains: + +- `speaker`: the raw numeric speaker slot; +- `clean_duration`: retained clean speech in seconds; +- `intervals`: `{start, end}` ranges in seconds in the original recording; +- `unavailable_reason`: null for usable profiles, otherwise a reason string; +- `embedding`: one vector for a usable speaker, omitted when unavailable. + +**UI association:** convert each profile's numeric `speaker` to a decimal string +and match segment/summary `label`. Do not use `SPEAKER_NN`, array position, or +human name. Slots may be sparse and out of order; display names may repeat. +Preview `intervals` against the original audio, not separated audio. Disable +saving unavailable profiles. Enrollment is explicit, never automatic; only +relabel after a successful registration response. See +[portable voice registration](/features/voice-recognition/#portable-profile-registration). + +Profiles are sensitive, unsigned biometric data, not proof of identity or consent. +Do not log their vectors. Obtain the speaker's consent before enrollment. + +API tracing excludes the entire exchange for `/v1/audio/diarization`, its +`/audio/diarization` alias, and `/v1/voice/register` before capturing bodies. +This also protects JSON base64 audio when profile export is off. These routes +produce no in-memory or persisted API trace; other routes keep their existing +tracing behavior. External proxies and client logs must apply the same privacy +policy. Existing trace files from older versions are not retroactively scrubbed. + +## Remember speakers in the Web UI + +Use a LocalAI build with portable enrollment support and a profile-capable `parakeet-cpp` backend. +The backend needs the profile APIs from merged upstream commit +[`bee7c14`](https://github.com/mudler/parakeet.cpp/commit/bee7c14dfcc23613df58176c59a40459e7b47095) or a compatible later build. +Installing the model weights alone does not update an older backend. + +1. Open **Models → Explore** and search for `parakeet-cpp-nemotron-3-diarization-asr-speakers`. +2. Select **Install** and wait for installation to complete. Check **Operate → Activity** for progress or errors. +3. Open **Studio → Diarization** (or `/app/diarization`). Select that model and upload your recording. + +Obtain the speaker's consent before enrollment. To remember a speaker from that recording: + +1. Select **Prepare speakers to remember**, then select **Diarize**. This + requests profiles, transcript text, and speaker summaries. Use a + profile-capable parakeet-cpp model configured with a speaker encoder. +2. In **Speakers**, select **Preview 1**, **Preview 2**, or another available + interval to listen to clean speech from the original recording. Playback + stops at the end of that interval. **Stop preview** stops it earlier. + Your browser must support the recording's audio format. +3. For an unknown speaker, select **Name and remember**. Enter a name and + select **Remember**. No second recording or audio upload is needed. +4. After the server confirms registration, the name appears on all turns for + that speaker. A failed save keeps the entered name so you can retry. + +Upload another recording and select **Diarize** to match remembered voices. +You can turn off **Prepare speakers to remember**; recognition does not require another profile export. +Matches show their names; unmatched speakers keep their speaker labels. +With preparation off, the UI requests speaker turns without transcript text. +Use the API example below to request text without exporting profiles. + +Speakers without a usable profile cannot be +remembered; try longer speech without overlapping speakers. Duplicate names +are allowed: each save creates a separate registration, not a merged voice. +Changing the model or recording clears the current results and save dialog. +A save already sent to the server can still complete, but cannot rename turns +in a different recording. + +The page requires the **Audio Diarization** permission and access to the selected model. Preparing profiles and +remembering speakers additionally require **Voice Recognition**. Users without +that permission can still run normal diarization. If the backend does not +support profiles, the page reports an error: choose a compatible model or +turn off **Prepare speakers to remember**. It does not silently retry without +profiles. + +{{% notice warning %}} +Remembered voices are shared globally on this server and are lost when it +restarts. Nothing is enrolled automatically. The browser stores only the new +registration's ID, name, and registration time for the existing voice +management list, not its embedding or recording. That list is local to the +browser and is not a durable server registry. +{{% /notice %}} + +Use **Manage remembered voices**, then the **Enrollment** tab, to see or +remove registrations saved in this browser. Clean-clip voice enrollment stays +available there and does not require diarization. + +### API example: install, export, and remember + +This example uses the same complete gallery entry and requires `curl` and `jq`. +The commands assume a local server without authentication. +If authentication is enabled, add `-H "Authorization: Bearer "` to every request using your authorized key. +Keep keys out of shared scripts, logs, and shell history; see [Authentication]({{% relref "authentication" %}}). +Installation requires model-management access; inference and enrollment require the permissions described above. + +Install the model if it is not already installed: + +```bash +LOCALAI=http://localhost:8080 +MODEL=parakeet-cpp-nemotron-3-diarization-asr-speakers +curl --fail-with-body "$LOCALAI/models/apply" \ + -H 'Content-Type: application/json' \ + -d '{"id":"localai@parakeet-cpp-nemotron-3-diarization-asr-speakers"}' +``` + +Installation is asynchronous. Wait for successful completion in **Operate → Activity** before continuing. +API clients can query the returned job `status` URL; see the [model gallery API]({{% relref "model-gallery" %}}). + +{{% notice warning %}} +Exported profiles contain biometric vectors. Obtain consent before enrollment. +Keep the recording, response, and registration files private. Do not log or share their contents. +Use a new private directory so existing files cannot retain broader permissions. Delete these files when no longer needed. +{{% /notice %}} + +Export profiles and transcript text from your recording, keeping the complete JSON response: + +```bash +umask 077 +WORK=$(mktemp -d) +curl --fail-with-body "$LOCALAI/v1/audio/diarization" \ + -F "model=$MODEL" -F file=@conversation.wav \ + -F include_text=true -F include_speaker_profiles=true \ + -F response_format=verbose_json > "$WORK/diarization.json" + +# Inspect raw slots, clean intervals, and transcript labels without printing vectors. +jq '.speaker_profiles.speakers[] | {speaker, clean_duration, intervals, unavailable_reason}' \ + "$WORK/diarization.json" +jq '.segments[] | {label, start, end, text}' "$WORK/diarization.json" +``` + +Choose a usable raw `speaker` slot whose decimal string matches the intended segment `label`. +Listen to its `intervals` in the original recording before assigning a name. +Do not select by array position, normalized `SPEAKER_NN`, or display name. +If `unavailable_reason` indicates insufficient speech, try a longer recording without overlapping speakers. + +Replace `0` below with your chosen raw slot. Zero is valid, but does not mean “the first array element.” +Keep the complete `speaker_profiles` object unchanged: + +```bash +SLOT=0 +NAME=Ada +jq --arg model "$MODEL" --arg name "$NAME" --argjson slot "$SLOT" \ + '{model: $model, name: $name, speaker_slot: $slot, speaker_profiles: .speaker_profiles}' \ + "$WORK/diarization.json" > "$WORK/register.json" +curl --fail-with-body "$LOCALAI/v1/voice/register" \ + -H 'Content-Type: application/json' \ + --data-binary @"$WORK/register.json" +``` + +After successful registration, submit another recording with the same model: + +```bash +curl --fail-with-body "$LOCALAI/v1/audio/diarization" \ + -F "model=$MODEL" -F file=@next-conversation.wav \ + -F include_text=true -F response_format=verbose_json > "$WORK/next.json" +jq '.segments[] | {label, name, start, end, text}' "$WORK/next.json" + +# Remove private example outputs when no longer needed. +rm -f "$WORK/diarization.json" "$WORK/register.json" "$WORK/next.json" +rmdir "$WORK" +``` + +Matching speakers can now carry `name`, even though this request omits `include_speaker_profiles`. +Keep `include_text=true` and `verbose_json` when you want transcript text. +Recognition is not proof of identity. Registrations remain global and disappear on server restart. +See [portable profile registration](/features/voice-recognition/#portable-profile-registration) for encoder compatibility and validation rules. diff --git a/docs/content/features/voice-recognition.md b/docs/content/features/voice-recognition.md index 0c2938a1c..d3c866ba0 100644 --- a/docs/content/features/voice-recognition.md +++ b/docs/content/features/voice-recognition.md @@ -402,3 +402,68 @@ default only applies when omitted. both the face and voice 1:N recognition pipelines. - [Embeddings](/features/embeddings/) - text-only OpenAI-compatible embedding endpoint; for audio embeddings use `/v1/voice/embed`. + +## Portable profile registration + +`POST /v1/voice/register` also accepts a JSON alternative to `audio`: + +```javascript +// result is the parsed diarization response; slot is a selected raw speaker slot. +const request = { + model: "parakeet-diarization", + name: "Ada", + labels: {team: "research"}, + speaker_slot: slot, + speaker_profiles: result.speaker_profiles +}; +// POST JSON.stringify(request) with Content-Type: application/json. +``` + +Copy the complete `speaker_profiles` object returned by diarization unchanged. +Select `speaker_slot` explicitly, including for slot zero. It is the raw numeric +slot whose decimal string matches the diarization `label`, not a normalized +`SPEAKER_NN`, array index, or display name. `audio` and `speaker_profiles` are +mutually exclusive. `speaker_slot` without profiles is also invalid. Audio-only +registration keeps its existing JSON shape and behavior. + +The server loads the requested, authorized model and obtains encoder identity and +dimension from backend metadata. It validates the complete profile export and +selects the requested usable slot. Missing slots, unavailable speech, unsupported +versions, non-finite/zero/wrong-size vectors and encoder mismatch return 400. +A backend without trusted encoder metadata returns 501. Success returns the +existing `{id, name, registered_at}` response. + +Portable registrations store the **server-derived SHA-256 identity**, not a +caller-provided filename tag. Offline/live recognition admits these registrations +only when the loaded encoder has the same identity and dimension. Legacy audio +registrations retain their filename-tag compatibility rules. `/v1/voice/identify` +filters incompatible matches; a backend unable to report trusted identity cannot +match portable registrations, even when vector dimensions agree. Filtering can +return fewer than `top_k` results. The parakeet diarization model need not support +the separate audio-only VoiceEmbed RPC used by `/v1/voice/identify`. + +Each successful enrollment inserts a new registration with its own ID and vector. +Duplicate display names do not merge embeddings or update an earlier enrollment. +There is no automatic enrollment or sample aggregation. + +The recognition registry is **global, in-memory and per LocalAI instance**; +registrations are lost on restart and are not synchronized across frontends. +This is not durable “remembering” and not a per-user private address book. The +persistent `/api/voice-profiles` TTS-cloning feature is unrelated. Export and +registration use the existing voice-recognition permission, with existing model +access restrictions; permission does not establish biometric consent. + +API tracing excludes the entire exchange for `/v1/audio/diarization`, its +`/audio/diarization` alias, and `/v1/voice/register` before capturing bodies. +This also protects JSON base64 audio when profile export is off. These routes +produce no in-memory or persisted API trace; other routes keep their existing +tracing behavior. External proxies and client logs must apply the same privacy +policy. Existing trace files from older versions are not retroactively scrubbed. + +For offline and live diarization replay, registry tags never determine the +encoder dimension. LocalAI orders candidates by registration ID (tagged first), +then uses loaded encoder metadata to filter dimensions. Portable registrations +require an exact SHA-256 identity match as well. Older backends without trusted +metadata reject portable candidates and retain their native legacy dimension +checks. Identification filters compatibility after the store's `top_k` query; +incompatible results can crowd out compatible candidates within that window. diff --git a/swagger/docs.go b/swagger/docs.go index 628a57796..f155d5283 100644 --- a/swagger/docs.go +++ b/swagger/docs.go @@ -2910,7 +2910,8 @@ const docTemplate = `{ "/v1/audio/diarization": { "post": { "consumes": [ - "multipart/form-data" + "multipart/form-data", + "application/json" ], "tags": [ "audio" @@ -2933,7 +2934,7 @@ const docTemplate = `{ }, { "type": "integer", - "description": "exact speaker count (\u003e0 forces; 0 = auto)", + "description": "exact speaker count (>0 forces; 0 = auto)", "name": "num_speakers", "in": "formData" }, @@ -2984,6 +2985,12 @@ const docTemplate = `{ "description": "json (default), verbose_json, or rttm", "name": "response_format", "in": "formData" + }, + { + "type": "boolean", + "description": "Export portable biometric profiles; requires voice-recognition permission. Omitted by default.", + "name": "include_speaker_profiles", + "in": "formData" } ], "responses": { @@ -2993,7 +3000,8 @@ const docTemplate = `{ "$ref": "#/definitions/schema.DiarizationResult" } } - } + }, + "description": "JSON accepts model, file (raw base64 audio), include_text, include_speaker_profiles and response_format. Profiles require voice-recognition permission and json or verbose_json; unsupported backends return 501." } }, "/v1/audio/speech": { @@ -4220,7 +4228,8 @@ const docTemplate = `{ "$ref": "#/definitions/schema.VoiceRegisterResponse" } } - } + }, + "description": "Supply either audio or speaker_profiles plus an explicit numeric speaker_slot. The selected model must expose matching trusted encoder metadata for portable enrollment. Registrations are global and ephemeral, with a fresh ID for each request." } }, "/v1/voice/verify": { @@ -4328,6 +4337,73 @@ const docTemplate = `{ } }, "definitions": { + "schema.SpeakerProfiles": { + "type": "object", + "properties": { + "version": { + "type": "integer" + }, + "encoder": { + "$ref": "#/definitions/schema.SpeakerEncoder" + }, + "speakers": { + "type": "array", + "items": { + "$ref": "#/definitions/schema.SpeakerProfile" + } + } + } + }, + "schema.SpeakerProfile": { + "type": "object", + "properties": { + "speaker": { + "type": "integer", + "description": "Raw numeric slot matching the decimal segment/summary label, not SPEAKER_NN." + }, + "clean_duration": { + "type": "number" + }, + "intervals": { + "type": "array", + "items": { + "$ref": "#/definitions/schema.SpeakerProfileInterval" + } + }, + "unavailable_reason": { + "type": "string", + "x-nullable": true + }, + "embedding": { + "type": "array", + "items": { + "type": "number" + } + } + } + }, + "schema.SpeakerProfileInterval": { + "type": "object", + "properties": { + "start": { + "type": "number" + }, + "end": { + "type": "number" + } + } + }, + "schema.SpeakerEncoder": { + "type": "object", + "properties": { + "identity": { + "type": "string" + }, + "dimension": { + "type": "integer" + } + } + }, "config.Gallery": { "type": "object", "properties": { @@ -5836,6 +5912,9 @@ const docTemplate = `{ }, "task": { "type": "string" + }, + "speaker_profiles": { + "$ref": "#/definitions/schema.SpeakerProfiles" } } }, @@ -8899,6 +8978,13 @@ const docTemplate = `{ }, "store": { "type": "string" + }, + "speaker_profiles": { + "$ref": "#/definitions/schema.SpeakerProfiles" + }, + "speaker_slot": { + "type": "integer", + "description": "Required with speaker_profiles, mutually exclusive with audio; explicit raw numeric speaker slot." } } }, diff --git a/swagger/swagger.json b/swagger/swagger.json index fbaa509b5..dd5380831 100644 --- a/swagger/swagger.json +++ b/swagger/swagger.json @@ -2907,7 +2907,8 @@ "/v1/audio/diarization": { "post": { "consumes": [ - "multipart/form-data" + "multipart/form-data", + "application/json" ], "tags": [ "audio" @@ -2930,7 +2931,7 @@ }, { "type": "integer", - "description": "exact speaker count (\u003e0 forces; 0 = auto)", + "description": "exact speaker count (>0 forces; 0 = auto)", "name": "num_speakers", "in": "formData" }, @@ -2981,6 +2982,12 @@ "description": "json (default), verbose_json, or rttm", "name": "response_format", "in": "formData" + }, + { + "type": "boolean", + "description": "Export portable biometric profiles; requires voice-recognition permission. Omitted by default.", + "name": "include_speaker_profiles", + "in": "formData" } ], "responses": { @@ -2990,7 +2997,8 @@ "$ref": "#/definitions/schema.DiarizationResult" } } - } + }, + "description": "JSON accepts model, file (raw base64 audio), include_text, include_speaker_profiles and response_format. Profiles require voice-recognition permission and json or verbose_json; unsupported backends return 501." } }, "/v1/audio/speech": { @@ -4217,7 +4225,8 @@ "$ref": "#/definitions/schema.VoiceRegisterResponse" } } - } + }, + "description": "Supply either audio or speaker_profiles plus an explicit numeric speaker_slot. The selected model must expose matching trusted encoder metadata for portable enrollment. Registrations are global and ephemeral, with a fresh ID for each request." } }, "/v1/voice/verify": { @@ -4325,6 +4334,73 @@ } }, "definitions": { + "schema.SpeakerProfiles": { + "type": "object", + "properties": { + "version": { + "type": "integer" + }, + "encoder": { + "$ref": "#/definitions/schema.SpeakerEncoder" + }, + "speakers": { + "type": "array", + "items": { + "$ref": "#/definitions/schema.SpeakerProfile" + } + } + } + }, + "schema.SpeakerProfile": { + "type": "object", + "properties": { + "speaker": { + "type": "integer", + "description": "Raw numeric slot matching the decimal segment/summary label, not SPEAKER_NN." + }, + "clean_duration": { + "type": "number" + }, + "intervals": { + "type": "array", + "items": { + "$ref": "#/definitions/schema.SpeakerProfileInterval" + } + }, + "unavailable_reason": { + "type": "string", + "x-nullable": true + }, + "embedding": { + "type": "array", + "items": { + "type": "number" + } + } + } + }, + "schema.SpeakerProfileInterval": { + "type": "object", + "properties": { + "start": { + "type": "number" + }, + "end": { + "type": "number" + } + } + }, + "schema.SpeakerEncoder": { + "type": "object", + "properties": { + "identity": { + "type": "string" + }, + "dimension": { + "type": "integer" + } + } + }, "config.Gallery": { "type": "object", "properties": { @@ -5833,6 +5909,9 @@ }, "task": { "type": "string" + }, + "speaker_profiles": { + "$ref": "#/definitions/schema.SpeakerProfiles" } } }, @@ -8896,6 +8975,13 @@ }, "store": { "type": "string" + }, + "speaker_profiles": { + "$ref": "#/definitions/schema.SpeakerProfiles" + }, + "speaker_slot": { + "type": "integer", + "description": "Required with speaker_profiles, mutually exclusive with audio; explicit raw numeric speaker slot." } } }, diff --git a/swagger/swagger.yaml b/swagger/swagger.yaml index 728d0398b..f4d10dd12 100644 --- a/swagger/swagger.yaml +++ b/swagger/swagger.yaml @@ -1,5 +1,50 @@ basePath: / definitions: + schema.SpeakerProfiles: + type: object + properties: + version: + type: integer + encoder: + $ref: '#/definitions/schema.SpeakerEncoder' + speakers: + type: array + items: + $ref: '#/definitions/schema.SpeakerProfile' + schema.SpeakerProfile: + type: object + properties: + speaker: + type: integer + description: Raw numeric slot matching the decimal segment/summary label, not + SPEAKER_NN. + clean_duration: + type: number + intervals: + type: array + items: + $ref: '#/definitions/schema.SpeakerProfileInterval' + unavailable_reason: + type: string + x-nullable: true + embedding: + type: array + items: + type: number + schema.SpeakerProfileInterval: + type: object + properties: + start: + type: number + end: + type: number + schema.SpeakerEncoder: + type: object + properties: + identity: + type: string + dimension: + type: integer config.Gallery: properties: artifact_verification: @@ -1050,6 +1095,7 @@ definitions: type: string type: object schema.DiarizationResult: + type: object properties: duration: type: number @@ -1058,16 +1104,17 @@ definitions: num_speakers: type: integer segments: + type: array items: $ref: '#/definitions/schema.DiarizationSegment' - type: array speakers: + type: array items: $ref: '#/definitions/schema.DiarizationSpeaker' - type: array task: type: string - type: object + speaker_profiles: + $ref: '#/definitions/schema.SpeakerProfiles' schema.DiarizationSegment: properties: end: @@ -3252,20 +3299,26 @@ definitions: type: array type: object schema.VoiceRegisterRequest: + type: object properties: audio: type: string labels: + type: object additionalProperties: type: string - type: object model: type: string name: type: string store: type: string - type: object + speaker_profiles: + $ref: '#/definitions/schema.SpeakerProfiles' + speaker_slot: + type: integer + description: Required with speaker_profiles, mutually exclusive with audio; + explicit raw numeric speaker slot. schema.VoiceRegisterResponse: properties: id: @@ -5338,62 +5391,70 @@ paths: post: consumes: - multipart/form-data + - application/json + tags: + - audio + summary: Identify speakers in audio (who spoke when). parameters: - - description: model - in: formData + - type: string + description: model name: model - required: true - type: string - - description: audio file in: formData + required: true + - type: file + description: audio file name: file + in: formData required: true - type: file - - description: exact speaker count (>0 forces; 0 = auto) - in: formData + - type: integer + description: exact speaker count (>0 forces; 0 = auto) name: num_speakers - type: integer - - description: lower bound when auto-detecting in: formData + - type: integer + description: lower bound when auto-detecting name: min_speakers - type: integer - - description: upper bound when auto-detecting in: formData + - type: integer + description: upper bound when auto-detecting name: max_speakers - type: integer - - description: clustering distance threshold when num_speakers is unknown in: formData + - type: number + description: clustering distance threshold when num_speakers is unknown name: clustering_threshold - type: number - - description: discard segments shorter than this (seconds) in: formData + - type: number + description: discard segments shorter than this (seconds) name: min_duration_on - type: number - - description: merge gaps shorter than this (seconds) in: formData + - type: number + description: merge gaps shorter than this (seconds) name: min_duration_off - type: number - - description: audio language hint (only meaningful for backends that bundle - ASR) in: formData + - type: string + description: audio language hint (only meaningful for backends that bundle ASR) name: language - type: string - - description: include per-segment transcript when the backend supports it in: formData + - type: boolean + description: include per-segment transcript when the backend supports it name: include_text - type: boolean - - description: json (default), verbose_json, or rttm in: formData + - type: string + description: json (default), verbose_json, or rttm name: response_format - type: string + in: formData + - type: boolean + description: Export portable biometric profiles; requires voice-recognition + permission. Omitted by default. + name: include_speaker_profiles + in: formData responses: - "200": + '200': description: OK schema: $ref: '#/definitions/schema.DiarizationResult' - summary: Identify speakers in audio (who spoke when). - tags: - - audio + description: JSON accepts model, file (raw base64 audio), include_text, include_speaker_profiles + and response_format. Profiles require voice-recognition permission and json + or verbose_json; unsupported backends return 501. /v1/audio/speech: post: consumes: @@ -6176,21 +6237,25 @@ paths: - voice-recognition /v1/voice/register: post: + tags: + - voice-recognition + summary: Register a speaker for 1:N identification. parameters: - description: query params - in: body name: request + in: body required: true schema: $ref: '#/definitions/schema.VoiceRegisterRequest' responses: - "200": + '200': description: Response schema: $ref: '#/definitions/schema.VoiceRegisterResponse' - summary: Register a speaker for 1:N identification. - tags: - - voice-recognition + description: Supply either audio or speaker_profiles plus an explicit numeric + speaker_slot. The selected model must expose matching trusted encoder metadata + for portable enrollment. Registrations are global and ephemeral, with a fresh + ID for each request. /v1/voice/verify: post: parameters: diff --git a/website/content/blog/diarization-speaker-profiles.md b/website/content/blog/diarization-speaker-profiles.md new file mode 100644 index 000000000..91592f8ce --- /dev/null +++ b/website/content/blog/diarization-speaker-profiles.md @@ -0,0 +1,40 @@ +--- +title: "Remember speakers from your recordings in LocalAI" +date: 2026-10-01 +author: "Ettore Di Giacinto" +category: "Engineering" +tags: ["diarization", "transcription", "voice-recognition"] +summary: "Name speakers from an existing conversation and recognize them in later recordings." +extracss: ["blog.css"] +--- + +LocalAI is adding a way to remember speakers directly from a conversation, alongside speaker turns and transcription. You can upload a recording, read who said what, and name voices for recognition in later recordings without collecting separate samples from each person. + +For an interview, a transcript with speakers lets you follow the questions and answers and return to the audio to check a quote. In a recurring meeting, remembered voices can put names on returning participants' contributions. A podcast editor can use speaker turns to locate a host or guest's speech before listening back and choosing a cut. + +To try this workflow, use a build containing [PR #12414](https://github.com/mudler/LocalAI/pull/12414) and its updated audio backend. + +## Three ways to use a recording + +**Speaker turns only** marks when each person speaks, without transcribing the words. It distinguishes people within the recording without knowing their names. + +**A transcript with speakers** adds the words to those turns, so you can read the conversation with each contribution attributed to a speaker. + +**Remembered names** matches voices against people you have explicitly named and saved. LocalAI can attach a saved name when it recognizes someone in another recording. You choose whom to remember; preparing a recording does not save everyone automatically. + +Recognition can mistake one person for another, especially when people talk over each other. Check the audio before relying on an attribution or quoting someone. + +## Get started + +Ask participants for permission before preparing their voices or naming and remembering them. For the named-transcript workflow in the web interface: + +1. Open **Models → Explore** and install the option with diarization, transcription, and speaker recognition. The [setup guide](/docs/features/audio-diarization/) lists the model choices and installation requirements. Wait for installation to finish. +2. Go to **Studio → Diarization**, select that model, and upload your recording. +3. Enable **Prepare speakers to remember**, then select **Diarize**. This includes transcript text and prepares speakers for naming. +4. In **Speakers**, listen to an available **Preview** for the person you want to name. Choose **Name and remember**, enter their name, and select **Remember**. Once saved, the name appears on that person's turns. + +If someone has too little clear speech, remembering them may be unavailable. Try a recording where they speak for longer without interruptions. + +On a later recording, use the same recognition-capable model to match saved voices. Keep preparation enabled if you want transcript text in Studio; with it off, the current UI returns speaker turns without text. + +Saved voices are currently lost when the LocalAI server restarts and are shared across users of the same instance. Agree on whose voices to remember before using this on a shared server.