mirror of
https://github.com/mudler/LocalAI.git
synced 2026-10-05 12:34:43 -04:00
chore: merge master into distributed test branch
Preserve multipart language hints alongside JSON diarization and speaker profiles when resolving the endpoint conflict. Assisted-by: Codex:gpt-6
This commit is contained in:
commit
ba75324ae4
104 files changed
+5603
-291
No files matched your search
@@ -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
|
||||
|
||||
@@ -628,6 +628,7 @@ message TranscriptLiveConfig {
|
||||
string language = 1; // "" => model default
|
||||
int32 sample_rate = 2; // 0 => 16000; backends may reject others
|
||||
map<string, string> 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<string, uint64> 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 {
|
||||
|
||||
@@ -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))))
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
IK_LLAMA_VERSION?=0821d62a8b356bd1db3c6765551a30bfcc44a6de
|
||||
IK_LLAMA_VERSION?=d9e286846d6f8232db48ec5c111a4ea3aea675ef
|
||||
LLAMA_REPO?=https://github.com/ikawrakow/ik_llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
LLAMA_VERSION?=4da6337767f973e2b4d0797e5b323d77d8565e4a
|
||||
LLAMA_VERSION?=a868c3e3c56657f7e8a6231190dbbe90e7dd86c0
|
||||
LLAMA_REPO?=https://github.com/ggerganov/llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -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<int> 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"));
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
#pragma once
|
||||
|
||||
#include <optional>
|
||||
#include <string>
|
||||
|
||||
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<int>& 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
|
||||
@@ -0,0 +1,29 @@
|
||||
#include "parallel_params.h"
|
||||
|
||||
#include <cstdio>
|
||||
|
||||
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;
|
||||
}
|
||||
@@ -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)) {
|
||||
|
||||
@@ -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/
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}))
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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 <EOU> and reports it", func() {
|
||||
text, eou := stripEouMarker("it is certainly very like the old portrait<EOU>")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"))
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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}))
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
|
||||
File diff suppressed because one or more lines are too long.
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -2,3 +2,4 @@
|
||||
torch
|
||||
torchaudio
|
||||
funasr
|
||||
transformers>=4.32.0,<5
|
||||
@@ -3,3 +3,4 @@ protobuf
|
||||
certifi
|
||||
packaging==24.1
|
||||
setuptools
|
||||
transformers>=4.32.0,<5
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
+120
-28
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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))
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -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(),
|
||||
})
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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)),
|
||||
}
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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())
|
||||
|
||||
@@ -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 <file> 1 <start> <duration> <NA> <NA> <speaker> <NA> <NA>
|
||||
// Field separators are spaces; one row per segment.
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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 != ""
|
||||
|
||||
@@ -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"))
|
||||
})
|
||||
})
|
||||
@@ -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)}
|
||||
}
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
@@ -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."
|
||||
}
|
||||
}
|
||||
@@ -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' }) {
|
||||
<div
|
||||
role="dialog"
|
||||
aria-modal="true"
|
||||
aria-label={ariaLabel}
|
||||
className="modal-backdrop"
|
||||
onClick={onClose}
|
||||
>
|
||||
|
||||
@@ -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 (
|
||||
<div className="page-pad">
|
||||
<PageHeader title={text('title')} supporting={text('subtitle')} />
|
||||
<form onSubmit={submit} className="card stack">
|
||||
<div className="form-group" role="group" aria-label={text('model')}>
|
||||
<span className="form-label">{text('model')}</span>
|
||||
<ModelSelector value={model} capability={CAP_DIARIZATION} onChange={value => { if (value !== model) { invalidate(); setModel(value) } }} />
|
||||
</div>
|
||||
<div className="form-group">
|
||||
<label className="form-label" htmlFor="diarization-file">{text('recording')}</label>
|
||||
<input id="diarization-file" className="input" type="file" accept="audio/*,video/*" onChange={e => { invalidate(); setFile(e.target.files?.[0] || null) }} />
|
||||
</div>
|
||||
{canRemember && <div className="form-group">
|
||||
<label><input type="checkbox" checked={optIn} onChange={e => { invalidate(); setOptIn(e.target.checked) }} /> {text('optIn')}</label>
|
||||
<p className="form-help">{text('warning')} <Link to="/app/voice">{text('manage')}</Link></p>
|
||||
</div>}
|
||||
<button className="btn btn-primary" disabled={busy || !file || !model}>{text(busy ? 'running' : 'run')}</button>
|
||||
</form>
|
||||
{error && <p role="alert">{error}</p>}
|
||||
{url && <audio ref={audio} src={url} preload="metadata" onTimeUpdate={() => { if (end.current !== null && audio.current.currentTime >= end.current) stop() }} />}
|
||||
{result && <>
|
||||
<div className="hstack"><h2>{text('speakers')}</h2>
|
||||
{canRemember && result.speaker_profiles && <button type="button" className="btn btn-secondary" onClick={stop}>{text('stop')}</button>}
|
||||
</div>
|
||||
<ul className="lanes">
|
||||
{summaries.map(summary => {
|
||||
const profile = result.speaker_profiles?.speakers.find(p => String(p.speaker) === String(summary.label))
|
||||
const knownName = summary.name || result.segments?.find(s => String(s.label) === String(summary.label) && s.name)?.name
|
||||
const usable = profile?.embedding?.length > 0 && !profile.unavailable_reason
|
||||
return <li className="card stack" key={summary.label} data-testid={`speaker-${summary.label}`}>
|
||||
<h3>{knownName || summary.id}</h3>
|
||||
{profile && canRemember && <>
|
||||
<p>{text('duration', { seconds: profile.clean_duration })}</p>
|
||||
<div className="hstack">{profile.intervals.map((interval, i) => <button key={i} type="button" className="btn btn-secondary" onClick={() => preview(interval)}>{text('preview', { number: i + 1 })}</button>)}</div>
|
||||
{!usable && <p>{text('insufficient')} {profile.unavailable_reason && <small>({profile.unavailable_reason})</small>}</p>}
|
||||
{!knownName && <button type="button" className="btn btn-primary" disabled={!usable} onClick={() => { stop(); setSelected(profile.speaker); setName(''); setSaveError('') }}>{text('nameAndRemember')}</button>}
|
||||
</>}
|
||||
</li>
|
||||
})}
|
||||
</ul>
|
||||
<h2>{text('segments')}</h2>
|
||||
<ol data-testid="segments" className="lanes">
|
||||
{result.segments?.map((segment, i) => <li key={segment.id ?? i} className="card"><strong>{segment.name || segment.speaker}</strong> <span>{segment.start}–{segment.end}s</span><p>{segment.text}</p></li>)}
|
||||
</ol>
|
||||
</>}
|
||||
{selected !== null && canRemember && <Modal ariaLabel={text('nameAndRemember')} onClose={() => { if (!saving) setSelected(null) }}>
|
||||
<form className="stack" onSubmit={save} aria-label={text('nameAndRemember')}>
|
||||
<h2>{text('nameAndRemember')}</h2>
|
||||
<p>{text('warning')}</p>
|
||||
<label className="form-label" htmlFor="diarization-name">{text('name')}</label>
|
||||
<input className="input" id="diarization-name" required value={name} disabled={saving} onChange={e => setName(e.target.value)} />
|
||||
{saveError && <p role="alert">{saveError}</p>}
|
||||
<button type="submit" className="btn btn-primary" disabled={saving || !name.trim()}>{text(saving ? 'saving' : 'remember')}</button>
|
||||
<button type="button" className="btn btn-secondary" disabled={saving} onClick={() => setSelected(null)}>{text('cancel')}</button>
|
||||
</form>
|
||||
</Modal>}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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: <TTS /> },
|
||||
{ path: 'sound', element: <Sound /> },
|
||||
{ path: 'sound/:model', element: <Sound /> },
|
||||
{ path: 'diarization', element: <Feature feature="audio_diarization"><Diarization /></Feature> },
|
||||
{ path: 'diarization/:model', element: <Feature feature="audio_diarization"><Diarization /></Feature> },
|
||||
{ path: 'transform', element: <Feature feature="audio_transform"><AudioTransform /></Feature> },
|
||||
{ path: 'transform/:model', element: <Feature feature="audio_transform"><AudioTransform /></Feature> },
|
||||
{ path: 'studio', element: <Studio /> },
|
||||
|
||||
Vendored
+15
@@ -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 }))
|
||||
},
|
||||
}
|
||||
@@ -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()])
|
||||
}
|
||||
@@ -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)),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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:"-"`
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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}))
|
||||
})
|
||||
})
|
||||
@@ -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:<file> 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
|
||||
}
|
||||
@@ -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))
|
||||
})
|
||||
})
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"))
|
||||
})
|
||||
})
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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": "<raw base64 audio bytes>",
|
||||
"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 <key>"` 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.
|
||||
@@ -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:<path>` | an ASR model | a `speaker` on transcript segments (and words), and speaker segments during realtime live transcription |
|
||||
| `sound_model:<path>` | an ASR model | sound events during realtime live transcription |
|
||||
| `diarization_latency:<model\|low\|very_low\|ultra_low>` | a model with a diarization companion | latency mode for the live speaker stream; default `low` |
|
||||
| `speaker_model:<path>` | a model with a diarization model | names registered speakers (see [Voice Recognition]({{% relref "voice-recognition" %}}#naming-speakers-in-diarization-and-live-transcription)) |
|
||||
| `speaker_threshold:<float>` | a model with `speaker_model` | distance (1 minus cosine similarity) under which a speaker is named, in (0, 2); default `0.5` |
|
||||
| `speaker_margin:<float>` | 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 "<path>" is a <kind> model, expected a speaker model` (the file is not a speaker encoder GGUF), or `parakeet-cpp: speaker_threshold "<value>" must be a distance in (0, 2) (1 minus cosine similarity)` / `parakeet-cpp: speaker_margin "<value>" 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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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:<path>` | none | speaker encoder GGUF; needs a diarization model (the primary one, or `diarization_model:`) |
|
||||
| `speaker_threshold:<float>` | `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:<float>` | `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.
|
||||
@@ -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
|
||||
|
||||
@@ -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 |
|
||||
|
||||
@@ -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
|
||||
|
||||
+358
-2
@@ -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:
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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.
|
||||
|
||||
+55
-35
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
+100
-4
@@ -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."
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
+100
-4
@@ -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."
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
Loaded 100 of 104 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user