diff --git a/Makefile b/Makefile index fb9ef79f3..64aff207a 100644 --- a/Makefile +++ b/Makefile @@ -850,7 +850,7 @@ test-extra-backend: protogen-go ## Convenience wrappers: build the image, then exercise it. test-extra-backend-llama-cpp: docker-build-llama-cpp BACKEND_IMAGE=local-ai-backend:llama-cpp \ - BACKEND_TEST_CAPS=health,load,predict,stream,logprobs,logit_bias \ + BACKEND_TEST_CAPS=health,load,predict,stream,logprobs,logit_bias,context_overflow \ $(MAKE) test-extra-backend ## Raw llama.cpp embeddings are required by Go-side pooling. This exercises the diff --git a/backend/backend.proto b/backend/backend.proto index 6a09b98eb..676e6222c 100644 --- a/backend/backend.proto +++ b/backend/backend.proto @@ -628,6 +628,7 @@ message TranscriptLiveConfig { string language = 1; // "" => model default int32 sample_rate = 2; // 0 => 16000; backends may reject others map params = 3; // backend-specific tuning + repeated KnownVoice known_voices = 4; // see DiarizeRequest.known_voices } message TranscriptLiveAudio { @@ -649,6 +650,7 @@ message LiveSpeakerSegment { string speaker = 1; // decimal speaker index int64 start = 2; // stream-relative nanoseconds int64 end = 3; + string name = 4; // registered speaker name when the backend identified the speaker, else empty } message LiveSoundEvent { @@ -825,6 +827,10 @@ message DiarizeRequest { // PredictOptions.ModelIdentity for the full rationale. Empty means "no // identity supplied" and backends MUST skip the check. string ModelIdentity = 11; + // 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 { @@ -833,9 +839,22 @@ message DiarizeSegment { float end = 3; // seconds string speaker = 4; // backend-emitted speaker label (e.g. "0", "SPEAKER_00") string text = 5; // optional per-segment transcript (empty unless include_text and supported) + string name = 6; // registered speaker name, empty when unknown or not identified + float name_score = 7; // match score of that name (cosine similarity), 0 when unnamed +} + +// KnownVoice is one registered voice: a name and its speaker embedding. `model` +// 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) @@ -891,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; @@ -900,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/cpp/audio-cpp/Makefile b/backend/cpp/audio-cpp/Makefile index 9a1928e52..d66d53628 100644 --- a/backend/cpp/audio-cpp/Makefile +++ b/backend/cpp/audio-cpp/Makefile @@ -9,7 +9,7 @@ # recipe is a make target (not a prepare.sh) so 'make purge && make' is a clean # rebuild and so the bump bot can see the pin. -AUDIO_CPP_VERSION?=ed96b7307c8daba2ebcf7912af928825f6b14cb9 +AUDIO_CPP_VERSION?=9a02e61326aaaf9d462b584ca5e0daba22c0abfc AUDIO_CPP_REPO?=https://github.com/0xShug0/audio.cpp CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST)))) diff --git a/backend/cpp/ik-llama-cpp/Makefile b/backend/cpp/ik-llama-cpp/Makefile index 167bf61ac..b8a9297cc 100644 --- a/backend/cpp/ik-llama-cpp/Makefile +++ b/backend/cpp/ik-llama-cpp/Makefile @@ -1,5 +1,5 @@ -IK_LLAMA_VERSION?=0821d62a8b356bd1db3c6765551a30bfcc44a6de +IK_LLAMA_VERSION?=d9e286846d6f8232db48ec5c111a4ea3aea675ef LLAMA_REPO?=https://github.com/ikawrakow/ik_llama.cpp CMAKE_ARGS?= diff --git a/backend/cpp/llama-cpp/CMakeLists.txt b/backend/cpp/llama-cpp/CMakeLists.txt index 50982b45d..3c195c37d 100644 --- a/backend/cpp/llama-cpp/CMakeLists.txt +++ b/backend/cpp/llama-cpp/CMakeLists.txt @@ -126,6 +126,11 @@ if(LLAMA_GRPC_BUILD_TESTS) target_compile_features(thread_params_test PRIVATE cxx_std_17) add_test(NAME thread_params_test COMMAND thread_params_test) + add_executable(parallel_params_test parallel_params_test.cpp parallel_params.h) + target_include_directories(parallel_params_test PRIVATE ${CMAKE_CURRENT_SOURCE_DIR}) + target_compile_features(parallel_params_test PRIVATE cxx_std_17) + add_test(NAME parallel_params_test COMMAND parallel_params_test) + add_executable(model_load_error_test model_load_error_test.cpp model_load_error.h) target_include_directories(model_load_error_test PRIVATE ${CMAKE_CURRENT_SOURCE_DIR}) target_compile_features(model_load_error_test PRIVATE cxx_std_17) diff --git a/backend/cpp/llama-cpp/Makefile b/backend/cpp/llama-cpp/Makefile index af5529d48..c8eb3dbc3 100644 --- a/backend/cpp/llama-cpp/Makefile +++ b/backend/cpp/llama-cpp/Makefile @@ -1,5 +1,5 @@ -LLAMA_VERSION?=4da6337767f973e2b4d0797e5b323d77d8565e4a +LLAMA_VERSION?=a868c3e3c56657f7e8a6231190dbbe90e7dd86c0 LLAMA_REPO?=https://github.com/ggerganov/llama.cpp CMAKE_ARGS?= diff --git a/backend/cpp/llama-cpp/grpc-server.cpp b/backend/cpp/llama-cpp/grpc-server.cpp index 8b206612c..3b908736e 100644 --- a/backend/cpp/llama-cpp/grpc-server.cpp +++ b/backend/cpp/llama-cpp/grpc-server.cpp @@ -55,6 +55,7 @@ #include "llama_compat.h" // fork-skew switches, generated by prepare.sh #include "model_load_error.h" #include "thread_params.h" +#include "parallel_params.h" #include "message_content.h" #include "passthrough_options.h" #include "stream_peer.h" @@ -545,6 +546,9 @@ static void params_parse(server_context& /*ctx_server*/, const backend::ModelOpt // slot state in host RAM and the backend grows without bound. // Initialize n_parallel to 1 by default (can be overridden by options) params.n_parallel = 1; + // Set only when the model options carry parallel/n_parallel, so an + // explicit 1 can be told apart from the default. + std::optional parallel_option; // Initialize grpc_servers to empty (can be overridden by options) std::string grpc_servers_option = ""; @@ -664,12 +668,9 @@ static void params_parse(server_context& /*ctx_server*/, const backend::ModelOpt } else if (!strcmp(optname, "parallel") || !strcmp(optname, "n_parallel")) { if (optval != NULL) { try { - params.n_parallel = std::stoi(optval_str); - if (params.n_parallel > 1) { - params.cont_batching = true; - } + parallel_option = std::stoi(optval_str); } catch (const std::exception& e) { - // If conversion fails, keep default value (1) + // If conversion fails, fall back to the environment/default } } } else if (!strcmp(optname, "grpc_servers") || !strcmp(optname, "rpc_servers")) { @@ -1238,19 +1239,12 @@ static void params_parse(server_context& /*ctx_server*/, const backend::ModelOpt } } - // Set params.n_parallel from environment variable if not set via options (fallback) - if (params.n_parallel == 1) { - const char *env_parallel = std::getenv("LLAMACPP_PARALLEL"); - if (env_parallel != NULL) { - try { - params.n_parallel = std::stoi(env_parallel); - if (params.n_parallel > 1) { - params.cont_batching = true; - } - } catch (const std::exception& e) { - // If conversion fails, keep default value (1) - } - } + // The model options win over LLAMACPP_PARALLEL, including an explicit + // parallel:1 (previously indistinguishable from the default and replaced + // by the environment value). + params.n_parallel = llama_grpc::resolve_n_parallel(parallel_option, std::getenv("LLAMACPP_PARALLEL")); + if (params.n_parallel > 1) { + params.cont_batching = true; } // Add RPC devices from option or environment variable (fallback) @@ -2182,10 +2176,10 @@ public: // connection is closed return grpc::Status(grpc::StatusCode::CANCELLED, "Request cancelled by client"); } else if (first_result->is_error()) { + // Return the error only as the status. Writing it as a Reply first + // made it the first content chunk: LocalAI streamed the error text + // as assistant output on an HTTP 200 instead of failing the request. json error_json = first_result->to_json(); - backend::Reply reply; - reply.set_message(error_json.value("message", "")); - writer->Write(reply); return grpc::Status(grpc::StatusCode::INTERNAL, error_json.value("message", "Error occurred")); } diff --git a/backend/cpp/llama-cpp/parallel_params.h b/backend/cpp/llama-cpp/parallel_params.h new file mode 100644 index 000000000..de88ba386 --- /dev/null +++ b/backend/cpp/llama-cpp/parallel_params.h @@ -0,0 +1,25 @@ +#pragma once + +#include +#include + +namespace llama_grpc { + +// resolve_n_parallel picks the slot count. A value from the model options +// always wins, including an explicit 1: the YAML takes precedence over the +// environment, as documented. LLAMACPP_PARALLEL is only the fallback when +// the options do not set it; a value that does not parse is ignored. +inline int resolve_n_parallel(const std::optional& from_options, const char* env, int fallback = 1) { + if (from_options) { + return *from_options; + } + if (env != nullptr) { + try { + return std::stoi(env); + } catch (const std::exception&) { + } + } + return fallback; +} + +} // namespace llama_grpc diff --git a/backend/cpp/llama-cpp/parallel_params_test.cpp b/backend/cpp/llama-cpp/parallel_params_test.cpp new file mode 100644 index 000000000..27b55b2fb --- /dev/null +++ b/backend/cpp/llama-cpp/parallel_params_test.cpp @@ -0,0 +1,29 @@ +#include "parallel_params.h" + +#include + +int main() { + // parallel:1 in the model options must not be replaced by the environment. + if (llama_grpc::resolve_n_parallel(1, "4") != 1) { + std::fprintf(stderr, "explicit parallel:1 was overwritten by LLAMACPP_PARALLEL\n"); + return 1; + } + if (llama_grpc::resolve_n_parallel(8, "4") != 8) { + std::fprintf(stderr, "explicit parallel was overwritten by LLAMACPP_PARALLEL\n"); + return 1; + } + // Without an option the environment applies, otherwise the default. + if (llama_grpc::resolve_n_parallel(std::nullopt, "4") != 4) { + std::fprintf(stderr, "LLAMACPP_PARALLEL was not used as fallback\n"); + return 1; + } + if (llama_grpc::resolve_n_parallel(std::nullopt, nullptr) != 1) { + std::fprintf(stderr, "default slot count is not 1\n"); + return 1; + } + if (llama_grpc::resolve_n_parallel(std::nullopt, "many") != 1) { + std::fprintf(stderr, "unparsable LLAMACPP_PARALLEL was not ignored\n"); + return 1; + } + return 0; +} diff --git a/backend/cpp/llama-cpp/patches/0001-add-server-task-type-score.patch b/backend/cpp/llama-cpp/patches/0001-add-server-task-type-score.patch index 253e8da5f..f2a34a149 100644 --- a/backend/cpp/llama-cpp/patches/0001-add-server-task-type-score.patch +++ b/backend/cpp/llama-cpp/patches/0001-add-server-task-type-score.patch @@ -97,7 +97,7 @@ index 3b5f6a1..d0e18e6 100644 json_schema = json(); task_prev = std::move(task); -@@ -2271,6 +2302,229 @@ private: +@@ -2271,6 +2302,227 @@ private: queue_results.send(std::move(res)); } @@ -121,7 +121,7 @@ index 3b5f6a1..d0e18e6 100644 + // region can straddle ubatch boundaries for long prompts, so this + // accumulates view by view instead of reading everything when the + // prompt completes. -+ void collect_score_logprobs(server_slot & slot, const llama_batch & batch) { ++ void collect_score_logprobs(server_slot & slot, const common_batch & batch) { + const int32_t n_prompt = slot.task->n_score_prompt; + const int32_t n_total = slot.task->n_tokens(); + const auto & suffixes = slot.task->score_suffixes; @@ -140,14 +140,14 @@ index 3b5f6a1..d0e18e6 100644 + + const int32_t n_vocab = llama_vocab_n_tokens(vocab); + -+ for (int32_t i = 0; i < batch.n_tokens; ++i) { -+ if (!batch.logits[i] || batch.seq_id[i][0] != slot.id) { ++ for (int32_t i = 0; i < batch.size(); ++i) { ++ if (!batch.tokens[i].output || batch.tokens[i].seq_id != slot.id) { + continue; + } + + // the output at position p predicts the task token at index p + 1; + // score tasks are text-only, so positions equal token indices -+ const int32_t target = batch.pos[i] + 1; ++ const int32_t target = batch.tokens[i].pos[0] + 1; + if (target < n_prompt || target > n_total) { + continue; + } @@ -249,7 +249,7 @@ index 3b5f6a1..d0e18e6 100644 + next++; + } + -+ llama_batch fb = llama_batch_init(n_tok, 0, 1); ++ common_batch fb(ctx_tgt); + + for (size_t k = 0; k < chunk.size(); ++k) { + const llama_seq_id seq = seq_base + (llama_seq_id) k; @@ -259,11 +259,11 @@ index 3b5f6a1..d0e18e6 100644 + llama_memory_seq_cp(mem, slot.id, seq, -1, -1); + + for (size_t j = 0; j < sfx.size(); ++j) { -+ common_batch_add(fb, sfx[j], pos0 + (llama_pos) j, { seq }, j + 1 < sfx.size()); ++ fb.add(sfx[j], pos0 + (llama_pos) j, seq, j + 1 < sfx.size()); + } + } + -+ const int ret = llama_decode(ctx_tgt, fb); ++ const int ret = llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, fb.get()); + + if (ret == 0) { + int32_t i = 0; @@ -290,8 +290,6 @@ index 3b5f6a1..d0e18e6 100644 + llama_memory_seq_rm(mem, seq_base + (llama_seq_id) k, -1, -1); + } + -+ llama_batch_free(fb); -+ + if (ret != 0) { + SLT_ERR(slot, "score suffix decode failed, ret = %d\n", ret); + return false; @@ -476,7 +474,7 @@ index 3b5f6a1..d0e18e6 100644 + // their outputs, not just the one holding the final token + if (slot.task && slot.task->type == SERVER_TASK_TYPE_SCORE && + (slot.state == SLOT_STATE_PROCESSING_PROMPT || slot.state == SLOT_STATE_DONE_PROMPT)) { -+ collect_score_logprobs(slot, batch_view); ++ collect_score_logprobs(slot, batch.view); + } + if (!is_inside_view(slot.i_batch)) { diff --git a/backend/cpp/llama-cpp/prepare.sh b/backend/cpp/llama-cpp/prepare.sh index 8bae396cc..a7fe9ab0f 100644 --- a/backend/cpp/llama-cpp/prepare.sh +++ b/backend/cpp/llama-cpp/prepare.sh @@ -63,6 +63,10 @@ cp -r tts_request_options_test.cpp llama.cpp/tools/grpc-server/ # Thread-count default normalization and its standalone regression test. cp -r thread_params.h llama.cpp/tools/grpc-server/ cp -r thread_params_test.cpp llama.cpp/tools/grpc-server/ +# Slot-count resolution (option over LLAMACPP_PARALLEL) and its standalone +# regression test. +cp -r parallel_params.h llama.cpp/tools/grpc-server/ +cp -r parallel_params_test.cpp llama.cpp/tools/grpc-server/ # Parent-death watcher (included by grpc-server.cpp) and its standalone unit # test (run via backend/cpp/run-unit-tests.sh; also buildable under ctest). cp -r parent_watch.h llama.cpp/tools/grpc-server/ diff --git a/backend/go/cloud-proxy/provider_anthropic.go b/backend/go/cloud-proxy/provider_anthropic.go index aa39e1ccb..397229bf6 100644 --- a/backend/go/cloud-proxy/provider_anthropic.go +++ b/backend/go/cloud-proxy/provider_anthropic.go @@ -5,6 +5,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" "net/http" @@ -30,8 +31,8 @@ import ( // message_stop (terminates the stream). Others are ignored. type anthropicRequest struct { - Model string `json:"model"` - MaxTokens int32 `json:"max_tokens"` + Model string `json:"model"` + MaxTokens int32 `json:"max_tokens"` // System is `any`: a bare string normally, or []anthropicSystemBlock // when cache_prompt is on (the block form carries cache_control). System any `json:"system,omitempty"` @@ -115,7 +116,10 @@ type anthropicResponse struct { Role string `json:"role"` Content []anthropicContentBlock `json:"content"` Model string `json:"model"` - Usage *anthropicUsage `json:"usage,omitempty"` + // StopReason is "end_turn", "max_tokens", "tool_use", ... or + // "refusal" when the model declines to answer. + StopReason string `json:"stop_reason,omitempty"` + Usage *anthropicUsage `json:"usage,omitempty"` } type anthropicUsage struct { @@ -140,8 +144,20 @@ type anthropicStreamDelta struct { Type string `json:"type,omitempty"` Text string `json:"text,omitempty"` PartialJSON string `json:"partial_json,omitempty"` + // StopReason is set on message_delta events. + StopReason string `json:"stop_reason,omitempty"` } +// anthropicStopRefusal is the stop_reason Anthropic returns when the +// model declines to answer. Such a response carries no (or only partial) +// content. Passing it through as a normal reply makes a refusal look like +// an empty, successful completion (finish_reason "stop", no content), so +// routers and agents cannot tell "declined" from "nothing to say" and +// never fall back. Surface it as an error instead. +const anthropicStopRefusal = "refusal" + +var errAnthropicRefusal = errors.New("cloud-proxy: upstream model refused to answer (stop_reason=refusal)") + // Anthropic requires max_tokens. If the caller didn't set it, use a // generous-but-bounded default so the request doesn't 400. const anthropicDefaultMaxTokens int32 = 4096 @@ -438,6 +454,9 @@ func (c *CloudProxy) predictAnthropicRich(ctx context.Context, cfg *proxyConfig, if err := json.NewDecoder(resp.Body).Decode(&parsed); err != nil { return nil, fmt.Errorf("cloud-proxy: decode response: %w", err) } + if parsed.StopReason == anthropicStopRefusal { + return nil, errAnthropicRefusal + } reply := &pb.Reply{} if parsed.Usage != nil { @@ -552,6 +571,9 @@ func (c *CloudProxy) predictAnthropicStreamRich(ctx context.Context, cfg *proxyC } } case "message_delta": + if ev.Delta != nil && ev.Delta.StopReason == anthropicStopRefusal { + return errAnthropicRefusal + } // Anthropic sends final usage in message_delta.usage. Emit // a usage-only Reply so the consumer can record totals. if ev.Usage != nil { diff --git a/backend/go/cloud-proxy/provider_anthropic_test.go b/backend/go/cloud-proxy/provider_anthropic_test.go index 2ed9a2127..17065201e 100644 --- a/backend/go/cloud-proxy/provider_anthropic_test.go +++ b/backend/go/cloud-proxy/provider_anthropic_test.go @@ -387,3 +387,72 @@ func TestPredict_Anthropic_PromptCache(t *testing.T) { g.Expect(off).NotTo(ContainSubstring("cache_control")) g.Expect(off).To(ContainSubstring(`"system":"be brief"`)) } + +// A refusal must not look like an empty, successful completion. +func TestPredict_Anthropic_RefusalIsAnError(t *testing.T) { + g := NewWithT(t) + srv, _ := fakeAnthropicUpstream(t, func(_ anthropicRequest) (int, string, string) { + return 200, `{"type":"message","role":"assistant","content":[],"stop_reason":"refusal","usage":{"input_tokens":9,"output_tokens":0}}`, "application/json" + }) + defer srv.Close() + cp := newAnthropicTranslateCloudProxy(t, srv.URL) + + _, err := cp.Predict(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}, Tokens: 16}) + g.Expect(err).To(MatchError(errAnthropicRefusal)) +} + +// Counter-check: an empty end_turn reply stays a normal, successful reply. +func TestPredict_Anthropic_EmptyEndTurnIsNotAnError(t *testing.T) { + g := NewWithT(t) + srv, _ := fakeAnthropicUpstream(t, func(_ anthropicRequest) (int, string, string) { + return 200, `{"type":"message","role":"assistant","content":[],"stop_reason":"end_turn"}`, "application/json" + }) + defer srv.Close() + cp := newAnthropicTranslateCloudProxy(t, srv.URL) + + got, err := cp.Predict(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}, Tokens: 16}) + g.Expect(err).NotTo(HaveOccurred()) + g.Expect(got).To(Equal("")) +} + +func streamAnthropic(t *testing.T, frames []string) ([]string, error) { + t.Helper() + srv, _ := fakeAnthropicUpstream(t, func(_ anthropicRequest) (int, string, string) { + return 200, strings.Join(frames, ""), "text/event-stream" + }) + defer srv.Close() + cp := newAnthropicTranslateCloudProxy(t, srv.URL) + results := make(chan string, 8) + done := make(chan error, 1) + go func() { + done <- cp.PredictStream(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}, Tokens: 16}, results) + }() + var got []string + for s := range results { + got = append(got, s) + } + return got, <-done +} + +func TestPredictStream_Anthropic_RefusalIsAnError(t *testing.T) { + g := NewWithT(t) + _, err := streamAnthropic(t, []string{ + "event: message_start\ndata: {\"type\":\"message_start\"}\n\n", + "event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"refusal\"},\"usage\":{\"output_tokens\":0}}\n\n", + "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n", + }) + g.Expect(err).To(MatchError(errAnthropicRefusal)) +} + +// Counter-check: a regular stream ending with stop_reason end_turn succeeds. +func TestPredictStream_Anthropic_EndTurnIsNotAnError(t *testing.T) { + g := NewWithT(t) + got, err := streamAnthropic(t, []string{ + "event: message_start\ndata: {\"type\":\"message_start\"}\n\n", + "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"ok\"}}\n\n", + "event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"},\"usage\":{\"output_tokens\":1}}\n\n", + "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n", + }) + g.Expect(err).NotTo(HaveOccurred()) + g.Expect(strings.Join(got, "")).To(Equal("ok")) +} diff --git a/backend/go/crispasr/Makefile b/backend/go/crispasr/Makefile index 947023262..9e3863b44 100644 --- a/backend/go/crispasr/Makefile +++ b/backend/go/crispasr/Makefile @@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1) # CrispASR version (release tag) CRISPASR_REPO?=https://github.com/CrispStrobe/CrispASR -CRISPASR_VERSION?=be202c472503a5c7f1d3e568c420865cad02f1c3 +CRISPASR_VERSION?=fdc3a0007d68f8d3905f20e9cd4d18b5f193e096 SO_TARGET?=libgocrispasr.so CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF diff --git a/backend/go/nemo-speech-cpp/Makefile b/backend/go/nemo-speech-cpp/Makefile index 81a713178..9d76b91f7 100644 --- a/backend/go/nemo-speech-cpp/Makefile +++ b/backend/go/nemo-speech-cpp/Makefile @@ -12,7 +12,7 @@ # runs 'make -C backend/go/$(BACKEND) build' and then copies package/), so it # has to produce the binary and the package, not just the shared libraries. -NEMO_SPEECH_VERSION?=0f706e43cf1fbc031bad1423e05460d3acaeaa1c +NEMO_SPEECH_VERSION?=4c101bc7113f49101a3e11d2c994c519f41939f6 NEMO_SPEECH_REPO?=https://github.com/NVIDIA/NeMo-Speech.cpp GOCMD?=go diff --git a/backend/go/parakeet-cpp/Makefile b/backend/go/parakeet-cpp/Makefile index 59fe56717..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?=623a968bccbd2214588df398fcce687cd4218dea +# 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?=623a968bccbd2214588df398fcce687cd4218dea +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 1853aa666..8cbfd33f6 100644 --- a/backend/go/parakeet-cpp/diarize.go +++ b/backend/go/parakeet-cpp/diarize.go @@ -27,7 +27,28 @@ type diarizeSegmentJSON struct { // not the count of speakers actually present, so it is not read here; the // response's num_speakers is computed from distinct segment labels instead. type diarizePCMDoc struct { - Segments []diarizeSegmentJSON `json:"segments"` + Segments []diarizeSegmentJSON `json:"segments"` + Names map[string]speakerNameJSON `json:"names"` +} + +// speakerNameJSON mirrors one value of the "names" map the named C-API functions add: +// {"0":{"name":"Ada","score":0.93}}. +type speakerNameJSON struct { + Name string `json:"name"` + Score float32 `json:"score"` +} + +// nameFor returns the registered name of a diarization slot, or "" for an unknown slot, a slot +// with no matching voice, or speaker -1 (no diarized speaker). +func nameFor(names map[string]speakerNameJSON, speaker int) (string, float32) { + if speaker < 0 || len(names) == 0 { + return "", 0 + } + n, ok := names[strconv.Itoa(speaker)] + if !ok || n.Name == "" { + return "", 0 + } + return n.Name, n.Score } // diarizeUtteranceJSON mirrors one element of @@ -45,7 +66,8 @@ type diarizeUtteranceJSON struct { // consumed here; the per-word "words" detail belongs to a speaker-attributed // transcript RPC, not Diarize. type transcribeAndDiarizeDoc struct { - Utterances []diarizeUtteranceJSON `json:"utterances"` + Utterances []diarizeUtteranceJSON `json:"utterances"` + Names map[string]speakerNameJSON `json:"names"` } // speakerLabel renders a 0-based speaker index as the decimal string @@ -87,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") @@ -114,62 +141,166 @@ 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 + } - raw, err := p.diarizeCall(pcm, wantText) + var reg uintptr + if len(req.GetKnownVoices()) > 0 && p.spkCtx != 0 { + reg, err = p.buildSpeakerRegistry(req.GetKnownVoices()) + if err != nil { + return pb.DiarizeResponse{}, err + } + defer p.freeSpeakerRegistry(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) (string, error) { +func (p *ParakeetCpp) diarizeCall(pcm []float32, wantText bool, reg uintptr, wantProfiles bool) (string, error) { p.engineMu.Lock() defer p.engineMu.Unlock() - if p.diarCtx == 0 || (wantText && p.ctxPtr == 0) { + if p.diarCtx == 0 || (wantText && p.ctxPtr == 0) || (reg != 0 && p.spkCtx == 0) { 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 - if wantText { + 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, + "parakeet-cpp: naming speakers needs libparakeet.so ABI 10 (parakeet_capi_transcribe_and_diarize_named_json)") + } + // This C function takes no threshold or margin, so the text path uses the C side's + // defaults rather than speaker_threshold / speaker_margin. + cstr = CppTranscribeAndDiarizeNamedJSON(p.ctxPtr, p.diarCtx, p.spkCtx, reg, &pcm[0], int32(len(pcm)), 16000) + case reg != 0: + if CppDiarizeNamedPCMJSON == nil { + return "", status.Error(codes.Unimplemented, + "parakeet-cpp: naming speakers needs libparakeet.so ABI 10 (parakeet_capi_diarize_named_pcm_json)") + } + cstr = CppDiarizeNamedPCMJSON(p.diarCtx, p.spkCtx, reg, &pcm[0], int32(len(pcm)), 16000, p.speakerAccept, p.speakerMargin) + case wantText: cstr = CppTranscribeAndDiarizeJSON(p.ctxPtr, p.diarCtx, &pcm[0], int32(len(pcm)), 16000) - } else { + default: cstr = CppDiarizePCM(p.diarCtx, &pcm[0], int32(len(pcm)), 16000) } if cstr == 0 { - return "", fmt.Errorf("parakeet-cpp: diarize failed: %s", diarizeLastError(p, wantText)) + 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 // whichever came back non-empty. Called under the same engineMu as the // failing call (last_error is ctx-shared state). -func diarizeLastError(p *ParakeetCpp, wantText bool) string { +func diarizeLastError(p *ParakeetCpp, wantText, named bool) string { var msgs []string if m := CppLastError(p.diarCtx); m != "" { msgs = append(msgs, m) @@ -179,6 +310,11 @@ func diarizeLastError(p *ParakeetCpp, wantText bool) string { msgs = append(msgs, m) } } + if named { + if m := CppLastError(p.spkCtx); m != "" { + msgs = append(msgs, m) + } + } if len(msgs) == 0 { return "unknown error" } @@ -196,11 +332,14 @@ func parseDiarizeDoc(raw string, wantText bool) ([]*pb.DiarizeSegment, error) { } segs := make([]*pb.DiarizeSegment, 0, len(doc.Utterances)) for _, u := range doc.Utterances { + name, score := nameFor(doc.Names, u.Speaker) segs = append(segs, &pb.DiarizeSegment{ - Start: float32(u.Start), - End: float32(u.End), - Speaker: speakerLabel(u.Speaker), - Text: u.Text, + Start: float32(u.Start), + End: float32(u.End), + Speaker: speakerLabel(u.Speaker), + Text: u.Text, + Name: name, + NameScore: score, }) } return segs, nil @@ -212,10 +351,13 @@ func parseDiarizeDoc(raw string, wantText bool) ([]*pb.DiarizeSegment, error) { } segs := make([]*pb.DiarizeSegment, 0, len(doc.Segments)) for _, s := range doc.Segments { + name, score := nameFor(doc.Names, s.Speaker) segs = append(segs, &pb.DiarizeSegment{ - Start: float32(s.Start), - End: float32(s.End), - Speaker: speakerLabel(s.Speaker), + Start: float32(s.Start), + End: float32(s.End), + Speaker: speakerLabel(s.Speaker), + Name: name, + NameScore: score, }) } return segs, nil diff --git a/backend/go/parakeet-cpp/diarize_test.go b/backend/go/parakeet-cpp/diarize_test.go index 1db076771..c97d7baf6 100644 --- a/backend/go/parakeet-cpp/diarize_test.go +++ b/backend/go/parakeet-cpp/diarize_test.go @@ -40,7 +40,21 @@ func diarizeStubs() (restore func()) { savedTranscribeAndDiarize := CppTranscribeAndDiarizeJSON savedFreeString := CppFreeString savedLastError := CppLastError + savedNamedDiarize := CppDiarizeNamedPCMJSON + savedNamedText := CppTranscribeAndDiarizeNamedJSON + savedRegNew := CppSpeakerRegistryNew + savedRegFree := CppSpeakerRegistryFree + savedRegAdd := CppSpeakerRegistryAddEmbedding + savedRegLastError := CppSpeakerRegistryLastError + savedSpeakerDim := CppSpeakerDim return func() { + CppDiarizeNamedPCMJSON = savedNamedDiarize + CppTranscribeAndDiarizeNamedJSON = savedNamedText + CppSpeakerRegistryNew = savedRegNew + CppSpeakerRegistryFree = savedRegFree + CppSpeakerRegistryAddEmbedding = savedRegAdd + CppSpeakerRegistryLastError = savedRegLastError + CppSpeakerDim = savedSpeakerDim CppDiarizePCM = savedDiarize CppTranscribeAndDiarizeJSON = savedTranscribeAndDiarize CppFreeString = savedFreeString @@ -245,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) + _, 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") }) @@ -272,4 +286,161 @@ var _ = Describe("ParakeetCpp.Diarize", func() { Expect(resp.Segments[1].Start).To(BeNumerically("~", 1.05, 0.001)) Expect(resp.Segments[1].End).To(BeNumerically("~", 1.20, 0.001)) }) + Describe("with known voices", func() { + var freed []uintptr + var used string + ada := []*pb.KnownVoice{{Name: "Ada", Embedding: []float32{1, 0}}} + BeforeEach(func() { + freed, used = nil, "" + CppFreeString = func(uintptr) {} + CppSpeakerDim = func(uintptr) int32 { return 2 } + CppSpeakerRegistryNew = func() uintptr { return 9 } + CppSpeakerRegistryFree = func(r uintptr) { freed = append(freed, r) } + CppSpeakerRegistryAddEmbedding = func(uintptr, string, *float32, int32) int32 { return 0 } + CppDiarizePCM = func(uintptr, *float32, int32, int32) uintptr { + used = "plain" + return pool.cstr(`{"speakers":8,"segments":[{"speaker":0,"start":0.5,"end":2}]}`) + } + CppDiarizeNamedPCMJSON = func(diar, spk, reg uintptr, s *float32, n, sr int32, accept, margin float32) uintptr { + used = "named" + Expect(reg).To(Equal(uintptr(9))) + Expect(accept).To(BeNumerically("~", 0.7, 1e-6)) + Expect(margin).To(BeNumerically("~", 0.05, 1e-6)) + return pool.cstr(`{"speakers":8,"segments":[{"speaker":0,"start":0.5,"end":2.0},{"speaker":1,"start":2.5,"end":4.0}],` + + `"names":{"0":{"name":"Ada","score":0.93},"1":{"name":"","score":0.2}}}`) + } + }) + + 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}) + Expect(err).ToNot(HaveOccurred()) + Expect(used).To(Equal("named")) + Expect(res.Segments).To(HaveLen(2)) + Expect(res.Segments[0].Name).To(Equal("Ada")) + Expect(res.Segments[0].NameScore).To(BeNumerically("~", 0.93, 1e-6)) + Expect(res.Segments[1].Name).To(BeEmpty()) + Expect(freed).To(Equal([]uintptr{9})) + }) + It("uses the plain path when the request has no known voices", func() { + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + res, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5)}) + Expect(err).ToNot(HaveOccurred()) + Expect(used).To(Equal("plain")) + Expect(res.Segments[0].Name).To(BeEmpty()) + Expect(freed).To(BeEmpty()) + }) + It("uses the plain path when no speaker model is loaded, even with known voices", func() { + p := &ParakeetCpp{diarCtx: 1} + _, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), KnownVoices: ada}) + Expect(err).ToNot(HaveOccurred()) + Expect(used).To(Equal("plain")) + }) + It("uses the plain path, without a registry, when no voice has an embedding", func() { + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + _, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), + KnownVoices: []*pb.KnownVoice{{Name: "Ada"}, {Embedding: []float32{1, 0}}}}) + Expect(err).ToNot(HaveOccurred()) + Expect(used).To(Equal("plain")) + Expect(freed).To(Equal([]uintptr{9})) // the empty registry built for it is released, once + }) + It("takes the plain path, without failing, when the only voice has the wrong size", func() { + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + res, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), + KnownVoices: []*pb.KnownVoice{{Name: "Ada", Embedding: []float32{1, 0, 0}}}}) + Expect(err).ToNot(HaveOccurred()) + Expect(used).To(Equal("plain")) + Expect(res.Segments[0].Name).To(BeEmpty()) + Expect(freed).To(Equal([]uintptr{9})) + }) + It("takes the plain path when the C side refuses the only voice", func() { + CppSpeakerRegistryAddEmbedding = func(uintptr, string, *float32, int32) int32 { return 1 } + CppSpeakerRegistryLastError = func(uintptr) string { return "nope" } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + _, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), KnownVoices: ada}) + Expect(err).ToNot(HaveOccurred()) + Expect(used).To(Equal("plain")) + Expect(freed).To(Equal([]uintptr{9})) + }) + It("reports a missing v10 symbol instead of silently dropping the names", func() { + CppDiarizeNamedPCMJSON = nil + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + _, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), KnownVoices: ada}) + Expect(err).To(HaveOccurred()) + Expect(status.Code(err)).To(Equal(codes.Unimplemented)) + Expect(freed).To(Equal([]uintptr{9})) + }) + It("reports a missing named transcribe symbol on the include_text path", func() { + CppTranscribeAndDiarizeJSON = func(asr, diar uintptr, s *float32, n, sr int32) uintptr { return 0 } + CppTranscribeAndDiarizeNamedJSON = nil + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2, ctxPtr: 3} + _, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), IncludeText: true, KnownVoices: ada}) + Expect(err).To(HaveOccurred()) + Expect(status.Code(err)).To(Equal(codes.Unimplemented)) + Expect(freed).To(Equal([]uintptr{9})) + }) + It("names utterances on the include_text path", func() { + // wantText also requires the plain text symbol, present in any library that has the named one. + CppTranscribeAndDiarizeJSON = func(asr, diar uintptr, s *float32, n, sr int32) uintptr { return 0 } + CppTranscribeAndDiarizeNamedJSON = func(asr, diar, spk, reg uintptr, s *float32, n, sr int32) uintptr { + used = "named-text" + return pool.cstr(`{"speakers":8,"names":{"0":{"name":"Ada","score":0.9}},"utterances":[{"speaker":0,"name":"Ada","text":"hello","start":0.5,"end":2.0,"conf":0.9},{"speaker":-1,"text":"hm","start":2.5,"end":3.0,"conf":0.5}],"words":[]}`) + } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2, ctxPtr: 3} + res, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), IncludeText: true, KnownVoices: ada}) + Expect(err).ToNot(HaveOccurred()) + Expect(used).To(Equal("named-text")) + Expect(res.Segments[0].Name).To(Equal("Ada")) + Expect(res.Segments[0].Text).To(Equal("hello")) + Expect(res.Segments[1].Name).To(BeEmpty()) // speaker -1 has no name + Expect(freed).To(Equal([]uintptr{9})) + }) + It("keeps the name when close segments of one speaker are merged", func() { + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2, speakerAccept: 0.7, speakerMargin: 0.05} + CppDiarizeNamedPCMJSON = func(diar, spk, reg uintptr, s *float32, n, sr int32, a, m float32) uintptr { + return pool.cstr(`{"speakers":8,"segments":[{"speaker":0,"start":0.5,"end":2.0},{"speaker":0,"start":2.1,"end":3.0}],"names":{"0":{"name":"Ada","score":0.9}}}`) + } + res, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), MinDurationOff: 0.5, KnownVoices: ada}) + Expect(err).ToNot(HaveOccurred()) + Expect(res.Segments).To(HaveLen(1)) + Expect(res.Segments[0].Name).To(Equal("Ada")) + Expect(res.Segments[0].NameScore).To(BeNumerically("~", 0.9, 1e-6)) + }) + It("includes the speaker context message when the named call fails", func() { + CppDiarizeNamedPCMJSON = func(diar, spk, reg uintptr, s *float32, n, sr int32, a, m float32) uintptr { return 0 } + CppLastError = func(ctx uintptr) string { + switch ctx { + case 1: + return "diar side broke" + case 2: + return "speaker side broke" + } + return "" + } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + _, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(5), KnownVoices: ada}) + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("diar side broke")) + Expect(err.Error()).To(ContainSubstring("speaker side broke")) + Expect(freed).To(Equal([]uintptr{9})) + }) + }) }) diff --git a/backend/go/parakeet-cpp/goparakeetcpp.go b/backend/go/parakeet-cpp/goparakeetcpp.go index e8ed1c525..04c26bff6 100644 --- a/backend/go/parakeet-cpp/goparakeetcpp.go +++ b/backend/go/parakeet-cpp/goparakeetcpp.go @@ -104,6 +104,23 @@ var ( CppSceneStreamFeedJSON func(s uintptr, pcm *float32, n int32, isLast int32) uintptr CppSceneStreamLastError func(s uintptr) string CppSceneStreamFree func(s uintptr) + + // Speaker identification. CppSpeakerDim, the registry, CppSceneStreamBeginSpeaker and + // 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) + CppSpeakerRegistryAddEmbedding func(reg uintptr, name string, emb *float32, dim int32) int32 + CppSpeakerRegistryLastError func(reg uintptr) string + CppSceneStreamBeginSpeaker func(asr, diar, tagger, speaker, reg uintptr, o *cSceneOpts) uintptr + // CppDiarizeNamedPCMJSON takes two float32 arguments (acceptThreshold, margin), which + // purego passes in floating-point registers. Not exercised without the real library. + CppDiarizeNamedPCMJSON func(diar, speaker, reg uintptr, samples *float32, n, sampleRate int32, acceptThreshold, margin float32) uintptr + CppTranscribeAndDiarizeNamedJSON func(asr, diar, speaker, reg uintptr, samples *float32, n, sampleRate int32) uintptr ) // cSoundOpts and cSceneOpts mirror parakeet_sound_opts / parakeet_scene_opts @@ -121,6 +138,14 @@ type cSceneOpts struct { DiarLatency int32 Sound cSoundOpts Flags int32 + // Speaker identification (parakeet_scene_opts, ABI v9). The C side reads + // these only when Size covers them, so a Go struct built against a v9 + // library and run against a v8 one is still valid. + SpeakerAcceptThreshold float32 + SpeakerMargin float32 + SpeakerMinVoiceSec float32 + SpeakerRefreshSec float32 + SpeakerMaxVoiceSec float32 } // streamChunkSamples is how much 16 kHz mono PCM we hand to stream_feed per @@ -193,6 +218,12 @@ type ParakeetCpp struct { // companion. See roles.go. diarCtx uintptr tagCtx uintptr + // spkCtx is the speaker encoder context (speaker_model: companion); 0 when speaker + // naming is off. speakerAccept is the cosine acceptance threshold and speakerMargin + // the runner-up margin, both from the model options. + spkCtx uintptr + speakerAccept float32 + speakerMargin float32 // diarLatency is the PARAKEET_DIAR_LATENCY_* mode for diarization // streaming (diarization_latency: option, default "low"). Unused until // the diarization/scene streaming paths land. @@ -953,7 +984,7 @@ func (p *ParakeetCpp) Free() error { // re-checks ctxPtr under the lock) can never feed into a freed ctx. p.engineMu.Lock() defer p.engineMu.Unlock() - for _, ctxField := range [...]*uintptr{&p.ctxPtr, &p.diarCtx, &p.tagCtx} { + for _, ctxField := range [...]*uintptr{&p.ctxPtr, &p.diarCtx, &p.tagCtx, &p.spkCtx} { if *ctxField != 0 { CppFree(*ctxField) *ctxField = 0 diff --git a/backend/go/parakeet-cpp/goparakeetcpp_test.go b/backend/go/parakeet-cpp/goparakeetcpp_test.go index eda5eaa36..a1364b15d 100644 --- a/backend/go/parakeet-cpp/goparakeetcpp_test.go +++ b/backend/go/parakeet-cpp/goparakeetcpp_test.go @@ -75,6 +75,24 @@ func ensureLibLoaded() { purego.RegisterLibFunc(&CppSoundStreamDrainScoresJSON, lib, "parakeet_capi_sound_stream_drain_scores_json") purego.RegisterLibFunc(&CppFreeSoundSegments, lib, "parakeet_capi_free_sound_segments") purego.RegisterLibFunc(&CppSoundStreamFree, lib, "parakeet_capi_sound_stream_free") + purego.RegisterLibFunc(&CppSceneOptsDefault, lib, "parakeet_capi_scene_opts_default") + purego.RegisterLibFunc(&CppSceneStreamBegin, lib, "parakeet_capi_scene_stream_begin") + purego.RegisterLibFunc(&CppSceneStreamFeedJSON, lib, "parakeet_capi_scene_stream_feed_json") + purego.RegisterLibFunc(&CppSceneStreamLastError, lib, "parakeet_capi_scene_stream_last_error") + purego.RegisterLibFunc(&CppSceneStreamFree, lib, "parakeet_capi_scene_stream_free") + } + // Speaker identification (ABI v9 and v10), registered exactly as main.go does. + if sym, err := purego.Dlsym(lib, "parakeet_capi_scene_stream_begin_speaker"); err == nil && sym != 0 { + purego.RegisterLibFunc(&CppSpeakerDim, lib, "parakeet_capi_speaker_dim") + purego.RegisterLibFunc(&CppSpeakerRegistryNew, lib, "parakeet_capi_speaker_registry_new") + purego.RegisterLibFunc(&CppSpeakerRegistryFree, lib, "parakeet_capi_speaker_registry_free") + purego.RegisterLibFunc(&CppSpeakerRegistryLastError, lib, "parakeet_capi_speaker_registry_last_error") + purego.RegisterLibFunc(&CppSceneStreamBeginSpeaker, lib, "parakeet_capi_scene_stream_begin_speaker") + purego.RegisterLibFunc(&CppTranscribeAndDiarizeNamedJSON, lib, "parakeet_capi_transcribe_and_diarize_named_json") + } + if sym, err := purego.Dlsym(lib, "parakeet_capi_diarize_named_pcm_json"); err == nil && sym != 0 { + purego.RegisterLibFunc(&CppSpeakerRegistryAddEmbedding, lib, "parakeet_capi_speaker_registry_add_embedding") + purego.RegisterLibFunc(&CppDiarizeNamedPCMJSON, lib, "parakeet_capi_diarize_named_pcm_json") } purego.RegisterLibFunc(&CppFreeString, lib, "parakeet_capi_free_string") purego.RegisterLibFunc(&CppLastError, lib, "parakeet_capi_last_error") diff --git a/backend/go/parakeet-cpp/live.go b/backend/go/parakeet-cpp/live.go index 6497779d5..782c3e8a7 100644 --- a/backend/go/parakeet-cpp/live.go +++ b/backend/go/parakeet-cpp/live.go @@ -81,7 +81,7 @@ func (p *ParakeetCpp) AudioTranscriptionLive(in <-chan *pb.TranscriptLiveRequest // scene error. var scene sceneStreamHandle if p.sceneWanted() { - scene = p.sceneBegin() + scene = p.sceneBegin(cfg.GetKnownVoices()) if scene.s == 0 { xlog.Warn("parakeet-cpp: scene stream begin failed; live continues without speaker/sound events") } @@ -118,7 +118,7 @@ func (p *ParakeetCpp) AudioTranscriptionLive(in <-chan *pb.TranscriptLiveRequest if r.Delta != "" { full.WriteString(r.Delta) } - speakers := liveSpeakersToProto(sceneDoc.Speakers) + speakers := liveSpeakersToProto(sceneDoc.Speakers, sceneDoc.Names) sounds := liveSoundsToProto(sceneDoc.Sounds) if r.Delta != "" || r.Eou || r.Eob || len(r.Words) > 0 || len(speakers) > 0 || len(sounds) > 0 { out <- &pb.TranscriptLiveResponse{ @@ -154,7 +154,7 @@ func (p *ParakeetCpp) AudioTranscriptionLive(in <-chan *pb.TranscriptLiveRequest p.sceneFree(scene) scene = sceneStreamHandle{} if p.sceneWanted() { - scene = p.sceneBegin() + scene = p.sceneBegin(payload.Config.GetKnownVoices()) if scene.s == 0 { xlog.Warn("parakeet-cpp: scene stream begin failed; live continues without speaker/sound events") } diff --git a/backend/go/parakeet-cpp/live_test.go b/backend/go/parakeet-cpp/live_test.go index 11dc496b6..4ee4ee61e 100644 --- a/backend/go/parakeet-cpp/live_test.go +++ b/backend/go/parakeet-cpp/live_test.go @@ -47,7 +47,13 @@ func liveStubs() (restore func()) { savedSceneFeedJSON := CppSceneStreamFeedJSON savedSceneLastError := CppSceneStreamLastError savedSceneFree := CppSceneStreamFree + savedSceneBeginSpk := CppSceneStreamBeginSpeaker + savedRegNew, savedRegFree := CppSpeakerRegistryNew, CppSpeakerRegistryFree + savedRegAdd, savedSpkDim := CppSpeakerRegistryAddEmbedding, CppSpeakerDim return func() { + CppSceneStreamBeginSpeaker = savedSceneBeginSpk + CppSpeakerRegistryNew, CppSpeakerRegistryFree = savedRegNew, savedRegFree + CppSpeakerRegistryAddEmbedding, CppSpeakerDim = savedRegAdd, savedSpkDim CppStreamBegin, CppStreamBeginLang = savedBegin, savedBeginLang CppStreamFeed, CppStreamFeedJSON = savedFeed, savedFeedJSON CppStreamFinalize, CppStreamFinalizeJSON = savedFinalize, savedFinalizeJSON @@ -600,7 +606,7 @@ var _ = Describe("AudioTranscriptionLive scene events (stubbed C API)", func() { return pool.cstr(`{"speakers":[],"sounds":[]}`) } - h := p.sceneBegin() + h := p.sceneBegin(nil) Expect(h.s).NotTo(BeZero()) // Simulate a Free() racing in between the begin and the next feed: it @@ -672,6 +678,142 @@ var _ = Describe("AudioTranscriptionLive scene events (stubbed C API)", func() { }) }) +var _ = Describe("AudioTranscriptionLive named speakers (stubbed C API)", func() { + var ( + pool *liveCstrPool + restore func() + p *ParakeetCpp + regs []uintptr + plain int + spkBeg int + gotReg uintptr + order []string + ) + + liveVoicesConfig := func(voices ...*pb.KnownVoice) *pb.TranscriptLiveRequest { + return &pb.TranscriptLiveRequest{ + Payload: &pb.TranscriptLiveRequest_Config{Config: &pb.TranscriptLiveConfig{KnownVoices: voices}}, + } + } + ada := &pb.KnownVoice{Name: "Ada", Embedding: []float32{1, 0}} + + BeforeEach(func() { + pool = &liveCstrPool{} + restore = liveStubs() + p = &ParakeetCpp{ctxPtr: 1, diarCtx: 2, spkCtx: 3, speakerAccept: 0.7, speakerMargin: 0.05} + regs, plain, spkBeg, gotReg, order = nil, 0, 0, 0, nil + + CppStreamBeginLang = nil + CppStreamBegin = func(ctx uintptr) uintptr { return 7 } + CppStreamFree = func(s uintptr) {} + CppFreeString = func(s uintptr) {} + CppLastError = func(ctx uintptr) string { return "stub error" } + CppStreamFeed = nil + CppStreamFeedJSON = func(s uintptr, pcm []float32, n int32) uintptr { + return pool.cstr(`{"text":"","eou":0,"frame_sec":0.08,"words":[]}`) + } + CppStreamFinalize = nil + CppStreamFinalizeJSON = func(s uintptr) uintptr { + return pool.cstr(`{"text":"","eou":0,"frame_sec":0.08,"words":[]}`) + } + + CppSceneOptsDefault = func(o *cSceneOpts) { *o = cSceneOpts{} } + CppSceneStreamBegin = func(asr, diar, tagger uintptr, o *cSceneOpts) uintptr { plain++; return 100 } + CppSceneStreamBeginSpeaker = func(asr, diar, tag, spk, reg uintptr, o *cSceneOpts) uintptr { + spkBeg++ + gotReg = reg + return 200 + } + CppSceneStreamFeedJSON = func(s uintptr, pcm *float32, n int32, isLast int32) uintptr { + return pool.cstr(`{"speakers":[{"speaker":0,"start":0.1,"end":0.6}],"sounds":[],` + + `"names":{"0":{"name":"Ada","score":0.9}}}`) + } + CppSceneStreamLastError = func(s uintptr) string { return "" } + CppSceneStreamFree = func(s uintptr) { order = append(order, "stream") } + CppSpeakerDim = func(uintptr) int32 { return 2 } + CppSpeakerRegistryNew = func() uintptr { return 9 } + CppSpeakerRegistryAddEmbedding = func(uintptr, string, *float32, int32) int32 { return 0 } + CppSpeakerRegistryFree = func(r uintptr) { regs = append(regs, r); order = append(order, "registry") } + }) + + AfterEach(func() { restore() }) + + It("begins a speaker scene stream and emits the slot name on the closed segment", func() { + in, out, errCh := runLive(p) + in <- liveVoicesConfig(ada) + in <- liveAudio(make([]float32, 10)) + close(in) + Expect(<-errCh).NotTo(HaveOccurred()) + + got := collectLive(out) + Expect(spkBeg).To(Equal(1)) + Expect(plain).To(Equal(0)) + Expect(gotReg).To(Equal(uintptr(9))) + var named *pb.LiveSpeakerSegment + for _, r := range got { + if len(r.Speakers) > 0 { + named = r.Speakers[0] + } + } + Expect(named).NotTo(BeNil()) + Expect(named.Speaker).To(Equal("0")) + Expect(named.Name).To(Equal("Ada")) + Expect(regs).To(Equal([]uintptr{9}), "the registry is freed once when the session ends") + Expect(order).To(Equal([]string{"stream", "registry"})) + }) + + It("uses the plain scene begin and emits empty names without known voices", func() { + CppSceneStreamFeedJSON = func(s uintptr, pcm *float32, n int32, isLast int32) uintptr { + return pool.cstr(`{"speakers":[{"speaker":0,"start":0.1,"end":0.6}],"sounds":[]}`) + } + in, out, errCh := runLive(p) + in <- liveConfig("") + in <- liveAudio(make([]float32, 10)) + close(in) + Expect(<-errCh).NotTo(HaveOccurred()) + + got := collectLive(out) + Expect(plain).To(Equal(1)) + Expect(spkBeg).To(Equal(0)) + Expect(regs).To(BeEmpty()) + found := false + for _, r := range got { + for _, s := range r.Speakers { + found = true + Expect(s.Name).To(BeEmpty()) + } + } + Expect(found).To(BeTrue()) + }) + + It("names the restarted scene stream after a config reset and frees every registry once", func() { + in, out, errCh := runLive(p) + in <- liveVoicesConfig(ada) + in <- liveAudio(make([]float32, 10)) + in <- liveVoicesConfig(ada) // reset + in <- liveAudio(make([]float32, 10)) + close(in) + Expect(<-errCh).NotTo(HaveOccurred()) + collectLive(out) + + Expect(spkBeg).To(Equal(2)) + Expect(regs).To(HaveLen(2), "one registry per scene stream, each freed exactly once") + }) + + It("frees the registry once when a feed failure disables the scene stream", func() { + CppSceneStreamFeedJSON = func(s uintptr, pcm *float32, n int32, isLast int32) uintptr { return 0 } + in, out, errCh := runLive(p) + in <- liveVoicesConfig(ada) + in <- liveAudio(make([]float32, 10)) + in <- liveAudio(make([]float32, 10)) + close(in) + Expect(<-errCh).NotTo(HaveOccurred()) + collectLive(out) + + Expect(regs).To(Equal([]uintptr{9})) + }) +}) + var _ = Describe("stripEouMarker", func() { It("strips a trailing and reports it", func() { text, eou := stripEouMarker("it is certainly very like the old portrait") diff --git a/backend/go/parakeet-cpp/main.go b/backend/go/parakeet-cpp/main.go index 865d5b6f1..8c7a3b7c3 100644 --- a/backend/go/parakeet-cpp/main.go +++ b/backend/go/parakeet-cpp/main.go @@ -117,7 +117,30 @@ func main() { purego.RegisterLibFunc(&CppSceneStreamLastError, lib, "parakeet_capi_scene_stream_last_error") purego.RegisterLibFunc(&CppSceneStreamFree, lib, "parakeet_capi_scene_stream_free") } + // Speaker identification (ABI v9 and v10). Probed separately from model_kind so an older + // libparakeet.so still loads; speaker_model: is refused in roles.go unless the v10 symbols exist. + if sym, err := purego.Dlsym(lib, "parakeet_capi_scene_stream_begin_speaker"); err == nil && sym != 0 { + purego.RegisterLibFunc(&CppSpeakerDim, lib, "parakeet_capi_speaker_dim") + purego.RegisterLibFunc(&CppSpeakerRegistryNew, lib, "parakeet_capi_speaker_registry_new") + purego.RegisterLibFunc(&CppSpeakerRegistryFree, lib, "parakeet_capi_speaker_registry_free") + purego.RegisterLibFunc(&CppSpeakerRegistryLastError, lib, "parakeet_capi_speaker_registry_last_error") + purego.RegisterLibFunc(&CppSceneStreamBeginSpeaker, lib, "parakeet_capi_scene_stream_begin_speaker") + purego.RegisterLibFunc(&CppTranscribeAndDiarizeNamedJSON, lib, "parakeet_capi_transcribe_and_diarize_named_json") + } + if sym, err := purego.Dlsym(lib, "parakeet_capi_diarize_named_pcm_json"); err == nil && sym != 0 { + purego.RegisterLibFunc(&CppSpeakerRegistryAddEmbedding, lib, "parakeet_capi_speaker_registry_add_embedding") + 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/roles.go b/backend/go/parakeet-cpp/roles.go index c6a1e912e..6bb0e8c27 100644 --- a/backend/go/parakeet-cpp/roles.go +++ b/backend/go/parakeet-cpp/roles.go @@ -17,6 +17,7 @@ const ( modelKindASR = 1 modelKindDiarization = 2 modelKindSound = 3 + modelKindSpeaker = 4 ) // Diarization streaming latency modes (mirrors PARAKEET_DIAR_LATENCY_* in @@ -38,6 +39,8 @@ func modelKindName(kind int32) string { return "diarization" case modelKindSound: return "sound" + case modelKindSpeaker: + return "speaker" default: return "unknown" } @@ -129,14 +132,30 @@ func (p *ParakeetCpp) loadRoles(opts *pb.ModelOptions) error { diarModelOpt := optString(opts, "diarization_model") asrModelOpt := optString(opts, "asr_model") soundModelOpt := optString(opts, "sound_model") - hasCompanionOpts := diarModelOpt != "" || asrModelOpt != "" || soundModelOpt != "" + speakerModelOpt := optString(opts, "speaker_model") + hasCompanionOpts := diarModelOpt != "" || asrModelOpt != "" || soundModelOpt != "" || speakerModelOpt != "" if hasCompanionOpts && CppModelKind == nil { - return errors.New("parakeet-cpp: asr_model/diarization_model/sound_model options need " + + return errors.New("parakeet-cpp: asr_model/diarization_model/sound_model/speaker_model options need " + "parakeet_capi_model_kind (ABI v8) to verify what they load; the loaded libparakeet.so " + "is too old to report companion model roles") } + if speakerModelOpt != "" { + if CppSpeakerRegistryAddEmbedding == nil || CppSpeakerDim == nil || CppSceneStreamBeginSpeaker == nil { + return errors.New("parakeet-cpp: speaker_model needs libparakeet.so ABI 10 " + + "(parakeet_capi_speaker_registry_add_embedding); the loaded library is older") + } + } + accept, err := parseSpeakerThreshold(optString(opts, "speaker_threshold")) + if err != nil { + return err + } + margin, err := parseSpeakerMargin(optString(opts, "speaker_margin")) + if err != nil { + return err + } + latency, err := parseDiarLatency(optString(opts, "diarization_latency")) if err != nil { return err @@ -160,7 +179,7 @@ func (p *ParakeetCpp) loadRoles(opts *pb.ModelOptions) error { for _, c := range loaded { CppFree(c) } - p.ctxPtr, p.diarCtx, p.tagCtx = 0, 0, 0 + p.ctxPtr, p.diarCtx, p.tagCtx, p.spkCtx = 0, 0, 0, 0 p.companions = nil } @@ -177,6 +196,10 @@ func (p *ParakeetCpp) loadRoles(opts *pb.ModelOptions) error { p.diarCtx = primary case modelKindSound: p.tagCtx = primary + case modelKindSpeaker: + freeLoaded() + return errors.New("parakeet-cpp: a speaker model cannot be the primary model; " + + "use it as speaker_model: next to a diarization model") default: p.ctxPtr = primary } @@ -191,6 +214,9 @@ func (p *ParakeetCpp) loadRoles(opts *pb.ModelOptions) error { {"sound_model", soundModelOpt, modelKindSound, func(pp *ParakeetCpp, c uintptr) { pp.tagCtx = c }, func(pp *ParakeetCpp) uintptr { return pp.tagCtx }}, + {"speaker_model", speakerModelOpt, modelKindSpeaker, + func(pp *ParakeetCpp, c uintptr) { pp.spkCtx = c }, + func(pp *ParakeetCpp) uintptr { return pp.spkCtx }}, } for _, spec := range specs { if spec.value == "" { @@ -198,7 +224,7 @@ func (p *ParakeetCpp) loadRoles(opts *pb.ModelOptions) error { } // A companion whose role the primary already occupies (e.g. asr_model: // on an already-ASR primary) would overwrite that role field below, - // leaking the primary ctx: Free() only walks ctxPtr/diarCtx/tagCtx, so + // leaking the primary ctx: Free() walks ctxPtr/diarCtx/tagCtx/spkCtx, so // the overwritten pointer is never freed. Reject it before loading. if spec.current(p) != 0 { freeLoaded() @@ -221,6 +247,11 @@ func (p *ParakeetCpp) loadRoles(opts *pb.ModelOptions) error { p.companions = append(p.companions, cctx) } + if p.spkCtx != 0 && p.diarCtx == 0 { + freeLoaded() + return errors.New("parakeet-cpp: speaker_model needs a diarization model (the primary or diarization_model:)") + } + p.speakerAccept, p.speakerMargin = accept, margin p.diarLatency = latency return nil } diff --git a/backend/go/parakeet-cpp/roles_test.go b/backend/go/parakeet-cpp/roles_test.go index 7b2a892b3..9b72087f0 100644 --- a/backend/go/parakeet-cpp/roles_test.go +++ b/backend/go/parakeet-cpp/roles_test.go @@ -185,6 +185,131 @@ var _ = Describe("model roles (stubbed C API)", func() { Expect(f.freed).To(HaveLen(freedBeforeFree), "Free after a failed Load must not free anything again") }) + Describe("speaker_model", func() { + var savedAdd func(uintptr, string, *float32, int32) int32 + var savedDim func(uintptr) int32 + var savedBegin func(asr, diar, tagger, speaker, reg uintptr, o *cSceneOpts) uintptr + var setNew func(on bool) + setNew = func(on bool) { + if on { + CppSpeakerRegistryAddEmbedding = func(uintptr, string, *float32, int32) int32 { return 0 } + CppSpeakerDim = func(uintptr) int32 { return 3 } + CppSceneStreamBeginSpeaker = func(_, _, _, _, _ uintptr, _ *cSceneOpts) uintptr { return 0 } + } else { + CppSpeakerRegistryAddEmbedding, CppSpeakerDim, CppSceneStreamBeginSpeaker = nil, nil, nil + } + } + BeforeEach(func() { + savedAdd, savedDim, savedBegin = CppSpeakerRegistryAddEmbedding, CppSpeakerDim, CppSceneStreamBeginSpeaker + setNew(true) + }) + AfterEach(func() { + CppSpeakerRegistryAddEmbedding, CppSpeakerDim, CppSceneStreamBeginSpeaker = savedAdd, savedDim, savedBegin + }) + + It("loads a kind-4 companion into spkCtx and Free releases it once", func() { + f := newFakeLib(). + withModel("diar.gguf", modelKindDiarization). + withModel("/models/spk.gguf", modelKindSpeaker) + restore = f.install() + + p := &ParakeetCpp{} + err := p.Load(&pb.ModelOptions{ + ModelFile: "diar.gguf", + ModelPath: "/models", + Options: []string{"speaker_model:spk.gguf", "speaker_threshold:0.3", "speaker_margin:0.1"}, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(p.spkCtx).ToNot(BeZero()) + Expect(p.speakerAccept).To(BeNumerically("~", 0.7, 1e-6)) + Expect(p.speakerMargin).To(BeNumerically("~", 0.1, 1e-6)) + + diar, spk := p.diarCtx, p.spkCtx + Expect(p.Free()).To(Succeed()) + Expect(f.freed).To(ConsistOf(diar, spk)) + Expect(p.spkCtx).To(BeZero()) + Expect(p.Free()).To(Succeed()) + Expect(f.freed).To(HaveLen(2), "a second Free frees nothing") + }) + + It("is rejected when the library lacks ABI 10, and nothing is loaded", func() { + setNew(false) + f := newFakeLib().withModel("diar.gguf", modelKindDiarization) + restore = f.install() + + p := &ParakeetCpp{} + err := p.Load(&pb.ModelOptions{ModelFile: "diar.gguf", Options: []string{"speaker_model:spk.gguf"}}) + Expect(err).To(MatchError(ContainSubstring("ABI 10"))) + Expect(f.loadedPaths).To(BeEmpty()) + }) + + It("is rejected without a diarization model and frees every context", func() { + f := newFakeLib(). + withModel("asr.gguf", modelKindASR). + withModel("spk.gguf", modelKindSpeaker) + restore = f.install() + + p := &ParakeetCpp{} + err := p.Load(&pb.ModelOptions{ModelFile: "asr.gguf", Options: []string{"speaker_model:spk.gguf"}}) + Expect(err).To(MatchError(ContainSubstring("diarization"))) + Expect(f.freed).To(HaveLen(2)) + Expect(p.ctxPtr).To(BeZero()) + Expect(p.spkCtx).To(BeZero()) + Expect(p.Free()).To(Succeed()) + Expect(f.freed).To(HaveLen(2)) + }) + + It("rejects a companion of the wrong kind", func() { + f := newFakeLib(). + withModel("diar.gguf", modelKindDiarization). + withModel("wrong.gguf", modelKindSound) + restore = f.install() + + p := &ParakeetCpp{} + err := p.Load(&pb.ModelOptions{ModelFile: "diar.gguf", Options: []string{"speaker_model:wrong.gguf"}}) + Expect(err).To(MatchError(ContainSubstring("is a sound model, expected a speaker model"))) + Expect(f.freed).To(HaveLen(2)) + Expect(p.spkCtx).To(BeZero()) + }) + + It("fails Load on an invalid speaker_threshold or speaker_margin", func() { + for _, bad := range []string{"speaker_threshold:abc", "speaker_threshold:2", "speaker_margin:1", "speaker_margin:-1"} { + f := newFakeLib().withModel("diar.gguf", modelKindDiarization) + r := f.install() + p := &ParakeetCpp{} + Expect(p.Load(&pb.ModelOptions{ModelFile: "diar.gguf", Options: []string{bad}})).ToNot(Succeed(), bad) + r() + } + }) + + It("rejects the speaker kind as the primary model and frees it", func() { + f := newFakeLib().withModel("spk.gguf", modelKindSpeaker) + restore = f.install() + + p := &ParakeetCpp{} + err := p.Load(&pb.ModelOptions{ModelFile: "spk.gguf"}) + Expect(err).To(MatchError(ContainSubstring("cannot be the primary model"))) + Expect(f.freed).To(HaveLen(1)) + Expect(p.spkCtx).To(BeZero()) + Expect(p.ctxPtr).To(BeZero()) + }) + + It("loads exactly as before without speaker_model on a library with none of the new symbols", func() { + setNew(false) + f := newFakeLib(). + withModel("asr.gguf", modelKindASR). + withModel("diar.gguf", modelKindDiarization) + restore = f.install() + + p := &ParakeetCpp{} + Expect(p.Load(&pb.ModelOptions{ModelFile: "asr.gguf", Options: []string{"diarization_model:diar.gguf"}})).To(Succeed()) + Expect(p.spkCtx).To(BeZero()) + Expect(p.ctxPtr).ToNot(BeZero()) + Expect(p.diarCtx).ToNot(BeZero()) + Expect(p.speakerAccept).To(BeNumerically("~", 0.5, 1e-6)) + }) + }) + It("rejects an asr_model companion on an already-ASR primary and frees everything it opened", func() { f := newFakeLib(). withModel("asr.gguf", modelKindASR). diff --git a/backend/go/parakeet-cpp/scene.go b/backend/go/parakeet-cpp/scene.go index 590740f13..099b968e2 100644 --- a/backend/go/parakeet-cpp/scene.go +++ b/backend/go/parakeet-cpp/scene.go @@ -39,6 +39,11 @@ type sceneSoundJSON struct { type sceneFeedJSON struct { Speakers []sceneSpeakerJSON `json:"speakers"` Sounds []sceneSoundJSON `json:"sounds"` + // Names is the CURRENT name of each speaker slot (keyed by the slot + // index as a string) at the time of this feed. Each closed segment takes + // its slot's current name, so a segment that closes before its slot is + // identified carries an empty name. Absent without a speaker model. + Names map[string]speakerNameJSON `json:"names"` } // sceneWanted reports whether AudioTranscriptionLive should run a companion @@ -62,10 +67,17 @@ func (p *ParakeetCpp) sceneWanted() bool { // caught instead of handed to the C side — mirroring streamFeedDoc's re-check // of p.ctxPtr (see the "Per-C-call engine serialization" comment in // goparakeetcpp.go). The zero value (s == 0) means "no scene stream". +// +// spk and reg are the speaker model and the known-voice registry of a +// 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 + 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 @@ -74,7 +86,7 @@ type sceneStreamHandle struct { // contexts 0 (defensive: sceneWanted() already guards this). A zero handle // means the C call itself failed; the caller logs a warning and continues // the live session without speaker/sound events. -func (p *ParakeetCpp) sceneBegin() sceneStreamHandle { +func (p *ParakeetCpp) sceneBegin(voices []*pb.KnownVoice) sceneStreamHandle { p.engineMu.Lock() defer p.engineMu.Unlock() diar, tag := p.diarCtx, p.tagCtx @@ -90,6 +102,29 @@ func (p *ParakeetCpp) sceneBegin() sceneStreamHandle { // whole lifetime. 0 disables per-class score retention; sound EVENTS // (onset/offset, what the live path actually consumes) are unaffected. opts.Sound.TopK = 0 + + // With a speaker model and at least one usable known voice the stream + // names speakers through a registry it borrows. engineMu is already + // held, so use the Locked builder. Any failure keeps the plain stream. + var reg uintptr + if diar != 0 && p.spkCtx != 0 && CppSceneStreamBeginSpeaker != nil { + r, err := p.buildSpeakerRegistryLocked(voices) + if err != nil { + xlog.Warn("parakeet-cpp: could not build the speaker registry for a live session; speakers stay unnamed", "err", err) + } else { + reg = r + } + } + if reg != 0 { + opts.SpeakerAcceptThreshold = p.speakerAccept + opts.SpeakerMargin = p.speakerMargin + s := CppSceneStreamBeginSpeaker(0, diar, tag, p.spkCtx, reg, &opts) + if s == 0 { + p.freeSpeakerRegistry(reg) + return sceneStreamHandle{} + } + 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 { return sceneStreamHandle{} @@ -111,6 +146,8 @@ func (p *ParakeetCpp) sceneFree(h sceneStreamHandle) { p.engineMu.Lock() defer p.engineMu.Unlock() CppSceneStreamFree(h.s) + // The stream borrowed the registry: free it only after the stream. + p.freeSpeakerRegistry(h.reg) } // sceneFeed runs one scene-stream feed (or the is_last flush) under @@ -126,7 +163,9 @@ func (p *ParakeetCpp) sceneFeed(h sceneStreamHandle, pcm []float32, isLast bool) p.engineMu.Lock() defer p.engineMu.Unlock() - if p.diarCtx != h.diar || p.tagCtx != h.tag { + if p.diarCtx != h.diar || p.tagCtx != h.tag || (h.spk != 0 && p.spkCtx != h.spk) { + // A plain stream (h.spk == 0) never borrows the speaker model, so a + // speaker model loaded or freed meanwhile does not concern it. return sceneFeedJSON{}, grpcerrors.ModelNotLoaded("parakeet-cpp") } @@ -155,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 } @@ -223,14 +263,19 @@ func (p *ParakeetCpp) feedSlicesScene(ctx context.Context, stream uintptr, scene // TranscriptLiveResponse.speakers (stream-relative nanoseconds). Reuses // diarize.go's speakerLabel so the live path renders speaker indices the // same way the offline Diarize RPC does. -func liveSpeakersToProto(speakers []sceneSpeakerJSON) []*pb.LiveSpeakerSegment { +// +// names is the feed document's "names" map; each segment takes its slot's +// current name, empty if the slot was not yet identified when it closed. +func liveSpeakersToProto(speakers []sceneSpeakerJSON, names map[string]speakerNameJSON) []*pb.LiveSpeakerSegment { if len(speakers) == 0 { return nil } out := make([]*pb.LiveSpeakerSegment, len(speakers)) for i, s := range speakers { + name, _ := nameFor(names, s.Speaker) out[i] = &pb.LiveSpeakerSegment{ Speaker: speakerLabel(s.Speaker), + Name: name, Start: secondsToNanos(s.Start), End: secondsToNanos(s.End), } diff --git a/backend/go/parakeet-cpp/scene_test.go b/backend/go/parakeet-cpp/scene_test.go index 9dbfa2205..ca4005707 100644 --- a/backend/go/parakeet-cpp/scene_test.go +++ b/backend/go/parakeet-cpp/scene_test.go @@ -1,6 +1,9 @@ package main import ( + "unsafe" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) @@ -31,7 +34,7 @@ var _ = Describe("ParakeetCpp.sceneBegin", func() { } p := &ParakeetCpp{diarCtx: 42, diarLatency: diarLatencyVeryLow} - h := p.sceneBegin() + h := p.sceneBegin(nil) Expect(h.s).ToNot(BeZero()) Expect(gotOpts.Sound.TopK).To(Equal(int32(0)), "the live scene path never drains sound scores (see sound.go's SoundDetection, "+ @@ -40,3 +43,123 @@ var _ = Describe("ParakeetCpp.sceneBegin", func() { Expect(gotOpts.DiarLatency).To(Equal(diarLatencyVeryLow)) }) }) + +var _ = Describe("scene stream with speaker names", func() { + var ( + restore func() + pool *liveCstrPool + freedRegs []uintptr + ) + BeforeEach(func() { + sOpts, sBegin, sBeginSpk, sFeed, sFree := CppSceneOptsDefault, CppSceneStreamBegin, CppSceneStreamBeginSpeaker, CppSceneStreamFeedJSON, CppSceneStreamFree + sNew, sRegFree, sAdd, sDim, sStr := CppSpeakerRegistryNew, CppSpeakerRegistryFree, CppSpeakerRegistryAddEmbedding, CppSpeakerDim, CppFreeString + restore = func() { + CppSceneOptsDefault, CppSceneStreamBegin, CppSceneStreamBeginSpeaker, CppSceneStreamFeedJSON, CppSceneStreamFree = sOpts, sBegin, sBeginSpk, sFeed, sFree + CppSpeakerRegistryNew, CppSpeakerRegistryFree, CppSpeakerRegistryAddEmbedding, CppSpeakerDim, CppFreeString = sNew, sRegFree, sAdd, sDim, sStr + } + pool = &liveCstrPool{} + freedRegs = nil + CppSceneOptsDefault = func(o *cSceneOpts) { o.Size = int32(unsafe.Sizeof(*o)) } + CppSpeakerDim = func(uintptr) int32 { return 2 } + CppSpeakerRegistryNew = func() uintptr { return 9 } + CppSpeakerRegistryFree = func(r uintptr) { freedRegs = append(freedRegs, r) } + CppSpeakerRegistryAddEmbedding = func(uintptr, string, *float32, int32) int32 { return 0 } + CppFreeString = func(uintptr) {} + }) + 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 + CppSceneStreamBeginSpeaker = func(asr, diar, tag, spk, reg uintptr, o *cSceneOpts) uintptr { + gotOpts, gotReg, gotSpk = *o, reg, spk + return 55 + } + CppSceneStreamBegin = func(asr, diar, tag uintptr, o *cSceneOpts) uintptr { + Fail("plain begin must not be used") + return 0 + } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2, speakerAccept: 0.7, speakerMargin: 0.05} + h := p.sceneBegin([]*pb.KnownVoice{{Name: "Ada", Embedding: []float32{1, 0}}}) + Expect(h.s).To(Equal(uintptr(55))) + Expect(h.reg).To(Equal(uintptr(9))) + Expect(gotReg).To(Equal(uintptr(9))) + Expect(gotSpk).To(Equal(uintptr(2))) + Expect(gotOpts.SpeakerAcceptThreshold).To(BeNumerically("~", 0.7, 1e-6)) + Expect(gotOpts.SpeakerMargin).To(BeNumerically("~", 0.05, 1e-6)) + }) + + It("uses the plain scene begin without voices or without a speaker model", func() { + plain := 0 + CppSceneStreamBegin = func(asr, diar, tag uintptr, o *cSceneOpts) uintptr { plain++; return 56 } + CppSceneStreamBeginSpeaker = func(asr, diar, tag, spk, reg uintptr, o *cSceneOpts) uintptr { + Fail("speaker begin must not be used") + return 0 + } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + Expect(p.sceneBegin(nil).s).To(Equal(uintptr(56))) + p2 := &ParakeetCpp{diarCtx: 1} + Expect(p2.sceneBegin([]*pb.KnownVoice{{Name: "Ada", Embedding: []float32{1, 0}}}).s).To(Equal(uintptr(56))) + Expect(plain).To(Equal(2)) + }) + + It("frees the registry when the speaker begin fails, and degrades to no scene stream", func() { + CppSceneStreamBeginSpeaker = func(asr, diar, tag, spk, reg uintptr, o *cSceneOpts) uintptr { return 0 } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + h := p.sceneBegin([]*pb.KnownVoice{{Name: "Ada", Embedding: []float32{1, 0}}}) + Expect(h.s).To(Equal(uintptr(0))) + Expect(freedRegs).To(Equal([]uintptr{9})) + }) + + It("frees the stream before its registry", func() { + var order []string + CppSceneStreamFree = func(uintptr) { order = append(order, "stream") } + CppSpeakerRegistryFree = func(uintptr) { order = append(order, "registry") } + (&ParakeetCpp{}).sceneFree(sceneStreamHandle{s: 55, reg: 9}) + Expect(order).To(Equal([]string{"stream", "registry"})) + }) + + It("turns the names of the feed document into live speaker segment names", func() { + segs := liveSpeakersToProto( + []sceneSpeakerJSON{{Speaker: 0, Start: 0, End: 1}, {Speaker: 1, Start: 1, End: 2}}, + map[string]speakerNameJSON{"0": {Name: "Ada", Score: 0.9}, "1": {Name: ""}}) + Expect(segs[0].Name).To(Equal("Ada")) + Expect(segs[0].Speaker).To(Equal("0")) + Expect(segs[1].Name).To(BeEmpty()) + Expect(liveSpeakersToProto([]sceneSpeakerJSON{{Speaker: 0}}, nil)[0].Name).To(BeEmpty()) + }) + + It("decodes the names map of a scene feed", func() { + CppSceneStreamFeedJSON = func(s uintptr, pcm *float32, n, last int32) uintptr { + return pool.cstr(`{"speakers":[{"speaker":0,"start":0.0,"end":1.0}],"sounds":[],"names":{"0":{"name":"Ada","score":0.9}}}`) + } + p := &ParakeetCpp{diarCtx: 1, spkCtx: 2} + doc, err := p.sceneFeed(sceneStreamHandle{s: 55, diar: 1, spk: 2}, []float32{0}, false) + Expect(err).ToNot(HaveOccurred()) + Expect(doc.Names["0"].Name).To(Equal("Ada")) + }) + + It("refuses to feed a stream whose speaker model was freed", func() { + p := &ParakeetCpp{diarCtx: 1} + _, err := p.sceneFeed(sceneStreamHandle{s: 55, diar: 1, spk: 2}, []float32{0}, false) + Expect(err).To(HaveOccurred()) + }) +}) diff --git a/backend/go/parakeet-cpp/speaker_real_test.go b/backend/go/parakeet-cpp/speaker_real_test.go new file mode 100644 index 000000000..c652ca15c --- /dev/null +++ b/backend/go/parakeet-cpp/speaker_real_test.go @@ -0,0 +1,184 @@ +package main + +import ( + "encoding/json" + "os" + "sort" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// fixtureVoices reads testdata/two_speakers_embeddings.json: WeSpeaker ResNet34 +// embeddings of voice A (two_speakers.wav 0.6-4.6 s) and voice B (6.9-10.9 s). +func fixtureVoices() (a, b *pb.KnownVoice) { + raw, err := os.ReadFile("testdata/two_speakers_embeddings.json") + Expect(err).ToNot(HaveOccurred()) + var doc struct { + Dim int `json:"dim"` + Voices []struct { + Name string `json:"name"` + Embedding []float32 `json:"embedding"` + } `json:"voices"` + } + Expect(json.Unmarshal(raw, &doc)).To(Succeed()) + Expect(doc.Voices).To(HaveLen(2)) + Expect(doc.Voices[0].Embedding).To(HaveLen(doc.Dim)) + mk := func(i int) *pb.KnownVoice { + return &pb.KnownVoice{Name: doc.Voices[i].Name, Embedding: doc.Voices[i].Embedding} + } + return mk(0), mk(1) +} + +func speakerFixturesOrSkip() (diarModel, speakerModel, wav string) { + diarModel = os.Getenv("PARAKEET_BACKEND_TEST_DIAR_MODEL") + speakerModel = os.Getenv("PARAKEET_BACKEND_TEST_SPEAKER_MODEL") + wav = os.Getenv("PARAKEET_BACKEND_TEST_WAV") + if diarModel == "" || speakerModel == "" || wav == "" { + Skip("set PARAKEET_BACKEND_TEST_DIAR_MODEL, PARAKEET_BACKEND_TEST_SPEAKER_MODEL and " + + "PARAKEET_BACKEND_TEST_WAV (parakeet.cpp tests/fixtures/two_speakers.wav)") + } + ensureLibLoaded() + if CppDiarizeNamedPCMJSON == nil { + Skip("libparakeet.so has no ABI 10 speaker naming (parakeet_capi_diarize_named_pcm_json)") + } + return +} + +// namesBySlot maps a speaker label to the set of names its segments carry. +func namesBySlot(segs []*pb.DiarizeSegment) map[string][]string { + out := map[string][]string{} + for _, s := range segs { + out[s.Speaker] = append(out[s.Speaker], s.Name) + } + return out +} + +var _ = Describe("ParakeetCpp speaker names (real libparakeet.so, ABI 10)", func() { + load := func(diarModel, speakerModel string, extra ...string) *ParakeetCpp { + p := &ParakeetCpp{} + opts := append([]string{"speaker_model:" + speakerModel}, extra...) + Expect(p.Load(&pb.ModelOptions{ModelFile: diarModel, Options: opts})).To(Succeed()) + return p + } + + // two_speakers.wav: voice A speaks 0.5-5.5 s and 14.8-18.7 s, voice B 6.9-13.5 s and + // 20.1-23.6 s; diarization slot 0 is voice A and slot 1 is voice B. + It("names diarized segments from registered voices, whatever order they arrive in", func() { + diarModel, speakerModel, wav := speakerFixturesOrSkip() + a, b := fixtureVoices() + p := load(diarModel, speakerModel) + defer func() { _ = p.Free() }() + + // Reversed on purpose: naming by arrival order would swap the speakers. + res, err := p.Diarize(&pb.DiarizeRequest{Dst: wav, KnownVoices: []*pb.KnownVoice{b, a}}) + Expect(err).ToNot(HaveOccurred()) + Expect(res.Segments).To(HaveLen(5)) + for _, s := range res.Segments { + GinkgoWriter.Printf("[%5.1f-%5.1f] slot %s name %q score %.3f\n", s.Start, s.End, s.Speaker, s.Name, s.NameScore) + want := map[string]string{"0": "voice_a", "1": "voice_b"}[s.Speaker] + Expect(want).ToNot(BeEmpty(), "unexpected slot %q", s.Speaker) + Expect(s.Name).To(Equal(want)) + Expect(s.NameScore).To(BeNumerically(">", 0.5)) + } + }) + + It("leaves a slot unnamed when its voice is not registered", func() { + diarModel, speakerModel, wav := speakerFixturesOrSkip() + _, b := fixtureVoices() + // Distance 0.3 means cosine 0.7; a different speaker scores near 0 with WeSpeaker. + p := load(diarModel, speakerModel, "speaker_threshold:0.3") + defer func() { _ = p.Free() }() + + res, err := p.Diarize(&pb.DiarizeRequest{Dst: wav, KnownVoices: []*pb.KnownVoice{b}}) + Expect(err).ToNot(HaveOccurred()) + names := namesBySlot(res.Segments) + Expect(names).To(HaveKey("0")) + Expect(names).To(HaveKey("1")) + for _, n := range names["0"] { + Expect(n).To(BeEmpty()) + } + for _, n := range names["1"] { + Expect(n).To(Equal("voice_b")) + } + }) + + // A distance of 0.01 asks for cosine 0.99, above any genuine score. If purego did not + // hand the float32 to C, the C side would read 0 and fall back to its 0.5 default, + // which names both slots, so an unnamed result proves the argument arrives. + It("passes the float32 accept threshold through purego (a 0.99 cosine names nobody)", func() { + diarModel, speakerModel, wav := speakerFixturesOrSkip() + a, b := fixtureVoices() + p := load(diarModel, speakerModel, "speaker_threshold:0.01") + defer func() { _ = p.Free() }() + + res, err := p.Diarize(&pb.DiarizeRequest{Dst: wav, KnownVoices: []*pb.KnownVoice{a, b}}) + Expect(err).ToNot(HaveOccurred()) + Expect(res.Segments).ToNot(BeEmpty()) + for _, s := range res.Segments { + Expect(s.Name).To(BeEmpty(), "slot %s was named at cosine 0.99", s.Speaker) + } + }) + + It("names live speaker segments from the registered voices", func() { + diarModel, speakerModel, wav := speakerFixturesOrSkip() + streamModel := os.Getenv("PARAKEET_BACKEND_TEST_STREAM_MODEL") + if streamModel == "" { + Skip("set PARAKEET_BACKEND_TEST_STREAM_MODEL (cache-aware streaming model) for the live path") + } + if CppSceneStreamBeginSpeaker == nil { + Skip("libparakeet.so has no scene_stream_begin_speaker") + } + a, b := fixtureVoices() + p := &ParakeetCpp{} + Expect(p.Load(&pb.ModelOptions{ + ModelFile: streamModel, + Options: []string{"diarization_model:" + diarModel, "speaker_model:" + speakerModel}, + })).To(Succeed()) + defer func() { _ = p.Free() }() + + pcm, _, err := decodeWavMono16k(wav) + Expect(err).ToNot(HaveOccurred()) + + in := make(chan *pb.TranscriptLiveRequest, 8) + out := make(chan *pb.TranscriptLiveResponse, 256) + errCh := make(chan error, 1) + go func() { errCh <- p.AudioTranscriptionLive(in, out) }() + in <- &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Config{ + Config: &pb.TranscriptLiveConfig{KnownVoices: []*pb.KnownVoice{b, a}}, + }} + go func() { + const chunk = 8000 // 0.5 s + for i := 0; i < len(pcm); i += chunk { + end := min(i+chunk, len(pcm)) + in <- liveAudio(pcm[i:end]) + } + close(in) + }() + + var segs []*pb.LiveSpeakerSegment + for r := range out { + segs = append(segs, r.GetSpeakers()...) + } + Expect(<-errCh).ToNot(HaveOccurred()) + + named := 0 + seen := map[string]bool{} + for _, s := range segs { + GinkgoWriter.Printf("live slot %s [%d-%d] name %q\n", s.Speaker, s.Start, s.End, s.Name) + seen[s.Speaker] = true + if s.Name == "" { + continue + } + named++ + Expect(s.Name).To(Equal(map[string]string{"0": "voice_a", "1": "voice_b"}[s.Speaker])) + } + keys := make([]string, 0, len(seen)) + for k := range seen { + keys = append(keys, k) + } + sort.Strings(keys) + Expect(named).To(BeNumerically(">", 0), "no live speaker segment was named; slots seen: %v", keys) + }) +}) diff --git a/backend/go/parakeet-cpp/speaker_registry.go b/backend/go/parakeet-cpp/speaker_registry.go new file mode 100644 index 000000000..2d56ebddf --- /dev/null +++ b/backend/go/parakeet-cpp/speaker_registry.go @@ -0,0 +1,143 @@ +package main + +import ( + "fmt" + "math" + "strconv" + "strings" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +const ( + defaultSpeakerDistance = 0.5 // cosine 0.5, parakeet.cpp's default accept_threshold + defaultSpeakerMargin = 0.05 // parakeet.cpp's default runner-up margin +) + +// parseSpeakerThreshold reads speaker_threshold, a distance (1 minus cosine, the unit +// /v1/voice/identify uses), and returns the cosine acceptance threshold the C-API takes. +// Empty means the default. A distance outside (0, 2) is an error. +func parseSpeakerThreshold(s string) (float32, error) { + d := defaultSpeakerDistance + if strings.TrimSpace(s) != "" { + v, err := strconv.ParseFloat(strings.TrimSpace(s), 64) + if err != nil || math.IsNaN(v) || v <= 0 || v >= 2 { + return 0, fmt.Errorf("parakeet-cpp: speaker_threshold %q must be a distance in (0, 2) (1 minus cosine similarity)", s) + } + d = v + } + return nonZero(float32(1 - d)), nil +} + +// minPositive stands in for an exact 0 threshold or margin. The C side reads 0 as +// "use the default", so a distance of 1 (cosine 0) or a margin of 0 would silently +// become 0.5 or 0.05. +const minPositive = float32(1e-6) + +func nonZero(v float32) float32 { + if v == 0 { + return minPositive + } + return v +} + +// parseSpeakerMargin reads speaker_margin, the runner-up margin in [0, 1). +func parseSpeakerMargin(s string) (float32, error) { + if strings.TrimSpace(s) == "" { + return defaultSpeakerMargin, nil + } + v, err := strconv.ParseFloat(strings.TrimSpace(s), 64) + if err != nil || math.IsNaN(v) || v < 0 || v >= 1 { + return 0, fmt.Errorf("parakeet-cpp: speaker_margin %q must be a number in [0, 1)", s) + } + return nonZero(float32(v)), nil +} + +// buildSpeakerRegistryLocked makes a parakeet_speaker_registry from the registered voices +// of one request or stream. Caller holds engineMu. It returns 0 (and no error) when there is +// nothing to build: no speaker model loaded or no usable voices. A voice whose embedding size +// differs from the speaker model's, or that the C side refuses, is skipped with a warning +// (without its name: the log is not for the caller who may not see voice names) so one bad +// voice cannot fail every request. The caller frees a non-zero result with freeSpeakerRegistry. +func (p *ParakeetCpp) buildSpeakerRegistryLocked(voices []*pb.KnownVoice) (uintptr, error) { + if p.spkCtx == 0 || CppSpeakerRegistryNew == nil || CppSpeakerRegistryAddEmbedding == nil || len(voices) == 0 { + return 0, nil + } + dim := 0 + if CppSpeakerDim != nil { + dim = int(CppSpeakerDim(p.spkCtx)) + } + reg := CppSpeakerRegistryNew() + if reg == 0 { + return 0, status.Error(codes.Internal, "parakeet-cpp: could not create a speaker registry") + } + added, skipped := 0, 0 + for _, v := range voices { + emb := v.GetEmbedding() + if v.GetName() == "" || len(emb) == 0 { + xlog.Warn("parakeet-cpp: skipping a known voice with no name or embedding") + continue + } + if dim > 0 && len(emb) != dim { + xlog.Warn("parakeet-cpp: skipped a registered voice: embedding size does not match the speaker model's", + "voice_size", len(emb), "speaker_model_size", dim) + skipped++ + continue + } + 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 + } + added++ + } + if added == 0 { + if skipped > 0 { + xlog.Warn("parakeet-cpp: no registered voice is usable with this speaker model; speakers stay unnamed", "skipped", skipped) + } + CppSpeakerRegistryFree(reg) + return 0, nil + } + return reg, nil +} + +// buildSpeakerRegistry is buildSpeakerRegistryLocked under engineMu. +func (p *ParakeetCpp) buildSpeakerRegistry(voices []*pb.KnownVoice) (uintptr, error) { + p.engineMu.Lock() + defer p.engineMu.Unlock() + return p.buildSpeakerRegistryLocked(voices) +} + +// freeSpeakerRegistry releases a registry built above. A zero handle is a no-op. +func (p *ParakeetCpp) freeSpeakerRegistry(reg uintptr) { + if reg == 0 || CppSpeakerRegistryFree == nil { + return + } + 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 new file mode 100644 index 000000000..4297b2978 --- /dev/null +++ b/backend/go/parakeet-cpp/speaker_registry_test.go @@ -0,0 +1,134 @@ +package main + +import ( + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("speaker options", func() { + It("turns a distance threshold into the cosine the C side takes", func() { + a, err := parseSpeakerThreshold("") + Expect(err).ToNot(HaveOccurred()) + Expect(a).To(BeNumerically("~", 0.5, 1e-6)) // default distance 0.5 + a, err = parseSpeakerThreshold("0.3") + Expect(err).ToNot(HaveOccurred()) + Expect(a).To(BeNumerically("~", 0.7, 1e-6)) + }) + It("rejects a threshold that is not a distance in (0, 2)", func() { + for _, bad := range []string{"abc", "0", "-0.1", "2", "2.5", "NaN"} { + _, err := parseSpeakerThreshold(bad) + Expect(err).To(HaveOccurred(), bad) + } + }) + It("never hands the C side an exact zero, which it reads as use the default", func() { + a, err := parseSpeakerThreshold("1") // distance 1 is cosine 0 + Expect(err).ToNot(HaveOccurred()) + Expect(a).To(Equal(float32(1e-6))) + m, err := parseSpeakerMargin("0") + Expect(err).ToNot(HaveOccurred()) + Expect(m).To(Equal(float32(1e-6))) + a, _ = parseSpeakerThreshold("") + Expect(a).To(BeNumerically("~", 0.5, 1e-6)) + m, _ = parseSpeakerMargin("") + Expect(m).To(BeNumerically("~", 0.05, 1e-6)) + }) + It("parses the margin, default 0.05, within [0, 1)", func() { + m, err := parseSpeakerMargin("") + Expect(err).ToNot(HaveOccurred()) + Expect(m).To(BeNumerically("~", 0.05, 1e-6)) + _, err = parseSpeakerMargin("-1") + Expect(err).To(HaveOccurred()) + _, err = parseSpeakerMargin("1") + Expect(err).To(HaveOccurred()) + }) +}) + +var _ = Describe("buildSpeakerRegistry", func() { + var restore func() + var added []string + var freed []uintptr + BeforeEach(func() { + sNew, sFree, sAdd, sDim, sErr := CppSpeakerRegistryNew, CppSpeakerRegistryFree, CppSpeakerRegistryAddEmbedding, CppSpeakerDim, CppSpeakerRegistryLastError + restore = func() { + CppSpeakerRegistryNew, CppSpeakerRegistryFree, CppSpeakerRegistryAddEmbedding, CppSpeakerDim, CppSpeakerRegistryLastError = sNew, sFree, sAdd, sDim, sErr + } + added, freed = nil, nil + CppSpeakerDim = func(uintptr) int32 { return 3 } + CppSpeakerRegistryNew = func() uintptr { return 77 } + CppSpeakerRegistryFree = func(r uintptr) { freed = append(freed, r) } + CppSpeakerRegistryAddEmbedding = func(r uintptr, name string, emb *float32, dim int32) int32 { + added = append(added, name) + return 0 + } + CppSpeakerRegistryLastError = func(uintptr) string { return "stub error" } + }) + AfterEach(func() { restore() }) + + voice := func(name string, n int) *pb.KnownVoice { + 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{ + {Name: "ada", Embedding: []float32{1, 0, 0}}, {Name: "ben", Embedding: []float32{0, 1, 0}}, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(reg).To(Equal(uintptr(77))) + Expect(added).To(Equal([]string{"ada", "ben"})) + Expect(freed).To(BeEmpty()) + }) + It("skips a voice of the wrong size without failing, and frees the registry when none is left", func() { + p := &ParakeetCpp{spkCtx: 5} + reg, err := p.buildSpeakerRegistry([]*pb.KnownVoice{voice("cy", 5)}) + Expect(err).ToNot(HaveOccurred()) + Expect(reg).To(Equal(uintptr(0))) + Expect(added).To(BeEmpty()) + Expect(freed).To(Equal([]uintptr{77})) + }) + It("keeps the usable voices when another one has the wrong size", func() { + p := &ParakeetCpp{spkCtx: 5} + reg, err := p.buildSpeakerRegistry([]*pb.KnownVoice{voice("cy", 5), {Name: "ada", Embedding: []float32{1, 0, 0}}}) + Expect(err).ToNot(HaveOccurred()) + Expect(reg).To(Equal(uintptr(77))) + Expect(added).To(Equal([]string{"ada"})) + Expect(freed).To(BeEmpty()) + }) + It("skips a voice the C side refuses, and frees the registry when none is left", func() { + CppSpeakerRegistryAddEmbedding = func(uintptr, string, *float32, int32) int32 { return 1 } + p := &ParakeetCpp{spkCtx: 5} + reg, err := p.buildSpeakerRegistry([]*pb.KnownVoice{{Name: "ada", Embedding: []float32{0, 0, 0}}}) + Expect(err).ToNot(HaveOccurred()) + Expect(reg).To(Equal(uintptr(0))) + Expect(freed).To(Equal([]uintptr{77})) + }) + It("returns no registry when no speaker model is loaded", func() { + p := &ParakeetCpp{} + reg, err := p.buildSpeakerRegistry([]*pb.KnownVoice{voice("ada", 3)}) + Expect(err).ToNot(HaveOccurred()) + Expect(reg).To(Equal(uintptr(0))) + }) + It("skips a voice with no name or embedding instead of failing the request", func() { + p := &ParakeetCpp{spkCtx: 5} + _, err := p.buildSpeakerRegistry([]*pb.KnownVoice{{Name: "", Embedding: []float32{1, 0, 0}}, {Name: "x"}, {Name: "ada", Embedding: []float32{1, 0, 0}}}) + Expect(err).ToNot(HaveOccurred()) + Expect(added).To(Equal([]string{"ada"})) + }) + It("frees a registry and ignores a zero handle", func() { + p := &ParakeetCpp{} + p.freeSpeakerRegistry(0) + p.freeSpeakerRegistry(9) + Expect(freed).To(Equal([]uintptr{9})) + }) +}) diff --git a/backend/go/parakeet-cpp/speakers.go b/backend/go/parakeet-cpp/speakers.go index ddd2c4292..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) + raw, err := p.diarizeCall(pcm, false, 0, false) if err != nil { return nil, err } diff --git a/backend/go/parakeet-cpp/testdata/two_speakers_embeddings.json b/backend/go/parakeet-cpp/testdata/two_speakers_embeddings.json new file mode 100644 index 000000000..6288242cb --- /dev/null +++ b/backend/go/parakeet-cpp/testdata/two_speakers_embeddings.json @@ -0,0 +1 @@ +{"model":"voice-detect-wespeaker-resnet34.gguf","dim":256,"voices":[{"name":"voice_a","clip":"two_speakers.wav 0.6-4.6 s","embedding":[0.011315,0.055604,-0.071566,0.00961,0.040386,-0.016421,-0.033611,-0.019875,0.012314,-0.048251,0.025119,-0.035261,-0.051,-0.065269,0.044619,-0.032789,-0.077872,-0.059979,0.077165,0.034743,-0.034322,0.051458,0.119578,-0.060482,-0.032685,-0.066351,-0.043352,0.014827,-0.008393,0.079046,0.075357,0.177104,-0.002038,0.020706,-0.018464,0.004545,0.026542,0.00754,0.126408,0.011555,-0.083333,0.046795,0.029405,0.039988,-0.051861,0.034671,0.063438,0.010646,0.089453,0.009784,0.161807,-0.060077,0.080398,0.0026,-0.179952,-0.092161,0.008115,-0.043436,0.032408,-0.009307,0.04841,0.029929,-0.066182,0.069376,0.017271,0.076435,-0.013476,-0.091337,0.027061,0.070812,0.039052,-0.025303,0.062032,-0.129511,-0.035932,0.123741,0.059323,0.084191,0.017091,-0.071467,0.00336,0.109306,-0.032808,0.071844,0.013484,0.01657,-0.04768,-0.087006,0.07233,-0.004725,-0.058554,-0.061997,0.088107,0.069736,0.008408,-0.104471,0.049595,0.070141,0.071505,0.03694,0.030285,-0.103312,-0.028439,0.053189,0.132217,0.061057,-0.115166,0.04954,-0.04504,-0.044936,0.041552,0.07976,0.005178,0.027334,0.036235,-0.096866,-0.087907,-0.095822,0.10443,0.041865,-0.034412,0.065312,0.008648,-0.075287,0.125829,0.046572,-0.006843,0.106429,0.024523,0.013693,0.011473,-0.053549,0.078732,0.094865,-0.06333,0.03013,-0.090245,-0.013319,-0.090633,0.04209,-0.012844,-0.055208,-0.096489,-0.093557,0.038501,0.033988,0.032093,-0.015028,-0.067271,-0.100653,-0.00632,0.055982,-0.026291,-0.088451,0.019825,0.18585,-0.013341,-0.016392,0.05616,0.022352,0.006584,-0.097643,-0.02696,-0.019335,0.089462,-0.082992,0.02947,-0.014762,-0.055757,-0.011026,0.046157,-0.01546,0.029421,0.005409,-0.080892,-0.009961,0.029963,0.013844,-0.021823,-0.036825,0.001308,0.044549,0.011387,0.058622,0.047924,-0.034883,-0.038361,0.092906,0.006849,-0.076708,0.080052,0.011746,0.168146,-0.017433,-0.017833,0.010594,0.05945,-0.034766,-0.008519,0.02914,0.062197,-0.032722,-0.029624,0.103967,-0.07483,-0.052394,0.022729,-0.036227,0.060867,0.144592,-0.054725,-0.009984,0.056824,-0.076816,-0.005553,-0.054548,-0.047348,0.011424,0.091265,-0.012273,0.013987,-0.019944,0.054455,0.050702,0.004049,0.013685,0.022356,0.003249,-0.081608,-0.06926,-0.077244,-0.048076,0.059592,0.009222,0.007737,0.014863,-0.070568,-0.070963,0.056541,-0.067229,-0.034265,0.072919,0.032005,-0.06045,-0.148719,0.050373,-0.06677,0.020233,-0.031594,-0.084966,0.033239,-0.020293,-0.049292,0.120206,-0.056133,0.090007]},{"name":"voice_b","clip":"two_speakers.wav 6.9-10.9 s","embedding":[-0.106236,0.026909,0.032522,-0.040702,0.06418,0.058867,-0.043462,0.029255,-0.069755,0.02625,-0.012982,0.011717,0.037221,0.004934,-0.148783,0.073643,0.065623,0.176729,0.000483,0.109778,-0.105427,0.02092,0.038877,-0.020431,0.06949,-0.027571,-0.023125,-0.036426,0.050819,-0.047508,-0.000816,-0.073777,0.060473,0.023495,-0.035937,-0.064082,-0.085108,-0.045895,0.029146,0.007555,0.003923,-0.081975,0.038573,0.072211,-0.063173,0.006889,-0.011189,-0.018597,-0.027534,-0.001236,0.056054,-0.055961,-0.016422,0.064765,-0.114341,-0.028348,-0.016785,-0.001291,0.00695,-0.073718,0.146327,-0.030716,-0.029492,0.066045,-0.049192,0.055391,0.105924,0.033385,0.049666,0.055198,-0.0327,-0.051139,-0.009429,0.113493,0.065922,-0.027484,-0.062549,-0.027209,0.006105,-0.016623,0.012859,-0.042844,0.046082,-0.029788,-0.132396,-0.005784,-0.020116,-0.042434,0.045863,-0.00114,0.022023,-0.117683,0.002397,0.187938,0.067782,-0.053011,0.018449,0.043227,-0.007033,0.034628,-0.036276,0.044694,-0.105351,-0.113562,0.051461,-0.040255,0.088879,0.101067,0.110482,-0.011473,0.030959,0.028627,0.042808,0.106958,-0.020069,-0.048371,0.002144,-0.063363,0.065032,-0.032523,0.025647,0.12314,0.026488,0.054077,0.026475,0.088169,-0.112363,-0.03501,-0.061824,-0.085917,0.027558,0.010396,-0.092906,-0.024477,0.027928,-0.012543,0.069895,0.01406,0.196947,0.006955,-0.012518,-0.048722,0.035787,-0.051044,0.007964,-0.009535,-0.045574,0.013616,-0.023598,-0.027236,0.023735,-0.023832,-0.053049,0.023769,0.080018,-0.039285,0.009429,-0.046602,0.023376,-0.011104,-0.118228,-0.089595,0.047087,-0.041417,-0.030876,-0.015441,-0.004501,-0.038595,0.070107,-0.078626,-0.128752,0.012564,-0.034697,-0.069278,0.058026,0.059071,0.082451,0.041674,0.031116,-0.007566,-0.007439,-0.079061,-0.059093,0.014351,-0.006355,-0.010857,-0.050468,0.059146,-0.047707,-0.08106,0.07763,0.065864,0.061567,-0.126551,-0.119901,-0.002173,0.076158,0.066434,0.053258,-0.065215,0.012435,0.015349,-0.00617,-0.033529,-0.024387,0.001493,0.113286,0.039088,0.051456,0.166047,-0.020073,-0.07846,-0.029372,0.022419,-0.06691,0.097686,-0.060915,-0.032454,-0.095959,-0.026277,0.026242,-0.123124,0.092298,-0.036004,0.031126,0.119571,0.09179,-0.060572,0.028013,0.016614,0.09969,-0.113456,-0.094285,-0.058706,-0.071713,0.131803,-0.066991,0.048721,-0.050628,-0.030655,0.010101,-0.046536,-0.024472,-2.4e-05,-0.04725,0.046721,0.036714,-0.082382,0.010479,0.053549,0.068698,0.006237,0.020804,-0.081368,-0.024734,0.051612]}]} \ No newline at end of file diff --git a/backend/go/vllm-cpp/Makefile b/backend/go/vllm-cpp/Makefile index 32733f11d..266d590a1 100644 --- a/backend/go/vllm-cpp/Makefile +++ b/backend/go/vllm-cpp/Makefile @@ -11,7 +11,7 @@ JOBS?=$(shell nproc --ignore=1 2>/dev/null || sysctl -n hw.ncpu 2>/dev/null || e # vllm.cpp version VLLM_CPP_REPO?=https://github.com/mudler/vllm.cpp -VLLM_CPP_VERSION?=96788348627b6a079fcc3ef7fc6676a970965d7b +VLLM_CPP_VERSION?=a19294a9ae2248e66d0206ecedee1b9041944e59 # MLX GEMM provider (darwin/metal only; see the metal branch below for why). # Consumed as the prebuilt pip wheel: building MLX from source needs `xcrun diff --git a/backend/go/voice-detect/Makefile b/backend/go/voice-detect/Makefile index 064c4b5a0..0e08f8cb7 100644 --- a/backend/go/voice-detect/Makefile +++ b/backend/go/voice-detect/Makefile @@ -1,6 +1,6 @@ # voice-detect backend Makefile. # -# Upstream pin lives below as VOICEDETECT_VERSION?=1db1759572c90faef6f3a78c36b5941a096a9f89 +# Upstream pin lives below as VOICEDETECT_VERSION?=b74a896f47c6d04fcca0a962ff317528fd0b0019 # can find and update it - matches the parakeet.cpp / whisper.cpp / ds4 convention). # # Local dev shortcut: if you already have an out-of-tree voice-detect.cpp build, @@ -13,7 +13,7 @@ # The default target below does the proper clone-at-pin + cmake build so CI does # not need a side-checkout. -VOICEDETECT_VERSION?=1db1759572c90faef6f3a78c36b5941a096a9f89 +VOICEDETECT_VERSION?=b74a896f47c6d04fcca0a962ff317528fd0b0019 VOICEDETECT_REPO?=https://github.com/localai-org/voice-detect.cpp GOCMD?=go diff --git a/backend/python/funasr/requirements-intel.txt b/backend/python/funasr/requirements-intel.txt index f00e3e925..a2cb6753a 100644 --- a/backend/python/funasr/requirements-intel.txt +++ b/backend/python/funasr/requirements-intel.txt @@ -2,3 +2,4 @@ torch torchaudio funasr +transformers>=4.32.0,<5 diff --git a/backend/python/funasr/requirements.txt b/backend/python/funasr/requirements.txt index 9ce0da738..e3d73cb2e 100644 --- a/backend/python/funasr/requirements.txt +++ b/backend/python/funasr/requirements.txt @@ -3,3 +3,4 @@ protobuf certifi packaging==24.1 setuptools +transformers>=4.32.0,<5 diff --git a/backend/python/whisperx/backend.py b/backend/python/whisperx/backend.py index dc6202286..c3f900e3e 100644 --- a/backend/python/whisperx/backend.py +++ b/backend/python/whisperx/backend.py @@ -16,7 +16,7 @@ import grpc sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'common')) sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common')) from grpc_auth import get_auth_interceptors -from transcript_utils import require_diarization_token, seconds_to_nanoseconds +from transcript_utils import diarize_or_keep, require_diarization_token, seconds_to_nanoseconds @@ -112,13 +112,15 @@ class BackendServicer(backend_pb2_grpc.BackendServicer): # Diarize if requested and HF token is available if request.diarize and self.hf_token: - if self.diarize_pipeline is None: - self.diarize_pipeline = DiarizationPipeline( - token=self.hf_token, - device=self.device, - ) - diarize_segments = self.diarize_pipeline(audio) - transcript = whisperx.assign_word_speakers(diarize_segments, transcript) + def _diarize(t): + if self.diarize_pipeline is None: + self.diarize_pipeline = DiarizationPipeline( + token=self.hf_token, + device=self.device, + ) + return whisperx.assign_word_speakers(self.diarize_pipeline(audio), t) + + transcript = diarize_or_keep(transcript, _diarize, lambda m: print(m, file=sys.stderr)) # Build result segments for idx, seg in enumerate(transcript["segments"]): @@ -137,8 +139,9 @@ class BackendServicer(backend_pb2_grpc.BackendServicer): text += seg_text except Exception as err: + # Report the failure instead of an empty, successful-looking result. print(f"Unexpected {err=}, {type(err)=}", file=sys.stderr) - return backend_pb2.TranscriptResult(segments=[], text="") + context.abort(grpc.StatusCode.INTERNAL, f"transcription failed: {err}") return backend_pb2.TranscriptResult(segments=resultSegments, text=text) diff --git a/backend/python/whisperx/test_transcript_utils.py b/backend/python/whisperx/test_transcript_utils.py index debe2ea6e..708d4bde2 100644 --- a/backend/python/whisperx/test_transcript_utils.py +++ b/backend/python/whisperx/test_transcript_utils.py @@ -21,5 +21,23 @@ class TestTranscriptUtils(unittest.TestCase): ) + def test_failed_diarization_keeps_the_transcript(self): + transcript = {"segments": [{"text": "Die Rechnung"}]} + logged = [] + + def refused(_): + raise RuntimeError("403 Client Error: gated repo") + + result = transcript_utils.diarize_or_keep(transcript, refused, logged.append) + self.assertIs(result, transcript) + self.assertIn("403", logged[0]) + + def test_successful_diarization_is_returned(self): + transcript = {"segments": [{"text": "Die Rechnung"}]} + with_speakers = {"segments": [{"text": "Die Rechnung", "speaker": "SPEAKER_00"}]} + result = transcript_utils.diarize_or_keep(transcript, lambda _: with_speakers, lambda _: None) + self.assertIs(result, with_speakers) + + if __name__ == "__main__": unittest.main() diff --git a/backend/python/whisperx/transcript_utils.py b/backend/python/whisperx/transcript_utils.py index a8ac57510..2211cf670 100644 --- a/backend/python/whisperx/transcript_utils.py +++ b/backend/python/whisperx/transcript_utils.py @@ -10,3 +10,17 @@ def require_diarization_token(diarize, token): def seconds_to_nanoseconds(seconds): """Convert WhisperX timestamps to the duration unit used by LocalAI.""" return int(seconds * 1_000_000_000) + + +def diarize_or_keep(transcript, diarize, log): + """Run diarization; if it fails, keep the transcript without speakers. + + Diarization is an add-on to a finished transcript. A refused download of + the gated pyannote pipeline (403) or any other diarization error must not + throw the transcript away. + """ + try: + return diarize(transcript) + except Exception as err: # noqa: BLE001 - any diarization failure degrades + log(f"Diarization failed, returning transcript without speakers: {err!r}") + return transcript diff --git a/core/backend/diarization.go b/core/backend/diarization.go index 241d1b20c..a87356cd4 100644 --- a/core/backend/diarization.go +++ b/core/backend/diarization.go @@ -2,11 +2,17 @@ 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" + "github.com/mudler/LocalAI/core/services/voicerecognition" grpcPkg "github.com/mudler/LocalAI/pkg/grpc" "github.com/mudler/LocalAI/pkg/grpc/proto" @@ -19,31 +25,41 @@ 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 } // modelIdentity: see the note on TranscriptionRequest.toProto. 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{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, + 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, } } @@ -76,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 { @@ -103,6 +132,7 @@ func diarizationResultFromProto(r *proto.DiarizeResponse) *schema.DiarizationRes idx int duration float64 segments int + name string } stats := map[string]*speakerStats{} order := []string{} @@ -126,14 +156,19 @@ func diarizationResultFromProto(r *proto.DiarizeResponse) *schema.DiarizationRes st.duration += dur } st.segments++ + if st.name == "" { + st.name = s.Name + } out.Segments = append(out.Segments, schema.DiarizationSegment{ - Id: i, - Speaker: fmt.Sprintf("SPEAKER_%02d", st.idx), - Label: raw, - Start: float64(s.Start), - End: float64(s.End), - Text: s.Text, + Id: i, + Speaker: fmt.Sprintf("SPEAKER_%02d", st.idx), + Label: raw, + Start: float64(s.Start), + End: float64(s.End), + Text: s.Text, + Name: s.Name, + NameScore: s.NameScore, }) } @@ -148,6 +183,7 @@ func diarizationResultFromProto(r *proto.DiarizeResponse) *schema.DiarizationRes out.Speakers = append(out.Speakers, schema.DiarizationSpeaker{ Id: fmt.Sprintf("SPEAKER_%02d", st.idx), Label: raw, + Name: st.name, TotalSpeechDuration: st.duration, SegmentCount: st.segments, }) @@ -158,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_names_test.go b/core/backend/diarization_names_test.go new file mode 100644 index 000000000..f01506256 --- /dev/null +++ b/core/backend/diarization_names_test.go @@ -0,0 +1,54 @@ +package backend + +import ( + "encoding/json" + + "github.com/mudler/LocalAI/core/services/voicerecognition" + "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("diarization names", func() { + It("adds name and name_score next to the normalized speaker, and keeps SPEAKER_NN", func() { + res := diarizationResultFromProto(&proto.DiarizeResponse{ + Duration: 10, + Segments: []*proto.DiarizeSegment{ + {Speaker: "0", Start: 0, End: 4, Name: "Ada", NameScore: 0.93}, + {Speaker: "1", Start: 4, End: 8}, + {Speaker: "0", Start: 8, End: 10, Name: "Ada", NameScore: 0.93}, + }, + }) + Expect(res.Segments[0].Speaker).To(Equal("SPEAKER_00")) + Expect(res.Segments[0].Name).To(Equal("Ada")) + Expect(res.Segments[0].NameScore).To(BeNumerically("~", 0.93, 1e-6)) + Expect(res.Segments[1].Speaker).To(Equal("SPEAKER_01")) + Expect(res.Segments[1].Name).To(BeEmpty()) + Expect(res.Speakers[0].Name).To(Equal("Ada")) + Expect(res.Speakers[1].Name).To(BeEmpty()) + }) + + It("marshals a result with no names byte for byte as before", func() { + res := diarizationResultFromProto(&proto.DiarizeResponse{ + Duration: 10, + Segments: []*proto.DiarizeSegment{{Speaker: "0", Start: 0, End: 1}}, + }) + b, err := json.Marshal(res) + Expect(err).ToNot(HaveOccurred()) + Expect(string(b)).To(Equal(`{"task":"diarize","duration":10,"num_speakers":1,` + + `"segments":[{"id":0,"speaker":"SPEAKER_00","label":"0","start":0,"end":1}],` + + `"speakers":[{"id":"SPEAKER_00","label":"0","total_speech_duration":1,"segment_count":1}]}`)) + }) + + It("sends the known voices to the backend", func() { + req := DiarizationRequest{Audio: "a.wav", KnownVoices: []voicerecognition.KnownVoice{ + {Name: "Ada", Embedding: []float32{1, 0}, Model: "m.gguf"}, + }} + p := req.toProto(0, "x") + Expect(p.KnownVoices).To(HaveLen(1)) + Expect(p.KnownVoices[0].Name).To(Equal("Ada")) + Expect(p.KnownVoices[0].Embedding).To(Equal([]float32{1, 0})) + Expect(p.KnownVoices[0].Model).To(Equal("m.gguf")) + Expect((&DiarizationRequest{Audio: "a.wav"}).toProto(0, "x").KnownVoices).To(BeEmpty()) + }) +}) 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 e64e54a57..7fcd27f82 100644 --- a/core/backend/transcript_live.go +++ b/core/backend/transcript_live.go @@ -11,6 +11,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/voicerecognition" "github.com/mudler/LocalAI/core/trace" grpcPkg "github.com/mudler/LocalAI/pkg/grpc" "github.com/mudler/LocalAI/pkg/grpc/proto" @@ -40,8 +41,11 @@ type LiveTranscriptionEvent struct { // are stream-relative seconds (mapped from the backend's nanoseconds). type LiveSpeakerSegment struct { Speaker string - Start float64 - End float64 + // Name is the registered speaker name the backend matched, empty when the + // speaker is unknown. + Name string + Start float64 + End float64 } // LiveSoundEvent is one closed sound event from a companion sound/scene @@ -222,6 +226,28 @@ func (ts *liveTraceState) record(closeErr error) { trace.RecordBackendTrace(bt) } +// LiveOption tunes a live transcription session. +type LiveOption func(*liveOptions) + +type liveOptions struct { + knownVoices []voicerecognition.KnownVoice +} + +// WithKnownVoices gives the backend the registered voices it may use to name +// the speakers it detects. Backends without speaker identification ignore them. +func WithKnownVoices(v []voicerecognition.KnownVoice) LiveOption { + return func(o *liveOptions) { o.knownVoices = v } +} + +// liveConfigProto builds the first message of a live session. +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{Id: v.ID, Name: v.Name, Embedding: v.Embedding, Model: v.Model}) + } + return cfg +} + // ModelTranscriptionLive loads the transcription backend, opens the // bidirectional AudioTranscriptionLive RPC, sends the session config, and // BLOCKS until the backend's ready ack. A grpcerrors. @@ -232,12 +258,18 @@ func (ts *liveTraceState) record(closeErr error) { // the backend streams, ending with the Final event triggered by Close. func ModelTranscriptionLive(ctx context.Context, language string, ml *model.ModelLoader, modelConfig config.ModelConfig, appConfig *config.ApplicationConfig, - onEvent func(LiveTranscriptionEvent)) (LiveTranscriptionSession, error) { + onEvent func(LiveTranscriptionEvent), opts ...LiveOption) (LiveTranscriptionSession, error) { + + lo := liveOptions{} + for _, f := range opts { + f(&lo) + } transcriptionModel, err := loadTranscriptionModel(ctx, ml, modelConfig, appConfig) if err != nil { return nil, err } + lo.knownVoices = compatiblePortableVoices(ctx, transcriptionModel, lo.knownVoices) release, err := AcquireGlobalBackendSlot() if err != nil { return nil, err @@ -262,10 +294,7 @@ func ModelTranscriptionLive(ctx context.Context, language string, } if err := stream.Send(&proto.TranscriptLiveRequest{ - Payload: &proto.TranscriptLiveRequest_Config{Config: &proto.TranscriptLiveConfig{ - Language: language, - SampleRate: liveSampleRate, - }}, + Payload: &proto.TranscriptLiveRequest_Config{Config: liveConfigProto(language, lo)}, }); err != nil { return fail(err) } @@ -329,6 +358,7 @@ func liveEventFromProto(r *proto.TranscriptLiveResponse) LiveTranscriptionEvent for _, s := range r.GetSpeakers() { ev.Speakers = append(ev.Speakers, LiveSpeakerSegment{ Speaker: s.GetSpeaker(), + Name: s.GetName(), Start: time.Duration(s.GetStart()).Seconds(), End: time.Duration(s.GetEnd()).Seconds(), }) diff --git a/core/backend/transcript_live_names_test.go b/core/backend/transcript_live_names_test.go new file mode 100644 index 000000000..34a76ff1c --- /dev/null +++ b/core/backend/transcript_live_names_test.go @@ -0,0 +1,36 @@ +package backend + +import ( + "github.com/mudler/LocalAI/core/services/voicerecognition" + "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("live speaker names", func() { + It("puts the known voices in the session config", func() { + o := liveOptions{} + WithKnownVoices([]voicerecognition.KnownVoice{{Name: "Ada", Embedding: []float32{1, 0}, Model: "m.gguf"}})(&o) + cfg := liveConfigProto("en", o) + Expect(cfg.Language).To(Equal("en")) + Expect(cfg.SampleRate).To(Equal(int32(liveSampleRate))) + Expect(cfg.KnownVoices).To(HaveLen(1)) + Expect(cfg.KnownVoices[0].Name).To(Equal("Ada")) + Expect(cfg.KnownVoices[0].Model).To(Equal("m.gguf")) + }) + It("sends no known voices by default", func() { + cfg := liveConfigProto("en", liveOptions{}) + Expect(cfg.KnownVoices).To(BeEmpty()) + Expect(cfg.Language).To(Equal("en")) + Expect(cfg.SampleRate).To(Equal(int32(liveSampleRate))) + }) + It("carries the speaker name from the backend event", func() { + ev := liveEventFromProto(&proto.TranscriptLiveResponse{ + Speakers: []*proto.LiveSpeakerSegment{{Speaker: "0", Name: "Ada", Start: 1e9, End: 2e9}, {Speaker: "1", Start: 2e9, End: 3e9}}, + }) + Expect(ev.Speakers).To(HaveLen(2)) + Expect(ev.Speakers[0].Name).To(Equal("Ada")) + Expect(ev.Speakers[0].Speaker).To(Equal("0")) + Expect(ev.Speakers[1].Name).To(BeEmpty()) + }) +}) 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 d8d97d619..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,22 +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(), voicerecognition.Metadata{ - Name: input.Name, - Labels: input.Labels, - }) + stored, err := registry.Register(c.Request().Context(), embedding, voiceMetadata(input.Name, input.Labels, encoder)) if err != nil { return err } @@ -59,3 +74,10 @@ func VoiceRegisterEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, }) } } + +// voiceMetadata is what a registration stores next to the embedding. Model is +// the speaker encoder that produced it, so a consumer with a different encoder +// can tell the vectors are not comparable. +func voiceMetadata(name string, labels map[string]string, embedderModel string) voicerecognition.Metadata { + return voicerecognition.Metadata{Name: name, Labels: labels, Model: embedderModel} +} diff --git a/core/http/endpoints/localai/voice_register_test.go b/core/http/endpoints/localai/voice_register_test.go new file mode 100644 index 000000000..b66d0e48d --- /dev/null +++ b/core/http/endpoints/localai/voice_register_test.go @@ -0,0 +1,18 @@ +package localai + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("voiceMetadata", func() { + It("carries the name, the labels and the encoder that embedded the voice", func() { + m := voiceMetadata("ada", map[string]string{"team": "a"}, "voice-detect-wespeaker-resnet34.gguf") + Expect(m.Name).To(Equal("ada")) + Expect(m.Labels).To(Equal(map[string]string{"team": "a"})) + Expect(m.Model).To(Equal("voice-detect-wespeaker-resnet34.gguf")) + }) + It("leaves the tag empty when the backend did not say", func() { + Expect(voiceMetadata("ada", nil, "").Model).To(BeEmpty()) + }) +}) diff --git a/core/http/endpoints/openai/audio_upload_test.go b/core/http/endpoints/openai/audio_upload_test.go index ee788a4fd..028bb39b1 100644 --- a/core/http/endpoints/openai/audio_upload_test.go +++ b/core/http/endpoints/openai/audio_upload_test.go @@ -27,7 +27,7 @@ var _ = Describe("audio upload endpoints reject bad uploads as client errors", f return TranscriptEndpoint(nil, nil, config.NewApplicationConfig()) }}, "diarization": {"/v1/audio/diarization", func() echo.HandlerFunc { - return DiarizationEndpoint(nil, nil, config.NewApplicationConfig()) + return DiarizationEndpoint(nil, nil, config.NewApplicationConfig(), nil) }}, "sound classification": {"/v1/audio/classifications", func() echo.HandlerFunc { return SoundClassificationEndpoint(nil, nil, config.NewApplicationConfig()) diff --git a/core/http/endpoints/openai/diarization.go b/core/http/endpoints/openai/diarization.go index c3128c126..1ab745e07 100644 --- a/core/http/endpoints/openai/diarization.go +++ b/core/http/endpoints/openai/diarization.go @@ -1,6 +1,9 @@ package openai import ( + "bytes" + "context" + "encoding/base64" "fmt" "io" "net/http" @@ -9,13 +12,19 @@ import ( "path/filepath" "strconv" "strings" + "sync" "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" ) @@ -32,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)" @@ -43,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) 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 == "" { @@ -60,8 +71,23 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap } req := backend.DiarizationRequest{ - Language: c.FormValue("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 language := c.FormValue("language"); language != "" { + req.Language = language + } + 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)) @@ -69,8 +95,18 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap req.ClusteringThreshold = float32(parseFormFloat(c, "clustering_threshold", 0)) req.MinDurationOn = float32(parseFormFloat(c, "min_duration_on", 0)) req.MinDurationOff = float32(parseFormFloat(c, "min_duration_off", 0)) + 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 } @@ -82,15 +118,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 { @@ -98,13 +149,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 } @@ -113,20 +164,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: @@ -137,6 +196,48 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap } } +// attachKnownVoices names the speakers from the voice registry when the model +// has a speaker model. Only voices made by that model's encoder are sent. A +// missing registry or speaker model, or a registry read error, leaves the +// request unnamed. +func attachKnownVoices(ctx context.Context, req *backend.DiarizationRequest, options []string, registry voicerecognition.Registry) { + req.KnownVoices = selectKnownVoices(ctx, "diarization", options, registry) +} + +// warned remembers the (feature, speaker model) pairs already warned about. +var warned sync.Map + +// warnOnce reports true the first time it sees key, false afterwards. +func warnOnce(key string) bool { + _, loaded := warned.LoadOrStore(key, struct{}{}) + return !loaded +} + +// selectKnownVoices returns the registered voices a backend may use to name +// speakers, or nil when the model has no speaker_model, there is no registry, +// or the registry cannot be read. It never fails the caller: unnamed speakers +// are the fallback. feature only prefixes the log messages. +func selectKnownVoices(ctx context.Context, feature string, options []string, registry voicerecognition.Registry) []voicerecognition.KnownVoice { + sm := voicerecognition.SpeakerModelFromOptions(options) + if sm == "" || registry == nil { + return nil + } + sel, err := voicerecognition.KnownVoicesFor(ctx, registry, sm) + if err != nil { + xlog.Warn(feature+": could not read the voice registry; speakers stay unnamed", "error", err) + return nil + } + if len(sel.Voices) == 0 && sel.OtherEncoder > 0 { + msg := feature + ": registered voices were made with a different encoder than this model's speaker_model; speakers stay unnamed" + if warnOnce(feature + "|" + sm) { + xlog.Warn(msg, "speaker_model", sm, "voices_from_other_encoder", sel.OtherEncoder) + } else { + xlog.Debug(msg, "speaker_model", sm, "voices_from_other_encoder", sel.OtherEncoder) + } + } + return sel.Voices +} + // renderRTTM emits NIST RTTM rows. Each row: // SPEAKER 1 // Field separators are spaces; one row per segment. diff --git a/core/http/endpoints/openai/diarization_test.go b/core/http/endpoints/openai/diarization_test.go index 9cba206a3..6a27a04d5 100644 --- a/core/http/endpoints/openai/diarization_test.go +++ b/core/http/endpoints/openai/diarization_test.go @@ -1,9 +1,13 @@ package openai import ( + "context" + "errors" "strings" + "github.com/mudler/LocalAI/core/backend" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/voicerecognition" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -49,3 +53,60 @@ var _ = Describe("renderRTTM", func() { Expect(out).To(HavePrefix("SPEAKER audio 1 ")) }) }) + +type fakeVoiceRegistry struct { + voicerecognition.Registry + entries []voicerecognition.Entry + err error +} + +func (f fakeVoiceRegistry) List(context.Context) ([]voicerecognition.Entry, error) { + return f.entries, f.err +} + +var _ = Describe("attachKnownVoices", func() { + ada := voicerecognition.Entry{ + Metadata: voicerecognition.Metadata{Name: "Ada", Model: "enc.gguf"}, + Embedding: []float32{1, 0}, + } + + It("sends the voices made by the speaker model's encoder", func() { + req := backend.DiarizationRequest{} + attachKnownVoices(context.Background(), &req, []string{"speaker_model:/models/enc.gguf"}, + fakeVoiceRegistry{entries: []voicerecognition.Entry{ada}}) + Expect(req.KnownVoices).To(HaveLen(1)) + Expect(req.KnownVoices[0].Name).To(Equal("Ada")) + }) + + It("warns once per key", func() { + Expect(warnOnce("diarization|warn-once-test.gguf")).To(BeTrue()) + Expect(warnOnce("diarization|warn-once-test.gguf")).To(BeFalse()) + Expect(warnOnce("live|warn-once-test.gguf")).To(BeTrue()) + }) + It("leaves the request alone without a speaker_model option", func() { + req := backend.DiarizationRequest{} + attachKnownVoices(context.Background(), &req, []string{"other:x"}, + fakeVoiceRegistry{entries: []voicerecognition.Entry{ada}}) + Expect(req.KnownVoices).To(BeEmpty()) + }) + + It("leaves the request alone without a registry", func() { + req := backend.DiarizationRequest{} + attachKnownVoices(context.Background(), &req, []string{"speaker_model:enc.gguf"}, nil) + Expect(req.KnownVoices).To(BeEmpty()) + }) + + It("leaves the request unnamed when the registry cannot be read", func() { + req := backend.DiarizationRequest{} + attachKnownVoices(context.Background(), &req, []string{"speaker_model:enc.gguf"}, + fakeVoiceRegistry{err: errors.New("boom")}) + Expect(req.KnownVoices).To(BeEmpty()) + }) + + It("skips voices made by another encoder", func() { + req := backend.DiarizationRequest{} + attachKnownVoices(context.Background(), &req, []string{"speaker_model:other.gguf"}, + fakeVoiceRegistry{entries: []voicerecognition.Entry{ada}}) + Expect(req.KnownVoices).To(BeEmpty()) + }) +}) diff --git a/core/http/endpoints/openai/realtime.go b/core/http/endpoints/openai/realtime.go index 11d1653cb..c152d2c22 100644 --- a/core/http/endpoints/openai/realtime.go +++ b/core/http/endpoints/openai/realtime.go @@ -35,6 +35,7 @@ import ( "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/routing/router" "github.com/mudler/LocalAI/core/services/voiceprofile" + "github.com/mudler/LocalAI/core/services/voicerecognition" "github.com/mudler/LocalAI/core/templates" laudio "github.com/mudler/LocalAI/pkg/audio" "github.com/mudler/LocalAI/pkg/functions" @@ -841,6 +842,7 @@ func runRealtimeSession(application *application.Application, t Transport, model application.ModelLoader(), application.ApplicationConfig(), application.FailoverManager(), + application.VoiceRegistry(), ); err != nil { xlog.Error("failed to update session", "error", err) // The cause is validation feedback on the client's own @@ -1158,7 +1160,7 @@ func sendTestTone(t Transport) { } } -func updateTransSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) error { +func updateTransSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager, voices voicerecognition.Registry) error { sessionLock.Lock() defer sessionLock.Unlock() @@ -1186,6 +1188,9 @@ func updateTransSession(session *Session, update *types.SessionUnion, cl *config return err } + if tm, ok := m.(*transcriptOnlyModel); ok { + tm.voiceRegistry = voices + } session.ModelInterface = m session.ModelConfig = cfg session.SoundDetectionEnabled = cfg.Pipeline.SoundDetection != "" diff --git a/core/http/endpoints/openai/realtime_live_voices_test.go b/core/http/endpoints/openai/realtime_live_voices_test.go new file mode 100644 index 000000000..c30679a69 --- /dev/null +++ b/core/http/endpoints/openai/realtime_live_voices_test.go @@ -0,0 +1,41 @@ +package openai + +import ( + "context" + "encoding/json" + "errors" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/voicerecognition" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("liveVoiceOptions", func() { + ada := voicerecognition.Entry{ + Metadata: voicerecognition.Metadata{Name: "Ada", Model: "enc.gguf"}, + Embedding: []float32{1, 0}, + } + cfgWith := func(opts ...string) *config.ModelConfig { return &config.ModelConfig{Options: opts} } + + It("opens the session with the registered voices", func() { + opts := liveVoiceOptions(context.Background(), fakeVoiceRegistry{entries: []voicerecognition.Entry{ada}}, cfgWith("speaker_model:enc.gguf")) + Expect(opts).To(HaveLen(1)) + }) + It("adds nothing without a speaker_model, a registry, or when the registry fails", func() { + reg := fakeVoiceRegistry{entries: []voicerecognition.Entry{ada}} + Expect(liveVoiceOptions(context.Background(), reg, cfgWith("x:y"))).To(BeEmpty()) + Expect(liveVoiceOptions(context.Background(), nil, cfgWith("speaker_model:enc.gguf"))).To(BeEmpty()) + Expect(liveVoiceOptions(context.Background(), fakeVoiceRegistry{err: errors.New("boom")}, cfgWith("speaker_model:enc.gguf"))).To(BeEmpty()) + }) +}) + +var _ = Describe("transcription segment event", func() { + It("carries speaker_name only when set", func() { + named, _ := json.Marshal(types.ConversationItemInputAudioTranscriptionSegmentEvent{Speaker: "0", SpeakerName: "Ada"}) + Expect(string(named)).To(ContainSubstring(`"speaker_name":"Ada"`)) + plain, _ := json.Marshal(types.ConversationItemInputAudioTranscriptionSegmentEvent{Speaker: "0"}) + Expect(string(plain)).ToNot(ContainSubstring("speaker_name")) + }) +}) diff --git a/core/http/endpoints/openai/realtime_model.go b/core/http/endpoints/openai/realtime_model.go index 678a57a0d..dfc49112a 100644 --- a/core/http/endpoints/openai/realtime_model.go +++ b/core/http/endpoints/openai/realtime_model.go @@ -22,6 +22,7 @@ import ( "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/routing/router" "github.com/mudler/LocalAI/core/services/voiceprofile" + "github.com/mudler/LocalAI/core/services/voicerecognition" "github.com/mudler/LocalAI/core/templates" "github.com/mudler/LocalAI/pkg/functions" "github.com/mudler/LocalAI/pkg/grpc/proto" @@ -86,6 +87,9 @@ type wrappedModel struct { routerSessionID string routerUserID string + // voiceRegistry names live speakers from registered voices; nil disables it. + voiceRegistry voicerecognition.Registry + stageRouter // tuneLLM applies the pipeline's LLM overrides (reasoning effort, // disable_thinking) to a chain target loaded per call. @@ -113,6 +117,9 @@ type transcriptOnlyModel struct { modelLoader *model.ModelLoader confLoader *config.ModelConfigLoader + // voiceRegistry names live speakers from registered voices; nil disables it. + voiceRegistry voicerecognition.Registry + stageRouter } @@ -187,7 +194,7 @@ func (m *transcriptOnlyModel) TranscribeLive(ctx context.Context, language strin // Only opening the live session can move to the next target. err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { var err error - live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent) + live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent, liveVoiceOptions(ctx, m.voiceRegistry, cfg)...) return err }) return live, err @@ -579,7 +586,7 @@ func (m *wrappedModel) TranscribeLive(ctx context.Context, language string, onEv // open, events flow to the client for the rest of the utterance. err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { var err error - live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent) + live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent, liveVoiceOptions(ctx, m.voiceRegistry, cfg)...) return err }) return live, err @@ -1018,6 +1025,8 @@ type RealtimeRoutingContext struct { UserID string // Failover resolves pipeline stages that name a failover chain. Failover *failover.Manager + // VoiceRegistry holds the voices registered through /v1/voice/register. + VoiceRegistry voicerecognition.Registry } // buildRealtimeRoutingContext assembles the routing dependencies the @@ -1040,6 +1049,8 @@ func buildRealtimeRoutingContext(a *application.Application, sessionID string) * SessionID: sessionID, UserID: userID, Failover: a.FailoverManager(), + + VoiceRegistry: a.VoiceRegistry(), } } @@ -1203,6 +1214,17 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model wm.routerStore = routing.Store wm.routerSessionID = routing.SessionID wm.routerUserID = routing.UserID + wm.voiceRegistry = routing.VoiceRegistry } return wm, nil } + +// liveVoiceOptions selects the registered voices a live session may name speakers +// with. It stays empty without a speaker_model or a voice registry. +func liveVoiceOptions(ctx context.Context, registry voicerecognition.Registry, cfg *config.ModelConfig) []backend.LiveOption { + voices := selectKnownVoices(ctx, "live transcription", cfg.Options, registry) + if len(voices) == 0 { + return nil + } + return []backend.LiveOption{backend.WithKnownVoices(voices)} +} diff --git a/core/http/endpoints/openai/realtime_semantic_vad.go b/core/http/endpoints/openai/realtime_semantic_vad.go index df56abf19..fa1ed0713 100644 --- a/core/http/endpoints/openai/realtime_semantic_vad.go +++ b/core/http/endpoints/openai/realtime_semantic_vad.go @@ -220,6 +220,7 @@ func (l *liveTurnState) drainEvents(audioSec float64) { ItemID: l.itemID, ContentIndex: 0, Speaker: seg.Speaker, + SpeakerName: seg.Name, Start: seg.Start, End: seg.End, }) diff --git a/core/http/endpoints/openai/realtime_voicegate_test.go b/core/http/endpoints/openai/realtime_voicegate_test.go index 3d9b458e1..a4c900d46 100644 --- a/core/http/endpoints/openai/realtime_voicegate_test.go +++ b/core/http/endpoints/openai/realtime_voicegate_test.go @@ -116,6 +116,9 @@ func (f *fakeRegistry) Identify(ctx context.Context, probe []float32, topK int) return f.matches, f.err } func (f *fakeRegistry) Forget(ctx context.Context, id string) error { return nil } +func (f *fakeRegistry) List(ctx context.Context) ([]voicerecognition.Entry, error) { + return nil, nil +} var _ = Describe("voiceGate identify mode", func() { stubEmbed := func(emb []float32, err error) func(context.Context, string) ([]float32, error) { diff --git a/core/http/endpoints/openai/types/server_events.go b/core/http/endpoints/openai/types/server_events.go index 4cab30a61..cdf88bf84 100644 --- a/core/http/endpoints/openai/types/server_events.go +++ b/core/http/endpoints/openai/types/server_events.go @@ -595,6 +595,10 @@ type ConversationItemInputAudioTranscriptionSegmentEvent struct { // The speaker label for the segment, if available. Speaker string `json:"speaker,omitempty"` + // The registered name of the speaker, when the backend recognised a + // voice registered through /v1/voice/register. + SpeakerName string `json:"speaker_name,omitempty"` + // The start time of the segment in seconds. Always present (not // omitempty: a segment starting at 0.0s must still carry "start"). Start float64 `json:"start"` 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 a2214f12e..8c8530c1d 100644 --- a/core/http/routes/openai.go +++ b/core/http/routes/openai.go @@ -190,7 +190,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()) + 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 cbac34f5d..5ad88c20e 100644 --- a/core/schema/diarization.go +++ b/core/schema/diarization.go @@ -12,6 +12,11 @@ type DiarizationSegment struct { Start float64 `json:"start"` End float64 `json:"end"` Text string `json:"text,omitempty"` + // Name is the registered speaker this segment was matched to, and NameScore + // the cosine similarity of the match. Both are omitted when the backend did + // not identify the speaker. Speaker stays the normalized SPEAKER_NN label. + Name string `json:"name,omitempty"` + NameScore float32 `json:"name_score,omitempty"` } // DiarizationSpeaker summarizes one speaker across the whole audio so @@ -20,6 +25,7 @@ type DiarizationSegment struct { type DiarizationSpeaker struct { Id string `json:"id"` Label string `json:"label,omitempty"` + Name string `json:"name,omitempty"` TotalSpeechDuration float64 `json:"total_speech_duration"` SegmentCount int `json:"segment_count"` } @@ -28,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 new file mode 100644 index 000000000..0f4cf8fba --- /dev/null +++ b/core/services/voicerecognition/known_voices.go @@ -0,0 +1,86 @@ +package voicerecognition + +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 +} + +// SpeakerModelFromOptions returns the value of a speaker_model: entry in +// a model config's options, or "" when there is none. +func SpeakerModelFromOptions(options []string) string { + for _, o := range options { + k, v, ok := strings.Cut(o, ":") + if ok && strings.TrimSpace(k) == "speaker_model" { + return strings.TrimSpace(v) + } + } + return "" +} + +// EncoderTag is how an encoder is identified: the lowercased base name of its +// model file. A voice registered through the voice-detect backend carries the +// backend's model name, which defaults to that base name. +func EncoderTag(modelPath string) string { + return strings.ToLower(path.Base(filepath.ToSlash(modelPath))) +} + +// KnownVoiceSelection is the result of SelectKnownVoices. +type KnownVoiceSelection struct { + Voices []KnownVoice + OtherEncoder int // voices skipped because another encoder made them + Untagged int // voices with no encoder tag that were included +} + +// 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 + var untagged []Entry + for _, e := range entries { + if e.Metadata.Name == "" || len(e.Embedding) == 0 { + continue + } + switch { + case e.Metadata.Model == "": + untagged = append(untagged, e) + // 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++ + } + } + 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 { + sel.Untagged++ + sel.Voices = append(sel.Voices, KnownVoice{ID: e.Metadata.ID, Name: e.Metadata.Name, Embedding: e.Embedding}) + } + return sel +} + +// KnownVoicesFor lists the registry and selects the voices for a speaker model. +func KnownVoicesFor(ctx context.Context, reg Registry, speakerModelPath string) (KnownVoiceSelection, error) { + entries, err := reg.List(ctx) + if err != nil { + return KnownVoiceSelection{}, err + } + return SelectKnownVoices(entries, speakerModelPath), nil +} diff --git a/core/services/voicerecognition/known_voices_test.go b/core/services/voicerecognition/known_voices_test.go new file mode 100644 index 000000000..fbe5510cf --- /dev/null +++ b/core/services/voicerecognition/known_voices_test.go @@ -0,0 +1,131 @@ +package voicerecognition_test + +import ( + "context" + "errors" + + "github.com/mudler/LocalAI/core/services/voicerecognition" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func entry(name, model string, emb ...float32) voicerecognition.Entry { + return voicerecognition.Entry{ + Metadata: voicerecognition.Metadata{ID: name, Name: name, Model: model}, + Embedding: emb, + } +} + +var _ = Describe("SpeakerModelFromOptions", func() { + It("reads the speaker_model option", func() { + Expect(voicerecognition.SpeakerModelFromOptions([]string{"diarization_model:d.gguf", "speaker_model: voice-detect-wespeaker-resnet34.gguf "})). + To(Equal("voice-detect-wespeaker-resnet34.gguf")) + }) + It("is empty without the option", func() { + Expect(voicerecognition.SpeakerModelFromOptions([]string{"diarization_model:d.gguf"})).To(BeEmpty()) + Expect(voicerecognition.SpeakerModelFromOptions(nil)).To(BeEmpty()) + }) +}) + +var _ = Describe("EncoderTag", func() { + It("is the lowercased basename", func() { + Expect(voicerecognition.EncoderTag("models/Voice-Detect-WeSpeaker.GGUF")).To(Equal("voice-detect-wespeaker.gguf")) + Expect(voicerecognition.EncoderTag("voice-detect-ecapa-tdnn-voxceleb.gguf")).To(Equal("voice-detect-ecapa-tdnn-voxceleb.gguf")) + }) +}) + +var _ = Describe("SelectKnownVoices", func() { + const wespeaker = "voice-detect-wespeaker-resnet34.gguf" + It("keeps the voices made by the speaker model's encoder", func() { + sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{ + entry("ada", wespeaker, 1, 0), + entry("ben", "voice-detect-ecapa-tdnn-voxceleb.gguf", 0, 1, 0), + entry("cy", "Voice-Detect-WeSpeaker-ResNet34.gguf", 0, 1), + }, "models/"+wespeaker) + Expect(sel.Voices).To(HaveLen(2)) + Expect(sel.Voices[0].Name).To(Equal("ada")) + Expect(sel.Voices[1].Name).To(Equal("cy")) + Expect(sel.OtherEncoder).To(Equal(1)) + }) + It("matches a speaker_model value that has a directory and upper-case letters", func() { + sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{entry("ada", wespeaker, 1, 0)}, "Some/Dir/Voice-Detect-WeSpeaker-ResNet34.GGUF") + Expect(sel.Voices).To(HaveLen(1)) + }) + It("sends nothing, and says why, when every voice is from another encoder", func() { + sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{entry("ben", "voice-detect-ecapa-tdnn-voxceleb.gguf", 0, 1, 0)}, wespeaker) + Expect(sel.Voices).To(BeEmpty()) + Expect(sel.OtherEncoder).To(Equal(1)) + }) + 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(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{ + entry("old", "", 0, 1), + entry("ada", wespeaker, 1, 0), + }, wespeaker) + Expect(sel.Voices).To(HaveLen(2)) + Expect(sel.Voices[0].Name).To(Equal("ada")) + Expect(sel.Voices[1].Name).To(Equal("old")) + }) + It("includes every untagged voice when no voice is tagged for this encoder", func() { + sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{entry("old", "", 1, 0)}, wespeaker) + Expect(sel.Voices).To(HaveLen(1)) + Expect(sel.Untagged).To(Equal(1)) + }) + It("includes untagged voices of different sizes when no voice is tagged for this encoder", func() { + sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{entry("old", "", 1, 0), entry("older", "", 1, 0, 0)}, wespeaker) + Expect(sel.Voices).To(HaveLen(2)) // the backend skips the ones whose size does not match + Expect(sel.Untagged).To(Equal(2)) + }) + It("skips voices without a name or an embedding", func() { + sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{entry("", wespeaker, 1, 0), entry("x", wespeaker)}, wespeaker) + Expect(sel.Voices).To(BeEmpty()) + }) + It("is empty for an empty registry", func() { + Expect(voicerecognition.SelectKnownVoices(nil, wespeaker).Voices).To(BeEmpty()) + }) + It("handles an empty speaker model path without matching tagged voices", func() { + sel := voicerecognition.SelectKnownVoices([]voicerecognition.Entry{entry("ada", wespeaker, 1, 0), entry("old", "", 0, 1)}, "") + Expect(sel.OtherEncoder).To(Equal(1)) + Expect(sel.Voices).To(HaveLen(1)) + Expect(sel.Voices[0].Name).To(Equal("old")) + }) + It("does not modify the input entries", func() { + in := []voicerecognition.Entry{entry("old", "", 0, 1), entry("ada", wespeaker, 1, 0)} + voicerecognition.SelectKnownVoices(in, wespeaker) + Expect(in[0].Metadata.Name).To(Equal("old")) + Expect(in[1].Metadata.Name).To(Equal("ada")) + }) +}) + +type listRegistry struct { + voicerecognition.Registry + entries []voicerecognition.Entry + err error +} + +func (r listRegistry) List(context.Context) ([]voicerecognition.Entry, error) { + return r.entries, r.err +} + +var _ = Describe("KnownVoicesFor", func() { + It("selects from the registry listing", func() { + sel, err := voicerecognition.KnownVoicesFor(context.Background(), listRegistry{entries: []voicerecognition.Entry{entry("ada", "m.gguf", 1)}}, "m.gguf") + Expect(err).ToNot(HaveOccurred()) + Expect(sel.Voices).To(HaveLen(1)) + }) + It("propagates a List error", func() { + boom := errors.New("boom") + _, err := voicerecognition.KnownVoicesFor(context.Background(), listRegistry{err: boom}, "m.gguf") + Expect(err).To(MatchError(boom)) + }) +}) diff --git a/core/services/voicerecognition/registry.go b/core/services/voicerecognition/registry.go index 85ed9e3b7..b76c0a16c 100644 --- a/core/services/voicerecognition/registry.go +++ b/core/services/voicerecognition/registry.go @@ -32,6 +32,19 @@ type Registry interface { // Forget removes a previously-registered embedding by ID. // Returns ErrNotFound if the ID is unknown. Forget(ctx context.Context, id string) error + + // List returns every registered voice with its embedding, oldest first + // (ties broken by ID). The embeddings are copies. The store registry + // answers from its in-process index, so it lists what this process + // registered since it started, which is also everything the in-memory + // store holds. + List(ctx context.Context) ([]Entry, error) +} + +// Entry is a registered voice together with its embedding, as returned by List. +type Entry struct { + Metadata Metadata + Embedding []float32 } // Metadata is the user-supplied payload stored alongside a speaker embedding. @@ -41,6 +54,10 @@ type Metadata struct { Name string `json:"name"` Labels map[string]string `json:"labels,omitempty"` RegisteredAt time.Time `json:"registered_at"` + // Model names the speaker encoder that produced the embedding (the voice + // backend's model name, by default the GGUF file name). Empty for voices + // registered before this field existed. + Model string `json:"model,omitempty"` } // Match is a single result from Identify, ranked by similarity. diff --git a/core/services/voicerecognition/store_registry.go b/core/services/voicerecognition/store_registry.go index 39df94619..94e2897f6 100644 --- a/core/services/voicerecognition/store_registry.go +++ b/core/services/voicerecognition/store_registry.go @@ -41,11 +41,11 @@ type storeRegistry struct { dim int // TODO(postgres): the local-store gRPC surface keys by embedding - // vector and exposes no "list all" method, so we cannot delete by - // ID without remembering the embedding. This in-memory index is - // rebuilt on every Register and lost on restart — acceptable while - // the only implementation is itself in-memory. - idIndex sync.Map // map[string][]float32 + // vector and exposes no "list all" method, so we cannot delete by ID + // or list voices without remembering them. This in-memory index holds + // every registration with its metadata. It is rebuilt on every Register + // and lost on restart, which matches the lifetime of the in-memory store. + idIndex sync.Map // map[string]Entry } func (r *storeRegistry) Register(ctx context.Context, embedding []float32, meta Metadata) (Metadata, error) { @@ -76,7 +76,7 @@ func (r *storeRegistry) Register(ctx context.Context, embedding []float32, meta } embCopy := append([]float32(nil), embedding...) - r.idIndex.Store(meta.ID, embCopy) + r.idIndex.Store(meta.ID, Entry{Metadata: meta, Embedding: embCopy}) return meta, nil } @@ -124,7 +124,7 @@ func (r *storeRegistry) Forget(ctx context.Context, id string) error { if !ok { return ErrNotFound } - embedding := raw.([]float32) + embedding := raw.(Entry).Embedding backend, err := r.resolve(ctx, r.storeName) if err != nil { @@ -136,3 +136,20 @@ func (r *storeRegistry) Forget(ctx context.Context, id string) error { r.idIndex.Delete(id) return nil } + +func (r *storeRegistry) List(ctx context.Context) ([]Entry, error) { + var out []Entry + r.idIndex.Range(func(_, v any) bool { + e := v.(Entry) + e.Embedding = append([]float32(nil), e.Embedding...) + out = append(out, e) + return true + }) + sort.SliceStable(out, func(i, j int) bool { + if !out[i].Metadata.RegisteredAt.Equal(out[j].Metadata.RegisteredAt) { + return out[i].Metadata.RegisteredAt.Before(out[j].Metadata.RegisteredAt) + } + return out[i].Metadata.ID < out[j].Metadata.ID + }) + return out, nil +} diff --git a/core/services/voicerecognition/store_registry_test.go b/core/services/voicerecognition/store_registry_test.go new file mode 100644 index 000000000..bd3559888 --- /dev/null +++ b/core/services/voicerecognition/store_registry_test.go @@ -0,0 +1,110 @@ +package voicerecognition_test + +import ( + "context" + "sync" + + "github.com/mudler/LocalAI/core/services/voicerecognition" + "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + ggrpc "google.golang.org/grpc" +) + +// fakeStore records Set and Delete calls; everything else panics (nil embedded interface). +type fakeStore struct { + grpc.Backend + mu sync.Mutex + sets int + deletes int + values [][]byte +} + +func (f *fakeStore) StoresSet(ctx context.Context, in *pb.StoresSetOptions, opts ...ggrpc.CallOption) (*pb.Result, error) { + f.mu.Lock() + defer f.mu.Unlock() + f.sets++ + for _, v := range in.Values { + f.values = append(f.values, v.Bytes) + } + return &pb.Result{Success: true}, nil +} + +func (f *fakeStore) StoresDelete(ctx context.Context, in *pb.StoresDeleteOptions, opts ...ggrpc.CallOption) (*pb.Result, error) { + f.mu.Lock() + defer f.mu.Unlock() + f.deletes++ + return &pb.Result{Success: true}, nil +} + +// StoresFind returns every stored value as a perfect match. +func (f *fakeStore) StoresFind(ctx context.Context, in *pb.StoresFindOptions, opts ...ggrpc.CallOption) (*pb.StoresFindResult, error) { + f.mu.Lock() + defer f.mu.Unlock() + res := &pb.StoresFindResult{} + for _, v := range f.values { + res.Keys = append(res.Keys, &pb.StoresKey{Floats: in.Key.Floats}) + res.Values = append(res.Values, &pb.StoresValue{Bytes: v}) + res.Similarities = append(res.Similarities, 1) + } + return res, nil +} + +var _ = Describe("storeRegistry List", func() { + var ( + fs *fakeStore + reg voicerecognition.Registry + ctx = context.Background() + ) + BeforeEach(func() { + fs = &fakeStore{} + reg = voicerecognition.NewStoreRegistry(func(context.Context, string) (grpc.Backend, error) { return fs, nil }, "t", 0) + }) + + It("lists what was registered, with the encoder tag and a copy of the embedding", func() { + a, err := reg.Register(ctx, []float32{1, 0, 0}, voicerecognition.Metadata{Name: "ada", Model: "voice-detect-wespeaker-resnet34.gguf"}) + Expect(err).ToNot(HaveOccurred()) + _, err = reg.Register(ctx, []float32{0, 1, 0}, voicerecognition.Metadata{Name: "ben"}) + Expect(err).ToNot(HaveOccurred()) + + got, err := reg.List(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(HaveLen(2)) + Expect(got[0].Metadata.ID).To(Equal(a.ID)) // oldest first + Expect(got[0].Metadata.Name).To(Equal("ada")) + Expect(got[0].Metadata.Model).To(Equal("voice-detect-wespeaker-resnet34.gguf")) + Expect(got[0].Embedding).To(Equal([]float32{1, 0, 0})) + Expect(got[1].Metadata.Model).To(BeEmpty()) // an untagged (legacy style) registration + + got[0].Embedding[0] = 99 // mutating the result must not touch the registry + again, _ := reg.List(ctx) + Expect(again[0].Embedding[0]).To(Equal(float32(1))) + }) + + It("forgets a voice from the list", func() { + a, _ := reg.Register(ctx, []float32{1, 0}, voicerecognition.Metadata{Name: "ada"}) + _, _ = reg.Register(ctx, []float32{0, 1}, voicerecognition.Metadata{Name: "ben"}) + Expect(reg.Forget(ctx, a.ID)).To(Succeed()) + got, err := reg.List(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(HaveLen(1)) + Expect(got[0].Metadata.Name).To(Equal("ben")) + Expect(fs.deletes).To(Equal(1)) + }) + + It("is empty for a fresh registry", func() { + got, err := reg.List(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(BeEmpty()) + }) + + It("returns the encoder tag from Identify", func() { + _, err := reg.Register(ctx, []float32{1, 0}, voicerecognition.Metadata{Name: "ada", Model: "enc.gguf"}) + Expect(err).ToNot(HaveOccurred()) + matches, err := reg.Identify(ctx, []float32{1, 0}, 1) + Expect(err).ToNot(HaveOccurred()) + Expect(matches).To(HaveLen(1)) + Expect(matches[0].Metadata.Model).To(Equal("enc.gguf")) + }) +}) diff --git a/core/services/voicerecognition/voicerecognition_suite_test.go b/core/services/voicerecognition/voicerecognition_suite_test.go new file mode 100644 index 000000000..3ebd43584 --- /dev/null +++ b/core/services/voicerecognition/voicerecognition_suite_test.go @@ -0,0 +1,13 @@ +package voicerecognition_test + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestVoiceRecognition(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "VoiceRecognition Suite") +} diff --git a/docs/content/features/audio-diarization.md b/docs/content/features/audio-diarization.md index 37f8c8159..53e5e487f 100644 --- a/docs/content/features/audio-diarization.md +++ b/docs/content/features/audio-diarization.md @@ -79,6 +79,26 @@ Adds per-speaker totals and (when the backend supports it and `include_text=true } ``` +### Speaker names + +With a parakeet-cpp model that has a `speaker_model:` and voices registered through `/v1/voice/register`, segments whose speaker matches a registered voice gain `name` and `name_score` (the cosine similarity of the match), and the matching `speakers` entry gains `name`. Both fields are omitted for a speaker that was not identified, so an unnamed response looks exactly as before. `speaker` stays `SPEAKER_NN`, and RTTM output still uses `SPEAKER_NN`. See [Voice Recognition]({{% relref "voice-recognition" %}}#naming-speakers-in-diarization-and-live-transcription) for the setup and the limits. + +```json +{ + "task": "diarize", + "duration": 12.34, + "num_speakers": 2, + "segments": [ + {"id": 0, "speaker": "SPEAKER_00", "label": "0", "start": 0.00, "end": 2.34, "text": "Hello, world.", "name": "Alice", "name_score": 0.82}, + {"id": 1, "speaker": "SPEAKER_01", "label": "1", "start": 2.34, "end": 4.10, "text": "How are you?"} + ], + "speakers": [ + {"id": "SPEAKER_00", "label": "0", "name": "Alice", "total_speech_duration": 5.6, "segment_count": 3}, + {"id": "SPEAKER_01", "label": "1", "total_speech_duration": 1.76, "segment_count": 1} + ] +} +``` + ### Response - `rttm` NIST RTTM, the standard interchange format used by `pyannote.metrics` / `dscore`: @@ -160,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 @@ -195,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/audio-to-text.md b/docs/content/features/audio-to-text.md index d01edec4e..e90423749 100644 --- a/docs/content/features/audio-to-text.md +++ b/docs/content/features/audio-to-text.md @@ -112,6 +112,8 @@ In addition to `file` and `model`, the endpoint accepts the following multipart | `stream` | When `true`, the endpoint emits an SSE stream of `transcript.text.delta` events followed by a final `transcript.text.done` event. | | `diarize` | LocalAI extension - speaker diarization. WhisperX requires `HF_TOKEN`; requests fail with `FailedPrecondition` when it is missing. | +If speaker diarization fails after transcription succeeded, the WhisperX backend logs the error and returns the transcript without speaker labels. Other transcription failures return an error instead of an empty transcript. Diarization still requires `HF_TOKEN`. + The response body for `verbose_json` includes `text`, `language`, `duration`, and `segments[]` (with `speaker` populated when diarization is enabled). ## Streaming transcriptions @@ -200,9 +202,14 @@ The same backend also serves the `/v1/audio/diarization` and `/v1/audio/classifi | `diarization_model:` | an ASR model | a `speaker` on transcript segments (and words), and speaker segments during realtime live transcription | | `sound_model:` | an ASR model | sound events during realtime live transcription | | `diarization_latency:` | a model with a diarization companion | latency mode for the live speaker stream; default `low` | +| `speaker_model:` | a model with a diarization model | names registered speakers (see [Voice Recognition]({{% relref "voice-recognition" %}}#naming-speakers-in-diarization-and-live-transcription)) | +| `speaker_threshold:` | a model with `speaker_model` | distance (1 minus cosine similarity) under which a speaker is named, in (0, 2); default `0.5` | +| `speaker_margin:` | a model with `speaker_model` | how much the best match must beat the runner-up, in [0, 1); default `0.05` | With a `diarization_model` companion, `/v1/audio/transcriptions` labels each segment with its `speaker` (`"0"`, `"1"`, ... in order of first appearance) and splits segments where the speaker changes; with `timestamp_granularities[]=word` each word carries its speaker too. With `stream=true` the closing `transcript.text.done` event lists the segments with their speakers. Pass `-F diarize=false` to skip diarization for one request. The diarization GGUF can also be imported directly: `local-ai models import https://huggingface.co/mudler/parakeet-cpp-gguf/resolve/main/nemotron-3-diarization-f16.gguf`. +`speaker_model:` needs libparakeet with C-API v10. A wrong setup fails at load time with one of these errors: `parakeet-cpp: speaker_model needs libparakeet.so ABI 10 (parakeet_capi_speaker_registry_add_embedding); the loaded library is older`, `parakeet-cpp: speaker_model needs a diarization model (the primary or diarization_model:)`, `parakeet-cpp: a speaker model cannot be the primary model; use it as speaker_model: next to a diarization model`, `parakeet-cpp: speaker_model "" is a model, expected a speaker model` (the file is not a speaker encoder GGUF), or `parakeet-cpp: speaker_threshold "" must be a distance in (0, 2) (1 minus cosine similarity)` / `parakeet-cpp: speaker_margin "" must be a number in [0, 1)` for a bad number. + The loader rejects a companion whose role duplicates the primary's own (for example `asr_model:` on an already-ASR primary, or `sound_model:` on a CED primary), and rejects a companion GGUF that does not match the role its option names (for example `sound_model:` pointing at an ASR GGUF fails to load, naming the kind it expected). See [Speaker Diarization]({{% relref "audio-diarization" %}}) for the `Diarize` RPC and [Sound Classification]({{% relref "audio-classification" %}}) for `SoundDetection`, and [Realtime API]({{% relref "openai-realtime" %}}) for the live speaker/sound events emitted during a realtime session. ### Segment timestamps diff --git a/docs/content/features/decisions.md b/docs/content/features/decisions.md index 2feb30c44..115403b5e 100644 --- a/docs/content/features/decisions.md +++ b/docs/content/features/decisions.md @@ -108,11 +108,19 @@ Install one from the gallery and filter on the `decisions` tag: | `tev1-4b-vllm-cpp` | Tev1 4B | Autoregressive Qwen3.5-4B fine-tune that answers with an option letter, about 9.3 GB | | `tev1-0.8b-vllm-cpp` | Tev1 0.8B | Autoregressive Qwen3.5-0.8B fine-tune that answers with an option letter, about 1.8 GB | | `kev-0.8b-vllm-cpp` | kev 0.8B | Qwen3.5-0.8B-Base with a merged LoRA and a PointerHead readout, converted for vllm.cpp only, about 1.53 GB | +| `nimble-9b-vllm-cpp` | Bespoke Nimble 9B | Qwen3.5-9B with the Nimble LoRA merged, reads the answer-letter logits, converted for vllm.cpp only, about 19.3 GB | +| `clm-v0.1-8b-vllm-cpp` | CLM v0.1 8B | Bi-encoder: Qwen3-8B backbone with state and action heads, answers by cosine similarity, converted for vllm.cpp only, about 16.5 GB | The engine, [vllm.cpp]({{% relref "features/vllm-cpp" %}}), also supports the -CLM and xor decision models. Those checkpoints need a conversion step, so -they are not gallery entries yet. The kev entry installs a checkpoint that was -already converted with the vllm.cpp `convert-kev.py` script. +xor decision model. That checkpoint needs a conversion step, so it is not a +gallery entry yet. The kev, Nimble and CLM entries install checkpoints that +were already converted with the vllm.cpp `convert-kev.py`, `convert-nimble.py` +and `convert-clm.py` scripts. + +Nimble refuses a question with more than 26 choices (the upstream release +allows 255). On CPU it needs about 20 GB of free RAM. For CLM, put the +question in `instructions`: the state head reads the state followed by the +instructions. On CPU it needs about 19 GB of free RAM. Tev1 is an autoregressive decision model. The engine answers each question by scoring the option letters, so its `confidence` is the entropy measure Ollama diff --git a/docs/content/features/model-gallery.md b/docs/content/features/model-gallery.md index 09063fdd7..3bdfc6d41 100644 --- a/docs/content/features/model-gallery.md +++ b/docs/content/features/model-gallery.md @@ -39,6 +39,27 @@ Both views use the same model selection and store the view, search, filter, and selection in the URL. Installing from Explore does not move you away from the catalog; the entry updates in place when the operation finishes. +## Cyber-Ornith 1.5 9B + +Cyber-Ornith 1.5 is a Qwen3.5 fine-tune for security auditing, terminal tasks, and tool use. +The gallery includes Q4_K_M and Q6_K GGUF builds for text chat with llama.cpp. +Both use the embedded chat template and a 32,768-token context by default. + +Install with automatic variant selection: + +```bash +local-ai models install cyber-ornith-1.5-9b-obliterated +``` + +To select Q6_K explicitly: + +```bash +local-ai models install cyber-ornith-1.5-9b-obliterated --variant cyber-ornith-1.5-9b-obliterated-q6 +``` + +See the [model card](https://huggingface.co/DuoNeural/Cyber-Ornith-1.5-9B-OBLITERATED) +and [GGUF downloads](https://huggingface.co/mradermacher/Cyber-Ornith-1.5-9B-OBLITERATED-i1-GGUF). + ## Cyber-Tiel-Coder Install `cyber-tiel-coder-35b-a3b-q4-mtp` for coding and image chat with llama.cpp. diff --git a/docs/content/features/openai-realtime.md b/docs/content/features/openai-realtime.md index 150966cc4..f66f498f2 100644 --- a/docs/content/features/openai-realtime.md +++ b/docs/content/features/openai-realtime.md @@ -166,12 +166,15 @@ Each closed speaker segment emits a `conversation.item.input_audio_transcription "item_id": "item_abc", "content_index": 0, "speaker": "0", + "speaker_name": "Alice", "start": 1.92, "end": 4.10, "text": "" } ``` +`speaker_name` is the name of a voice registered through `/v1/voice/register`, and is present only when the model has a `speaker_model:` and the speaker was identified. A segment that closes before its speaker is identified has none, and later segments of the same speaker do. See [Voice Recognition]({{% relref "voice-recognition" %}}#naming-speakers-in-diarization-and-live-transcription). The segments of the offline path below carry no `speaker_name`. + Each sound event emits a `conversation.item.sound_detection` event with one tag and the detection window's `start`/`end`: ```json diff --git a/docs/content/features/text-generation.md b/docs/content/features/text-generation.md index 4c558d8d1..ee10ccfdc 100644 --- a/docs/content/features/text-generation.md +++ b/docs/content/features/text-generation.md @@ -33,6 +33,8 @@ Available additional parameters: `top_p`, `top_k`, `max_tokens` Reasoning models return their thinking in the `reasoning` field. When a model reasons and calls a tool in the same turn, see [Interleaved Thinking with Tool Calls]({{%relref "features/interleaved-thinking" %}}). +When `stream: true` is set and the llama.cpp backend fails before the first chunk, for example because the prompt exceeds the context size, the request fails with an HTTP error. The error message is not streamed as assistant content. An error after streaming has started is reported inside the stream. + ### Edit completions https://platform.openai.com/docs/api-reference/edits @@ -627,6 +629,8 @@ options: **Note:** The `parallel` option can also be set via the `LLAMACPP_PARALLEL` environment variable, and `grpc_servers` can be set via the `LLAMACPP_GRPC_SERVERS` environment variable. Options specified in the YAML file take precedence over environment variables. +An explicit `parallel: 1` (or `n_parallel: 1`) in the model options takes precedence over `LLAMACPP_PARALLEL`, like any other value. The environment variable is only used when neither option is set; if it is missing or not a number, the backend uses one slot. + ##### Hardware auto-tuning (and how to override it) On a detected GPU, LocalAI fills a few performance-relevant defaults the model config leaves unset - a larger physical batch on NVIDIA Blackwell, and a VRAM-scaled `parallel` slot count for concurrent serving. Both are gated on **per-device** VRAM at the model's context: when a large context already fills a single card (e.g. a 27B model with a 200k context across 2×16 GiB), the batch boost and the extra parallel slots are suppressed so they can't tip the tighter GPU into CUDA out-of-memory. diff --git a/docs/content/features/voice-recognition.md b/docs/content/features/voice-recognition.md index 23bb16bd6..d3c866ba0 100644 --- a/docs/content/features/voice-recognition.md +++ b/docs/content/features/voice-recognition.md @@ -185,6 +185,85 @@ recognition - the voice-recognition HTTP API is designed to swap the backing store without changing the wire format. {{% /notice %}} +## Naming speakers in diarization and live transcription + +The parakeet-cpp backend can put the names of registered voices on +diarization results and on live transcription speaker segments. Without +this, speakers only carry labels such as `SPEAKER_00`. + +1. Register each voice with the WeSpeaker encoder. Install the model with + `local-ai models install voice-detect-wespeaker-resnet34`, then call + `/v1/voice/register` with `"model": "voice-detect-wespeaker-resnet34"` + (see the [1:N workflow](#1n-identification-workflow-register--identify--forget)). +2. Install one of the gallery models that loads the same encoder: + `parakeet-cpp-nemotron-3-diarization-speakers` (diarization), + `parakeet-cpp-nemotron-3-diarization-asr-speakers` (diarization with + `include_text`) or `parakeet-cpp-realtime-scene-speakers` (live + transcription). Each one adds + `speaker_model:voice-detect-wespeaker-resnet34.gguf` to a + parakeet-cpp model config. +3. Call `/v1/audio/diarization` with that model. Matched segments gain a + `name` and a `name_score`, and the matching entry in `speakers` gains a + `name`. `speaker` stays `SPEAKER_NN`, and RTTM output is unchanged. See + [Speaker Diarization]({{% relref "audio-diarization" %}}) for the + response. + +### Which voices are used + +LocalAI sends the backend only the registered voices made by the same +encoder as the model's `speaker_model:` file. Each registered voice is +tagged with the name of the voice-detect model that made it, which by +default is the GGUF file name (`voice-detect-wespeaker-resnet34.gguf` for the +gallery entry). The tag must equal the base name of the `speaker_model:` +file. Voices made with another encoder are ignored, and LocalAI logs a +warning when that leaves no usable voice. Voices registered before the tag +existed have no tag: they are used when their embedding size matches the +tagged ones (or all of them, when no voice carries a matching tag). The +backend skips a voice whose embedding size does not match the speaker model's, +with a warning in the LocalAI log. Naming then falls back to the remaining +voices, or to no names. + +{{% notice warning %}} +Do not set a `model_name:` option on the voice-detect model config. It +replaces the default name, the voices are then tagged with it, and they no +longer match the `speaker_model:` file. Keep the default name. +{{% /notice %}} + +### Options + +These go in the `options:` list of the parakeet-cpp model config (see +[Audio to Text]({{% relref "audio-to-text" %}}) for the other parakeet-cpp +options). + +| Option | Default | Meaning | +|---|---|---| +| `speaker_model:` | none | speaker encoder GGUF; needs a diarization model (the primary one, or `diarization_model:`) | +| `speaker_threshold:` | `0.5` | largest distance (1 minus cosine similarity, the unit `/v1/voice/identify` reports) at which a speaker is named; must be in (0, 2) | +| `speaker_margin:` | `0.05` | the best match must beat the runner-up by this much, otherwise the speaker stays unnamed; must be in [0, 1) | + +parakeet.cpp's measured starting values for `speaker_threshold` are 0.5 for +WeSpeaker ResNet34 and CAM++, and 0.3 for ECAPA. A lower value names fewer +speakers and makes fewer mistakes. + +### Limits + +- The voice registry is in memory and global. Registered names disappear when + LocalAI restarts, and every user of the instance shares them. +- Anyone who is allowed to call a model with `speaker_model:` can learn which + registered names match their audio, and their audio is matched against voices + registered by any user, because the voice registry is global. Restrict such + models with the per-user model allowlist. +- With `include_text=true` the names use the default threshold and margin: + `speaker_threshold` and `speaker_margin` only apply to diarization without + text. +- In live transcription, a speaker segment that closes before its speaker + is identified has no name. Later segments of that speaker do. +- Overlapping speech is not resolved. +- Accuracy was measured on one fixture (two read-speech voices). Check the + threshold on your own audio. +- The backend needs a libparakeet with C-API v10. With an older library a + model config that sets `speaker_model:` fails to load. + ## API reference ### `POST /v1/voice/verify` (1:1) @@ -323,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/docs/content/integrations.md b/docs/content/integrations.md index 6fb84a250..ca5ad73ea 100644 --- a/docs/content/integrations.md +++ b/docs/content/integrations.md @@ -74,9 +74,9 @@ availability may lag upstream releases. ### Chat Bots -- [Discord bot](https://github.com/mudler/LocalAGI/tree/main/examples/discord) -- [Slack bot](https://github.com/mudler/LocalAGI/tree/main/examples/slack) -- [Telegram bot](https://github.com/mudler/LocalAI/tree/master/examples/telegram-bot) +- [Discord bot](https://github.com/mudler/LocalAI-examples/tree/main/discord-bot) +- [Slack bot](https://github.com/mudler/LocalAI-examples/tree/main/slack-bot) +- [Telegram bot](https://github.com/mudler/LocalAI-examples/tree/main/telegram-bot) - [Hellper (Telegram)](https://github.com/JackBekket/Hellper) ### Home Automation diff --git a/docs/content/operations/backend-monitor.md b/docs/content/operations/backend-monitor.md index 1fa1750bf..359bab8f2 100644 --- a/docs/content/operations/backend-monitor.md +++ b/docs/content/operations/backend-monitor.md @@ -123,6 +123,8 @@ curl -X POST http://localhost:8080/backend/shutdown \ Returns `200 OK` with the shutdown confirmation message on success. +Stopping a backend removes its watchdog timers and eviction state. A timeout from a stopped backend does not shut down a replacement at a different address. + ## Error Responses | Status Code | Description | diff --git a/docs/content/operations/cloud-proxy.md b/docs/content/operations/cloud-proxy.md index 312fa2327..42f03bb98 100644 --- a/docs/content/operations/cloud-proxy.md +++ b/docs/content/operations/cloud-proxy.md @@ -198,6 +198,8 @@ image blocks, and per-request usage tokens are dropped through the internal `Predict()` signature. Use passthrough mode when your clients need the upstream's full feature set. +In translate mode, an Anthropic response with `stop_reason: "refusal"` is returned to the client as an error instead of an empty successful reply, for both non-streaming and streaming requests. A streamed response may already have delivered partial content when the refusal arrives. Responses that end normally (`end_turn`) are unaffected, even when their content is empty. + #### Anthropic prompt caching `proxy.cache_prompt: true` makes the translator add Anthropic diff --git a/gallery/index.yaml b/gallery/index.yaml index eb37a7d12..942d932a9 100644 --- a/gallery/index.yaml +++ b/gallery/index.yaml @@ -89,8 +89,84 @@ sha256: ad5811e291431bd0de1cec0c4004a5eac98daee9850882edac69a823209e88ab uri: https://huggingface.co/ukisai/Swift-Qwen3.8-27B-GGUF/resolve/main/Swift-Qwen3.8-27B-Q4_K_M.gguf - filename: llama-cpp/mmproj/Swift-Qwen3.8-27B-Q4_K_M/mmproj-Swift-Qwen3.8-27B-F16.gguf - sha256: daa1116c9422fa390cc8688495da0e91781f92841dfc3b31a378ff252571745a uri: https://huggingface.co/ukisai/Swift-Qwen3.8-27B-GGUF/resolve/main/mmproj-Swift-Qwen3.8-27B-F16.gguf + sha256: 9f3a5e81486b0c0270eff405b9a22911bf388aa7b58f3788eed9b0774904283b +- name: cyber-ornith-1.5-9b-obliterated + variants: + - model: cyber-ornith-1.5-9b-obliterated-q6 + url: "github:mudler/LocalAI/gallery/virtual.yaml@master" + urls: + - https://huggingface.co/DuoNeural/Cyber-Ornith-1.5-9B-OBLITERATED + - https://huggingface.co/mradermacher/Cyber-Ornith-1.5-9B-OBLITERATED-i1-GGUF + description: | + Cyber-Ornith 1.5 is a 9B Qwen3.5 fine-tune for security auditing, terminal tasks, and tool use. + This importance-matrix Q4_K_M GGUF build uses llama.cpp with the embedded chat template for text chat. + license: apache-2.0 + tags: + - llm + - gguf + - cpu + - gpu + - coding + last_checked: "2026-09-30" + overrides: + backend: llama-cpp + context_size: 32768 + function: + automatic_tool_parsing_fallback: true + grammar: + disable: true + known_usecases: + - chat + options: + - use_jinja:true + template: + use_tokenizer_template: true + parameters: + model: Cyber-Ornith-1.5-9B-OBLITERATED.i1-Q4_K_M.gguf + temperature: 0.6 + top_p: 0.9 + files: + - filename: Cyber-Ornith-1.5-9B-OBLITERATED.i1-Q4_K_M.gguf + sha256: 5b7bdec6919db868e517e9ab9a05ce8be096f6f139798a7b94a3aa3d2d1eab28 + uri: https://huggingface.co/mradermacher/Cyber-Ornith-1.5-9B-OBLITERATED-i1-GGUF/resolve/5936e88f1885b856b54520c372f2cd41e71819fd/Cyber-Ornith-1.5-9B-OBLITERATED.i1-Q4_K_M.gguf +- name: cyber-ornith-1.5-9b-obliterated-q6 + url: "github:mudler/LocalAI/gallery/virtual.yaml@master" + urls: + - https://huggingface.co/DuoNeural/Cyber-Ornith-1.5-9B-OBLITERATED + - https://huggingface.co/mradermacher/Cyber-Ornith-1.5-9B-OBLITERATED-i1-GGUF + description: | + Cyber-Ornith 1.5 is a 9B Qwen3.5 fine-tune for security auditing, terminal tasks, and tool use. + This importance-matrix Q6_K GGUF build uses llama.cpp with the embedded chat template for text chat. + license: apache-2.0 + tags: + - llm + - gguf + - cpu + - gpu + - coding + last_checked: "2026-09-30" + overrides: + backend: llama-cpp + context_size: 32768 + function: + automatic_tool_parsing_fallback: true + grammar: + disable: true + known_usecases: + - chat + options: + - use_jinja:true + template: + use_tokenizer_template: true + parameters: + model: Cyber-Ornith-1.5-9B-OBLITERATED.i1-Q6_K.gguf + temperature: 0.6 + top_p: 0.9 + files: + - filename: Cyber-Ornith-1.5-9B-OBLITERATED.i1-Q6_K.gguf + sha256: 0a43704d7b18c2eb62fdc6d67b4cf32d5fd4d7a147e1bf0b00d3dc75d88822fd + uri: https://huggingface.co/mradermacher/Cyber-Ornith-1.5-9B-OBLITERATED-i1-GGUF/resolve/5936e88f1885b856b54520c372f2cd41e71819fd/Cyber-Ornith-1.5-9B-OBLITERATED.i1-Q6_K.gguf - name: "ornith-1.5-9b-uncensored" url: "github:mudler/LocalAI/gallery/virtual.yaml@master" urls: @@ -297,7 +373,7 @@ files: - filename: ds4flash.gguf uri: https://huggingface.co/unsloth/DeepSeek-V4-Flash-Vision-Exp-GGUF - sha256: bfca14d287c9fd865529e02efa0ba6f572fb7bef623fa0bf4b0fb6214ee63a5e + sha256: 69bd869c7924086f5d8a3a86ca9663a1f331b7bd231121262e1989c59ee2179f - name: "qwopus3.8-27b-flash-v2" variants: - model: qwopus3.8-27b-flash-v2-q8 @@ -53768,6 +53844,167 @@ - filename: parakeet-cpp/tdt_ctc-110m-f16.gguf uri: huggingface://mudler/parakeet-cpp-gguf/tdt_ctc-110m-f16.gguf sha256: 7f9a6376edde6a74592ace48b2ebdc27a1ac972d0be9dfcc29e668d99381faf1 +- name: parakeet-cpp-nemotron-3-diarization-speakers + url: github:mudler/LocalAI/gallery/virtual.yaml@master + urls: + - https://huggingface.co/mudler/parakeet-cpp-gguf + - https://huggingface.co/mudler/voice-detect-gguf + - https://huggingface.co/nvidia/Nemotron-3-Diarization + - https://github.com/mudler/parakeet.cpp + description: | + Nemotron-3-Diarization (Sortformer) with WeSpeaker ResNet34 speaker + identification, for the parakeet-cpp backend. Speakers you register with + /v1/voice/register (using the voice-detect-wespeaker-resnet34 model) come back + by name in /v1/audio/diarization, next to the SPEAKER_NN label. Speakers that + are not registered keep only their SPEAKER_NN label. The diarization model is + OpenMDW-1.1, the speaker model is CC-BY-4.0. Naming was measured on one + two-voice fixture only; check the threshold on your own audio. + license: openmdw-1.1 + tags: + - parakeet + - parakeet-cpp + - nemotron + - sortformer + - diarization + - speaker-diarization + - speaker-recognition + - gguf + - ggml + - quantized + overrides: + backend: parakeet-cpp + known_usecases: + - diarization + name: parakeet-cpp-nemotron-3-diarization-speakers + options: + - speaker_model:voice-detect-wespeaker-resnet34.gguf + parameters: + model: parakeet-cpp/nemotron-3-diarization-q8_0.gguf + files: + - filename: parakeet-cpp/nemotron-3-diarization-q8_0.gguf + uri: huggingface://mudler/parakeet-cpp-gguf/nemotron-3-diarization-q8_0.gguf + sha256: 76c5bb1fb20d82706142ad32769b7ab496d2458489473a000fd7074c52ceec22 + - filename: voice-detect-wespeaker-resnet34.gguf + uri: https://huggingface.co/mudler/voice-detect-gguf/resolve/main/wespeaker-resnet34-voxceleb.gguf + sha256: 72040372494eafec299836bc1977cfc13c603cb486674ed59b0f4c03758d29da +- name: parakeet-cpp-nemotron-3-diarization-asr-speakers + url: github:mudler/LocalAI/gallery/virtual.yaml@master + urls: + - https://huggingface.co/mudler/parakeet-cpp-gguf + - https://huggingface.co/mudler/voice-detect-gguf + - https://huggingface.co/nvidia/Nemotron-3-Diarization + - https://huggingface.co/nvidia/parakeet-tdt_ctc-110m + - https://github.com/mudler/parakeet.cpp + description: | + Nemotron-3-Diarization (Sortformer) paired with the Parakeet TDT+CTC 110M + ASR model through the asr_model option, both Q8_0/F16 GGUF for the + parakeet-cpp backend (C++/ggml port of NVIDIA NeMo). Served through + /v1/audio/diarization with include_text: each speaker segment comes back + with its transcribed text in one call. Diarization model is + OpenMDW-1.1, ASR model is CC-BY-4.0. Also loads WeSpeaker ResNet34 + (CC-BY-4.0) through the speaker_model option: speakers registered with + /v1/voice/register (voice-detect-wespeaker-resnet34 model) come back by + name, next to the SPEAKER_NN label. + license: openmdw-1.1 + tags: + - parakeet + - parakeet-cpp + - nemotron + - sortformer + - asr + - diarization + - speaker-diarization + - speech-recognition + - stt + - gguf + - ggml + - quantized + overrides: + backend: parakeet-cpp + known_usecases: + - diarization + name: parakeet-cpp-nemotron-3-diarization-asr-speakers + options: + - asr_model:parakeet-cpp/tdt_ctc-110m-f16.gguf + - speaker_model:voice-detect-wespeaker-resnet34.gguf + parameters: + model: parakeet-cpp/nemotron-3-diarization-q8_0.gguf + files: + - filename: parakeet-cpp/nemotron-3-diarization-q8_0.gguf + uri: huggingface://mudler/parakeet-cpp-gguf/nemotron-3-diarization-q8_0.gguf + sha256: 76c5bb1fb20d82706142ad32769b7ab496d2458489473a000fd7074c52ceec22 + - filename: parakeet-cpp/tdt_ctc-110m-f16.gguf + uri: huggingface://mudler/parakeet-cpp-gguf/tdt_ctc-110m-f16.gguf + sha256: 7f9a6376edde6a74592ace48b2ebdc27a1ac972d0be9dfcc29e668d99381faf1 + - filename: voice-detect-wespeaker-resnet34.gguf + uri: https://huggingface.co/mudler/voice-detect-gguf/resolve/main/wespeaker-resnet34-voxceleb.gguf + sha256: 72040372494eafec299836bc1977cfc13c603cb486674ed59b0f4c03758d29da +- name: parakeet-cpp-realtime-scene-speakers + url: github:mudler/LocalAI/gallery/virtual.yaml@master + urls: + - https://huggingface.co/mudler/parakeet-cpp-gguf + - https://huggingface.co/mudler/ced-gguf + - https://huggingface.co/mudler/voice-detect-gguf + - https://huggingface.co/nvidia/parakeet_realtime_eou_120m-v1 + - https://huggingface.co/nvidia/Nemotron-3-Diarization + - https://huggingface.co/mispeech/ced-tiny + - https://github.com/mudler/parakeet.cpp + description: | + Cache-aware streaming RNNT FastConformer with end-of-utterance (EOU) + detection, 120M, paired with Nemotron-3-Diarization and CED-Tiny through + the diarization_model and sound_model options. F16/Q8_0 GGUF for the + parakeet-cpp backend (C++/ggml port of NVIDIA NeMo). Use with streaming + transcription: while a turn is live, closed speaker segments and sound + events are surfaced alongside the ASR text (realtime + conversation.item.input_audio_transcription.segment and + conversation.item.sound_detection events). Live speaker/sound events only + fire during speech turns under semantic_vad; sounds between turns are not + seen by this path. License per model: transcription model NVIDIA Open + Model License, diarization model OpenMDW-1.1, CED-Tiny Apache-2.0, + WeSpeaker ResNet34 CC-BY-4.0. Also loads WeSpeaker ResNet34 through the + speaker_model option, so live speaker segments carry the name of a voice + registered with /v1/voice/register (voice-detect-wespeaker-resnet34 model) + once the speaker is identified. + license: nvidia-open-model-license + tags: + - parakeet + - parakeet-cpp + - nemotron + - sortformer + - ced + - asr + - speech-recognition + - diarization + - sound-classification + - streaming + - realtime + - stt + - gguf + - ggml + overrides: + backend: parakeet-cpp + known_usecases: + - transcript + name: parakeet-cpp-realtime-scene-speakers + options: + - diarization_model:parakeet-cpp/nemotron-3-diarization-q8_0.gguf + - sound_model:parakeet-cpp/ced-tiny-q8_0.gguf + - speaker_model:voice-detect-wespeaker-resnet34.gguf + parameters: + model: parakeet-cpp/realtime_eou_120m-v1-f16.gguf + files: + - filename: parakeet-cpp/realtime_eou_120m-v1-f16.gguf + uri: huggingface://mudler/parakeet-cpp-gguf/realtime_eou_120m-v1-f16.gguf + sha256: d1a2b12f12b8a096a57499c9111ed13b442a2b786e17a292c168be45088f0edc + - filename: parakeet-cpp/nemotron-3-diarization-q8_0.gguf + uri: huggingface://mudler/parakeet-cpp-gguf/nemotron-3-diarization-q8_0.gguf + sha256: 76c5bb1fb20d82706142ad32769b7ab496d2458489473a000fd7074c52ceec22 + - filename: parakeet-cpp/ced-tiny-q8_0.gguf + uri: huggingface://mudler/ced-gguf/ced-tiny-q8_0.gguf + sha256: 48bee4e2fc3cc85d7806e03471db24e77fda6c2a2e81ffe9ef67caebaf2bd674 + - filename: voice-detect-wespeaker-resnet34.gguf + uri: https://huggingface.co/mudler/voice-detect-gguf/resolve/main/wespeaker-resnet34-voxceleb.gguf + sha256: 72040372494eafec299836bc1977cfc13c603cb486674ed59b0f4c03758d29da - name: parakeet-cpp-ced-tiny url: github:mudler/LocalAI/gallery/virtual.yaml@master urls: @@ -63875,6 +64112,125 @@ type: huggingface repo: mudler/kev-0.8b-vllm-cpp revision: c17e73666ded1e9d284470eae7e0de9a27294e77 +- name: nimble-9b-vllm-cpp + url: github:mudler/LocalAI/gallery/virtual.yaml@master + urls: + - https://huggingface.co/mudler/Bespoke-Nimble-9B-vllm-cpp + - https://huggingface.co/bespokelabs/Bespoke-Nimble-9B + - https://github.com/mudler/vllm.cpp + description: | + Nimble is a decision model from Bespoke Labs (the model Ollama serves as + nimble). For each question it runs one forward pass and reads the logits + of the answer letters at the last prompt position. It answers typed + choice, noul and score questions about a text state and does not + generate text. + + This entry installs a converted redistribution of + bespokelabs/Bespoke-Nimble-9B: the LoRA adapter is merged into its + Qwen3.5-9B base in BF16 and config.json names the NimbleModel + architecture. The checkpoint only works with vllm.cpp (the vllm-cpp + backend); transformers and vLLM cannot load it. + + In LocalAI, serve it via POST /v1/systemone. The vllm.cpp project compared + this converted directory on CPU against the Bespoke authors' own code over + seven questions: 7 of 7 answers equal, largest probability difference + 0.0054. This entry was installed and served through LocalAI on CPU and + gave the answers shown in the model card example. There is no accuracy + benchmark and no GPU (CUDA, ROCm, Metal) run. The engine refuses fields + with more than 26 choices (upstream allows 255). + + The entry sets an 8192-token context (Nimble's own prompt limit) and a + KV pool that holds four such sequences. BF16 weights, + about 19.3 GB, pinned to a revision. On CPU the model needs about 20 GB + of free RAM (measured peak 18.4 GB resident). On a 20-thread CPU a + request with three questions (988 prompt tokens) took 30 seconds warm. + license: apache-2.0 + tags: + - decisions + - systemone + - vllm-cpp + - cpu + - gpu + size: 19.33GB + last_checked: "2026-10-01" + overrides: + backend: vllm-cpp + known_usecases: + - decisions + context_size: 8192 + engine_args: + block_size: 32 + num_blocks: 1024 + max_num_seqs: 4 + parameters: + model: mudler/Bespoke-Nimble-9B-vllm-cpp + artifacts: + - name: model + target: model + source: + type: huggingface + repo: mudler/Bespoke-Nimble-9B-vllm-cpp + revision: 52eead25b6723710ec16378942c9ae914d179b6f +- name: clm-v0.1-8b-vllm-cpp + url: github:mudler/LocalAI/gallery/virtual.yaml@master + urls: + - https://huggingface.co/mudler/CLM-v0.1-8B-vllm-cpp + - https://huggingface.co/Contrastive-LM/CLM-v0.1-8B + - https://github.com/mudler/vllm.cpp + description: | + CLM is a bi-encoder decision model from Contrastive-LM. A frozen Qwen3-8B + backbone encodes the state and each candidate answer separately, two MLP + heads project them into a 512-dimensional space, and the answer + distribution is a softmax over the cosine similarities at scale 100. It + answers typed choice, noul and score questions and does not generate + text. + + This entry installs a converted redistribution: the CLM-v0.1-8B heads as + head.safetensors next to the unchanged Qwen3-8B backbone and tokenizer, + with config.json naming the ClmModel architecture. The checkpoint only + works with vllm.cpp (the vllm-cpp backend); transformers, vLLM and + llama.cpp cannot load it. + + In LocalAI, serve it via POST /v1/systemone. Put the question in + instructions: the state head reads the state followed by the + instructions. The vllm.cpp project compared this checkpoint on CPU with the + reference code over eight questions: 8 of 8 answers agree, largest + probability difference 0.029 against the reference in bf16 (transformers + stood in for the reference's GPU encoder). This entry was installed and + served through LocalAI on CPU and gave the model card's example answer + (person, 0.950). There is no accuracy benchmark and no GPU run. + + The entry sets a 4096-token context and a KV pool that holds four such + sequences. BF16 backbone with F32 heads, about 16.5 GB, pinned to a + revision. On CPU the model needs about 19 GB of free RAM (measured peak + 17.9 GB resident). + license: apache-2.0 + tags: + - decisions + - systemone + - vllm-cpp + - cpu + - gpu + size: 16.47GB + last_checked: "2026-10-01" + overrides: + backend: vllm-cpp + known_usecases: + - decisions + context_size: 4096 + engine_args: + block_size: 32 + num_blocks: 512 + max_num_seqs: 4 + parameters: + model: mudler/CLM-v0.1-8B-vllm-cpp + artifacts: + - name: model + target: model + source: + type: huggingface + repo: mudler/CLM-v0.1-8B-vllm-cpp + revision: 0d1903b14024c721594e1794aa48818f20f6bdb3 - name: qwen3-vl-4b-vllm-cpp url: github:mudler/LocalAI/gallery/virtual.yaml@master urls: diff --git a/pkg/model/loader.go b/pkg/model/loader.go index 12d46910a..ae0d25fe8 100644 --- a/pkg/model/loader.go +++ b/pkg/model/loader.go @@ -616,6 +616,32 @@ func (ml *ModelLoader) ShutdownModelForce(modelName string) error { return ml.shutdownModel(ctx, modelName, true) } +// ShutdownModelAtAddress ignores stale watchdog evictions after a reload. +// Address validation and teardown share the same lifecycle lock as loading. +func (ml *ModelLoader) ShutdownModelAtAddress(modelName, address string, force bool) error { + ctx, cancel := context.WithTimeout(context.Background(), gracefulShutdownTimeout) + defer cancel() + release, err := ml.operations.acquireContext(ctx, modelName, true) + if err != nil { + return fmt.Errorf("waiting to shut down model %q: %w", modelName, err) + } + defer release() + ml.mu.Lock() + store := ml.store + ml.mu.Unlock() + m, ok := store.Get(modelName) + if !ok || m.address != address { + return nil + } + err = ml.deleteProcess(ctx, modelName, force) + if errors.Is(err, ErrModelBusy) && !force && forceBackendShutdown { + forceCtx, forceCancel := context.WithTimeout(context.Background(), forcedShutdownTimeout) + defer forceCancel() + return ml.deleteProcess(forceCtx, modelName, true) + } + return err +} + // ShutdownModelContext is the cancellation-aware lifecycle primitive used by // both graceful and forced shutdown wrappers. func (ml *ModelLoader) ShutdownModelContext(ctx context.Context, modelName string, force bool) error { diff --git a/pkg/model/process.go b/pkg/model/process.go index 9ec13bb51..36bff8c80 100644 --- a/pkg/model/process.go +++ b/pkg/model/process.go @@ -164,6 +164,9 @@ func (ml *ModelLoader) deleteProcess(ctx context.Context, s string, force bool) // at a known-unreachable worker, while the distributed registry remains // the source of truth for anything that is still running remotely. store.Delete(s) + if wd != nil { + wd.Untrack(model.address) + } return unloadErr } @@ -310,6 +313,7 @@ func (ml *ModelLoader) startProcess(grpcProcess, id string, serverAddress string } if err := grpcControlProcess.Run(); err != nil { + ml.untrackProcess(grpcControlProcess) runtime.cleanup() return grpcControlProcess, err } @@ -366,6 +370,7 @@ func (ml *ModelLoader) startProcess(grpcProcess, id string, serverAddress string // whether the child is alive. go func() { <-grpcControlProcess.Done() + ml.untrackProcess(grpcControlProcess) // LoadAndDelete both reads the intentional-stop marker and frees the // map entry so it doesn't accumulate across the process's lifetime. _, intentional := ml.stoppingProcs.LoadAndDelete(grpcControlProcess) @@ -403,6 +408,7 @@ func (ml *ModelLoader) cleanupProcessRuntime(process *process.Process) { if process == nil { return } + ml.untrackProcess(process) value, ok := ml.processRuntimes.LoadAndDelete(process) if !ok { return @@ -414,6 +420,18 @@ func (ml *ModelLoader) cleanupProcessRuntime(process *process.Process) { }() } +// Use the current watchdog because settings updates can replace it while a +// backend is running. Match the process identity so a late exit notification +// cannot remove a replacement that happens to reuse the same address. +func (ml *ModelLoader) untrackProcess(p *process.Process) { + ml.mu.Lock() + wd := ml.wd + ml.mu.Unlock() + if wd != nil { + wd.untrackProcess(p) + } +} + // CleanupProcessRuntime releases state and scratch owned by a process started // through StartProcess. Callers that supervise processes outside ModelLoader's // model store must invoke it after they have consumed exit diagnostics. diff --git a/pkg/model/watchdog.go b/pkg/model/watchdog.go index 3a604668a..977eca1ff 100644 --- a/pkg/model/watchdog.go +++ b/pkg/model/watchdog.go @@ -468,6 +468,7 @@ type modelUsageInfo struct { // opted into forceEvictionWhenBusy — either way, waiting for the graceful // deadline would defeat prompt eviction. type evictionTarget struct { + address string model string wasBusy bool } @@ -575,7 +576,7 @@ func (wd *WatchDog) collectEvictionsLocked(candidates []modelUsageInfo, maxToEvi continue } xlog.Info("[WatchDog] evicting model", "model", m.model, "busy", isBusy) - evicted = append(evicted, evictionTarget{model: m.model, wasBusy: isBusy}) + evicted = append(evicted, evictionTarget{address: m.address, model: m.model, wasBusy: isBusy}) wd.untrack(m.address) } return evicted, skippedBusy @@ -586,12 +587,7 @@ func (wd *WatchDog) collectEvictionsLocked(candidates []modelUsageInfo, maxToEvi // the graceful shutdown deadline. func (wd *WatchDog) shutdownEvicted(targets []evictionTarget, label string) { for _, t := range targets { - var err error - if t.wasBusy { - err = wd.pm.ShutdownModelForce(t.model) - } else { - err = wd.pm.ShutdownModel(t.model) - } + err := wd.shutdownTarget(t) if err != nil { xlog.Error("[WatchDog] error shutting down model during "+label, "error", err, "model", t.model, "busy", t.wasBusy) } @@ -599,6 +595,21 @@ func (wd *WatchDog) shutdownEvicted(targets []evictionTarget, label string) { } } +// shutdownTarget retains the address selected under the watchdog lock. The +// loader checks it under its per-model lifecycle lock so a queued eviction +// cannot stop a backend loaded after the original target exited. +func (wd *WatchDog) shutdownTarget(t evictionTarget) error { + if pm, ok := wd.pm.(interface { + ShutdownModelAtAddress(string, string, bool) error + }); ok { + return pm.ShutdownModelAtAddress(t.model, t.address, t.wasBusy) + } + if t.wasBusy { + return wd.pm.ShutdownModelForce(t.model) + } + return wd.pm.ShutdownModel(t.model) +} + // EnforceGroupExclusivity evicts every loaded model that shares at least one // concurrency group with the requested model. The pinned/busy/retry semantics // match EnforceLRULimit so the loader's retry loop can stay generic. @@ -713,7 +724,7 @@ func (wd *WatchDog) checkIdle() { xlog.Debug("[WatchDog] Watchdog checks for idle connections") // Collect models to shutdown while holding the lock - var modelsToShutdown []string + var modelsToShutdown []evictionTarget for address, t := range wd.idleTime { xlog.Debug("[WatchDog] idle connection", "address", address) if time.Since(t) > wd.idletimeout { @@ -723,8 +734,8 @@ func (wd *WatchDog) checkIdle() { xlog.Debug("[WatchDog] Skipping idle eviction for pinned model", "model", model) continue } - xlog.Warn("[WatchDog] Address is idle for too long, killing it", "address", address) - modelsToShutdown = append(modelsToShutdown, model) + xlog.Warn("[WatchDog] Address is idle for too long, killing it", "address", address, "model", model) + modelsToShutdown = append(modelsToShutdown, evictionTarget{model: model, address: address}) } else { xlog.Warn("[WatchDog] Address unresolvable", "address", address) } @@ -733,13 +744,7 @@ func (wd *WatchDog) checkIdle() { } wd.Unlock() - // Now shutdown models without holding the watchdog lock to prevent deadlock - for _, model := range modelsToShutdown { - if err := wd.pm.ShutdownModel(model); err != nil { - xlog.Error("[watchdog] error shutting down model", "error", err, "model", model) - } - xlog.Debug("[WatchDog] model shut down", "model", model) - } + wd.shutdownEvicted(modelsToShutdown, "idle timeout") } func (wd *WatchDog) checkBusy() { @@ -747,7 +752,7 @@ func (wd *WatchDog) checkBusy() { xlog.Debug("[WatchDog] Watchdog checks for busy connections") // Collect models to shutdown while holding the lock - var modelsToShutdown []string + var modelsToShutdown []evictionTarget for address, t := range wd.busyTime { xlog.Debug("[WatchDog] active connection", "address", address) @@ -755,7 +760,7 @@ func (wd *WatchDog) checkBusy() { model, ok := wd.addressModelMap[address] if ok { xlog.Warn("[WatchDog] Model is busy for too long, killing it", "model", model) - modelsToShutdown = append(modelsToShutdown, model) + modelsToShutdown = append(modelsToShutdown, evictionTarget{model: model, address: address, wasBusy: true}) } else { xlog.Warn("[WatchDog] Address unresolvable", "address", address) } @@ -764,16 +769,7 @@ func (wd *WatchDog) checkBusy() { } wd.Unlock() - // The busy-killer targets backends whose in-flight gRPC call has been - // stuck past the busy timeout. Use the force path so the loader stops - // the process FIRST (dropping the stuck call's gRPC connection) instead - // of waiting for the graceful shutdown deadline. - for _, model := range modelsToShutdown { - if err := wd.pm.ShutdownModelForce(model); err != nil { - xlog.Error("[watchdog] error shutting down model", "error", err, "model", model) - } - xlog.Debug("[WatchDog] busy model shut down", "model", model) - } + wd.shutdownEvicted(modelsToShutdown, "busy timeout") } // checkMemory monitors memory usage (GPU VRAM if available, otherwise RAM) and evicts backends when usage exceeds threshold @@ -896,11 +892,7 @@ func (wd *WatchDog) evictLRUModel() { wd.Unlock() // Shutdown the model - shutdown := wd.pm.ShutdownModel - if wasBusy { - shutdown = wd.pm.ShutdownModelForce - } - if err := shutdown(lruModel.model); err != nil && !errors.Is(err, ErrModelNotFound) { + if err := wd.shutdownTarget(evictionTarget{model: lruModel.model, address: lruModel.address, wasBusy: wasBusy}); err != nil && !errors.Is(err, ErrModelNotFound) { xlog.Error("[WatchDog] error shutting down model during memory reclamation", "error", err, "model", lruModel.model) } else { // Untrack the model @@ -913,7 +905,18 @@ func (wd *WatchDog) evictLRUModel() { func (wd *WatchDog) untrack(address string) { if modelID, ok := wd.addressModelMap[address]; ok { - delete(wd.modelSizes, modelID) + // A stale address can coexist with the replacement's registration. + // Its cleanup must not erase the replacement's size estimate. + otherAddress := false + for addr, name := range wd.addressModelMap { + if addr != address && name == modelID { + otherAddress = true + break + } + } + if !otherAddress { + delete(wd.modelSizes, modelID) + } } delete(wd.busyTime, address) delete(wd.inFlight, address) @@ -924,3 +927,20 @@ func (wd *WatchDog) untrack(address string) { delete(wd.addressModelMap, address) delete(wd.addressMap, address) } + +// Untrack removes request and eviction state after a backend is removed. +func (wd *WatchDog) Untrack(address string) { + wd.Lock() + defer wd.Unlock() + wd.untrack(address) +} + +func (wd *WatchDog) untrackProcess(p *process.Process) { + wd.Lock() + defer wd.Unlock() + for address, tracked := range wd.addressMap { + if tracked == p { + wd.untrack(address) + } + } +} diff --git a/pkg/model/watchdog_lifecycle_test.go b/pkg/model/watchdog_lifecycle_test.go new file mode 100644 index 000000000..50e4e84d3 --- /dev/null +++ b/pkg/model/watchdog_lifecycle_test.go @@ -0,0 +1,157 @@ +// SPDX-License-Identifier: MIT + +package model + +import ( + "context" + "os" + "path/filepath" + "time" + + grpc "github.com/mudler/LocalAI/pkg/grpc" + "github.com/mudler/LocalAI/pkg/system" + process "github.com/mudler/go-processmanager" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type watchdogLifecycleBackend struct{ grpc.Backend } + +func (*watchdogLifecycleBackend) IsBusy() bool { return false } +func (*watchdogLifecycleBackend) Free(context.Context) error { return nil } + +var _ = Describe("Watchdog backend lifecycle", func() { + var loader *ModelLoader + var wd *WatchDog + BeforeEach(func() { + loader = NewModelLoader(&system.SystemState{Model: system.Model{ModelsPath: GinkgoT().TempDir()}}) + wd = NewWatchDog(WithProcessManager(loader), WithIdleTimeout(time.Second), WithBusyTimeout(time.Second), WithLRULimit(1)) + loader.SetWatchDog(wd) + }) + + It("removes shutdown tracking and ignores late request completions", func() { + loader.store.Set("model", NewModelWithClient("model", "old", &watchdogLifecycleBackend{})) + wd.AddAddressModelMap("old", "model") + finish := wd.TrackRequest("old") + Expect(loader.ShutdownModelForce("model")).To(Succeed()) + finish() + state := wd.GetState() + Expect(state.AddressModelMap).To(BeEmpty()) + Expect(state.BusyTime).To(BeEmpty()) + Expect(state.IdleTime).To(BeEmpty()) + Expect(state.InFlight).To(BeEmpty()) + Expect(state.LastUsed).To(BeEmpty()) + }) + + DescribeTable("does not evict a replacement backend from a stale address", + func(evict func(*WatchDog)) { + replacement := NewModelWithClient("model", "new", &watchdogLifecycleBackend{}) + loader.store.Set("model", replacement) + wd.AddAddressModelMap("old", "model") + wd.AddAddressModelMap("new", "model") + wd.RegisterModelSize("model", 123) + wd.lastUsed["old"] = time.Now().Add(-time.Hour) + wd.lastUsed["new"] = time.Now() + wd.idleTime["old"] = time.Now().Add(-time.Hour) + evict(wd) + resident, ok := loader.store.Get("model") + Expect(ok).To(BeTrue()) + Expect(resident).To(BeIdenticalTo(replacement)) + Expect(wd.GetState().AddressModelMap).To(HaveKeyWithValue("new", "model")) + Expect(wd.modelSizes).To(HaveKeyWithValue("model", int64(123))) + }, + Entry("idle timeout", func(w *WatchDog) { w.checkIdle() }), + Entry("busy timeout", func(w *WatchDog) { w.busyTime["old"] = time.Now().Add(-time.Hour); w.checkBusy() }), + Entry("memory eviction", func(w *WatchDog) { w.evictLRUModel() }), + ) + + DescribeTable("does not stop a replacement after eviction selection", + func(force bool) { + wd.AddAddressModelMap("old", "model") + if force { + wd.TrackRequest("old") + } + targets, _ := wd.collectEvictionsLocked([]modelUsageInfo{{model: "model", address: "old"}}, 1, force) + replacement := NewModelWithClient("model", "new", &watchdogLifecycleBackend{}) + loader.store.Set("model", replacement) + wd.AddAddressModelMap("new", "model") + wd.RegisterModelSize("model", 123) + wd.shutdownEvicted(targets, "test") + resident, ok := loader.store.Get("model") + Expect(ok).To(BeTrue()) + Expect(resident).To(BeIdenticalTo(replacement)) + Expect(wd.modelSizes).To(HaveKeyWithValue("model", int64(123))) + }, + Entry("graceful", false), + Entry("forced", true), + ) + + DescribeTable("still stops the matching backend", + func(force bool) { + loader.store.Set("model", NewModelWithClient("model", "current", &watchdogLifecycleBackend{})) + wd.AddAddressModelMap("current", "model") + Expect(wd.shutdownTarget(evictionTarget{model: "model", address: "current", wasBusy: force})).To(Succeed()) + _, ok := loader.store.Get("model") + Expect(ok).To(BeFalse()) + Expect(wd.GetState().AddressModelMap).To(BeEmpty()) + }, + Entry("graceful", false), + Entry("forced", true), + ) + + It("keeps replacement tracking when an old process cleanup arrives late", func() { + oldProcess := process.New() + newProcess := process.New() + wd.Add("same-address", oldProcess) + wd.AddAddressModelMap("same-address", "model") + wd.Add("same-address", newProcess) + wd.RegisterModelSize("model", 123) + finish := wd.TrackRequest("same-address") + loader.cleanupProcessRuntime(oldProcess) + Expect(wd.GetState().AddressMap).To(HaveKeyWithValue("same-address", newProcess)) + Expect(wd.modelSizes).To(HaveKeyWithValue("model", int64(123))) + finish() + Expect(wd.GetState().IdleTime).To(HaveKey("same-address")) + }) + + It("removes local process tracking synchronously during shutdown", func() { + p := process.New() + m := NewModelWithClient("model", "old", &watchdogLifecycleBackend{}) + m.process = p + loader.store.Set("model", m) + wd.Add("old", p) + wd.AddAddressModelMap("old", "model") + wd.TrackRequest("old")() + Expect(loader.ShutdownModelForce("model")).To(Succeed()) + Expect(wd.GetState().AddressMap).To(BeEmpty()) + Expect(wd.GetState().IdleTime).To(BeEmpty()) + }) + + It("cleans the current watchdog after its configuration is replaced", func() { + p := process.New() + wd.Add("old", p) + wd.AddAddressModelMap("old", "model") + replacement := NewWatchDog(WithProcessManager(loader)) + replacement.RestoreState(wd.GetState()) + loader.SetWatchDog(replacement) + loader.cleanupProcessRuntime(p) + Expect(replacement.GetState().AddressModelMap).To(BeEmpty()) + }) + + It("untracks a backend that fails to start", func() { + _, err := loader.startProcess(filepath.Join(GinkgoT().TempDir(), "missing"), "model", "old", nil) + Expect(err).To(HaveOccurred()) + Expect(wd.GetState().AddressModelMap).To(BeEmpty()) + Expect(wd.GetState().AddressMap).To(BeEmpty()) + }) + + It("untracks a backend after an unexpected exit", func() { + backend := filepath.Join(GinkgoT().TempDir(), "backend") + Expect(os.WriteFile(backend, []byte("#!/bin/sh\nexit 42\n"), 0o700)).To(Succeed()) + p, err := loader.startProcess(backend, "model", "old", nil) + Expect(err).NotTo(HaveOccurred()) + DeferCleanup(func() { loader.cleanupProcessRuntime(p) }) + Eventually(p.Done()).Should(BeClosed()) + Eventually(func() map[string]string { return wd.GetState().AddressModelMap }).Should(BeEmpty()) + }) +}) diff --git a/swagger/docs.go b/swagger/docs.go index c7cd3c8f3..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" } } }, @@ -5851,6 +5930,13 @@ const docTemplate = `{ "label": { "type": "string" }, + "name": { + "description": "Name is the registered speaker this segment was matched to, and NameScore\nthe cosine similarity of the match. Both are omitted when the backend did\nnot identify the speaker. Speaker stays the normalized SPEAKER_NN label.", + "type": "string" + }, + "name_score": { + "type": "number" + }, "speaker": { "type": "string" }, @@ -5871,6 +5957,9 @@ const docTemplate = `{ "label": { "type": "string" }, + "name": { + "type": "string" + }, "segment_count": { "type": "integer" }, @@ -8889,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 3f374d87c..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" } } }, @@ -5848,6 +5927,13 @@ "label": { "type": "string" }, + "name": { + "description": "Name is the registered speaker this segment was matched to, and NameScore\nthe cosine similarity of the match. Both are omitted when the backend did\nnot identify the speaker. Speaker stays the normalized SPEAKER_NN label.", + "type": "string" + }, + "name_score": { + "type": "number" + }, "speaker": { "type": "string" }, @@ -5868,6 +5954,9 @@ "label": { "type": "string" }, + "name": { + "type": "string" + }, "segment_count": { "type": "integer" }, @@ -8886,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 7eeb216fd..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: @@ -1076,6 +1123,14 @@ definitions: type: integer label: type: string + name: + description: |- + Name is the registered speaker this segment was matched to, and NameScore + the cosine similarity of the match. Both are omitted when the backend did + not identify the speaker. Speaker stays the normalized SPEAKER_NN label. + type: string + name_score: + type: number speaker: type: string start: @@ -1089,6 +1144,8 @@ definitions: type: string label: type: string + name: + type: string segment_count: type: integer total_speech_duration: @@ -3242,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: @@ -5328,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: @@ -6166,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/tests/e2e-backends/backend_test.go b/tests/e2e-backends/backend_test.go index 9ebd044c8..4afda646a 100644 --- a/tests/e2e-backends/backend_test.go +++ b/tests/e2e-backends/backend_test.go @@ -62,6 +62,10 @@ import ( // model output into ChatDelta.tool_calls. // "image" exercises the GenerateImage RPC and asserts a // non-empty file is written to the requested dst path. +// "context_overflow" streams a prompt longer than the +// context and asserts the backend fails the stream with +// an error status WITHOUT first sending the error text as +// a content chunk (which clients would read as model output). // "long_prefill" sends a prompt long enough to span more // than one prefill batch and asserts the answer still // reflects the prompt. Catches GPU backends whose kernels @@ -98,6 +102,7 @@ const ( capLoad = "load" capPredict = "predict" capStream = "stream" + capCtxOverflow = "context_overflow" capEmbeddings = "embeddings" capTools = "tools" capTranscription = "transcription" @@ -540,6 +545,39 @@ var _ = Describe("Backend container", Ordered, func() { GinkgoWriter.Printf("Stream: %d chunks, combined=%q\n", chunks, combined) }) + It("fails a stream whose prompt exceeds the context without emitting content", func() { + if !caps[capCtxOverflow] { + Skip("context_overflow capability not enabled") + } + ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) + defer cancel() + // Far more tokens than any test context (default 512). + stream, err := client.PredictStream(ctx, &pb.PredictOptions{ + Prompt: strings.Repeat("overflow ", 8*int(envInt32("BACKEND_TEST_CTX_SIZE", 512))+64), + Tokens: 8, + }) + Expect(err).NotTo(HaveOccurred()) + + var content string + var streamErr error + for { + msg, err := stream.Recv() + if err == io.EOF { + break + } + if err != nil { + streamErr = err + break + } + content += string(msg.GetMessage()) + } + Expect(streamErr).To(HaveOccurred(), "an over-long prompt must fail the stream") + Expect(streamErr.Error()).To(ContainSubstring("exceeds the available context size")) + // Before the fix the backend wrote the error text as a Reply message + // first. LocalAI forwarded it as assistant content on a 200 stream. + Expect(content).To(BeEmpty(), "error text was streamed as content: %q", content) + }) + // Logprobs: backends that wire OpenAI-compatible logprobs return a // JSON-encoded payload in Reply.logprobs (see backend.proto). The exact // shape is backend-specific; we only assert that the field is populated diff --git a/website/content/blog/decision-models.md b/website/content/blog/decision-models.md new file mode 100644 index 000000000..cded25cb9 --- /dev/null +++ b/website/content/blog/decision-models.md @@ -0,0 +1,58 @@ +--- +title: "Run decision models in LocalAI" +date: 2026-10-02 +author: "Ettore Di Giacinto" +category: "Engineering" +tags: ["decisions", "classification", "models"] +summary: "Ask named questions about text and get structured answers for routing, moderation, and model selection." +extracss: ["blog.css"] +--- + +LocalAI now supports Jev-style decision models that answer named questions about text and return structured results. You can choose a category, assess a statement, or assign a level on a scale, with several questions in one request. + +This is useful when your application needs a decision rather than a written reply. You supply the text and describe the choices. Your application reads the answers by question name, without extracting a label from a chat response. + +For example, a support system could choose a queue for an incoming ticket and separately check whether the customer asks to cancel. A moderation workflow could flag a comment for review and assign an urgency level. A model selector could classify a request as translation, code, or general assistance before sending it to a suitable model. + +These are application patterns, not built-in workflows. You define the categories and decide what to do with each answer. A result can route work automatically or put it in a review queue, depending on the consequences of a mistake. + +## What is supported + +The Decisions API accepts text and a set of questions at `POST /v1/systemone`. Questions can ask for a choice from named options, an assessment of whether a statement holds, or a level on a scale. The response groups answers under the question names you supplied. + +The current gallery includes Laya, GLiNER2.5-Decide, Tev1 in 0.8B and 4B sizes, kev 0.8B, Nimble 9B, and CLM v0.1 8B. Use a current master build and the **alpha** vllm-cpp backend for these entries. + +With the backend and model installed on your machine, inference runs there. Installation downloads the required files. Local execution does not remove the need to control access or review your server's logging and tracing settings. + +Scores are not guaranteed probabilities or calibrated measures of correctness. Test your categories on representative examples and keep human review for consequential decisions, including moderation. + +## Get started + +In **Models → Explore**, find and install `laya-vllm-cpp`. Wait for installation to finish. The [Decisions API guide](/docs/features/decisions/) lists the other gallery IDs, memory requirements, request limits, and model-specific differences. + +With authentication enabled, set `LOCALAI_API_KEY` to a valid key for a user with the `decisions` feature enabled. Otherwise, omit the Authorization header. Send only text you are authorized to process. + +This example asks which kind of model should receive a request: + +```bash +curl --fail-with-body http://localhost:8080/v1/systemone \ + -H "Authorization: Bearer ${LOCALAI_API_KEY}" \ + -H 'Content-Type: application/json' \ + --data-binary '{ + "model": "laya-vllm-cpp", + "state": "Translate the attached release notes from English into Italian.", + "questions": { + "route": { + "type": "choice", + "instructions": "What kind of assistance does this request need?", + "criteria": { + "translation": "Translate text between languages", + "code": "Write or debug software", + "general": "Answer general questions" + } + } + } + }' +``` + +Read the selected option from `answers.route.choice`. This call classifies the text; your application must send the original request to the chosen model. Add another named question when you need a separate assessment of the same text. Keep the first trial small enough to check each answer yourself. 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.