diff --git a/backend/go/silero-vad/vad.go b/backend/go/silero-vad/vad.go index f3e9f7be8..8ca239764 100644 --- a/backend/go/silero-vad/vad.go +++ b/backend/go/silero-vad/vad.go @@ -4,32 +4,70 @@ package main // It is meant to be used by the main executable that is the server for the specific backend type (falcon, gpt3, etc) import ( "fmt" + "math" + "strconv" + "strings" "github.com/mudler/LocalAI/pkg/grpc/base" pb "github.com/mudler/LocalAI/pkg/grpc/proto" "github.com/streamer45/silero-vad-go/speech" ) +const ( + defaultThreshold = 0.5 + defaultMinSilenceDurationMs = 100 + defaultSpeechPadMs = 30 +) + type VAD struct { base.SingleThread detector *speech.Detector } func (vad *VAD) Load(opts *pb.ModelOptions) error { - v, err := speech.NewDetector(speech.DetectorConfig{ - ModelPath: opts.ModelFile, - SampleRate: 16000, - //WindowSize: 1024, - Threshold: 0.5, - MinSilenceDurationMs: 100, - SpeechPadMs: 30, - }) + cfg := detectorConfigFromOptions(opts) + + v, err := speech.NewDetector(cfg) if err != nil { return fmt.Errorf("create silero detector: %w", err) } vad.detector = v - return err + return nil +} + +func detectorConfigFromOptions(opts *pb.ModelOptions) speech.DetectorConfig { + cfg := speech.DetectorConfig{ + ModelPath: opts.ModelFile, + SampleRate: 16000, + Threshold: defaultThreshold, + MinSilenceDurationMs: defaultMinSilenceDurationMs, + SpeechPadMs: defaultSpeechPadMs, + } + + for _, opt := range opts.Options { + key, value, ok := strings.Cut(opt, ":") + if !ok || value == "" { + continue + } + + switch strings.ToLower(strings.TrimSpace(key)) { + case "threshold": + if v, err := strconv.ParseFloat(strings.TrimSpace(value), 32); err == nil && !math.IsNaN(v) { + cfg.Threshold = float32(v) + } + case "min_silence_duration_ms": + if v, err := strconv.Atoi(strings.TrimSpace(value)); err == nil && v >= 0 { + cfg.MinSilenceDurationMs = v + } + case "speech_pad_ms": + if v, err := strconv.Atoi(strings.TrimSpace(value)); err == nil && v >= 0 { + cfg.SpeechPadMs = v + } + } + } + + return cfg } func (vad *VAD) VAD(req *pb.VADRequest) (pb.VADResponse, error) { diff --git a/backend/go/silero-vad/vad_test.go b/backend/go/silero-vad/vad_test.go new file mode 100644 index 000000000..d865a0800 --- /dev/null +++ b/backend/go/silero-vad/vad_test.go @@ -0,0 +1,67 @@ +package main + +import ( + "testing" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestVADOptions(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Silero VAD Options") +} + +var _ = Describe("Detector configuration", func() { + It("preserves the model path and defaults", func() { + cfg := detectorConfigFromOptions(&pb.ModelOptions{ModelFile: "silero_vad.onnx"}) + Expect(cfg.ModelPath).To(Equal("silero_vad.onnx")) + Expect(cfg.SampleRate).To(Equal(16000)) + Expect(cfg.Threshold).To(Equal(float32(0.5))) + Expect(cfg.MinSilenceDurationMs).To(Equal(100)) + Expect(cfg.SpeechPadMs).To(Equal(30)) + }) + + It("applies model options", func() { + cfg := detectorConfigFromOptions(&pb.ModelOptions{ + ModelFile: "silero_vad.onnx", + Options: []string{ + "threshold:0.55", + "min_silence_duration_ms:50", + "speech_pad_ms:450", + "ignored", + "bad_threshold:abc", + }, + }) + Expect(cfg.Threshold).To(Equal(float32(0.55))) + Expect(cfg.MinSilenceDurationMs).To(Equal(50)) + Expect(cfg.SpeechPadMs).To(Equal(450)) + }) + + It("ignores NaN and malformed values without replacing valid options", func() { + cfg := detectorConfigFromOptions(&pb.ModelOptions{ + ModelFile: "silero_vad.onnx", + Options: []string{ + "threshold:0.55", + "threshold:NaN", + "threshold:abc", + "min_silence_duration_ms:-1", + "speech_pad_ms:abc", + }, + }) + Expect(cfg.Threshold).To(Equal(float32(0.55))) + Expect(cfg.MinSilenceDurationMs).To(Equal(100)) + Expect(cfg.SpeechPadMs).To(Equal(30)) + Expect(cfg.IsValid()).To(Succeed()) + }) + + It("uses the default threshold when the only override is NaN", func() { + cfg := detectorConfigFromOptions(&pb.ModelOptions{ + ModelFile: "silero_vad.onnx", + Options: []string{"threshold:NaN"}, + }) + Expect(cfg.Threshold).To(Equal(float32(0.5))) + Expect(cfg.IsValid()).To(Succeed()) + }) +}) diff --git a/docs/content/features/voice-activity-detection.md b/docs/content/features/voice-activity-detection.md index 2fc3f6e5b..63f516f4f 100644 --- a/docs/content/features/voice-activity-detection.md +++ b/docs/content/features/voice-activity-detection.md @@ -93,9 +93,33 @@ name: silero-vad backend: silero-vad ``` +Detection parameters can be overridden via model `options` (`key:value` entries): + +```yaml +name: silero-vad +backend: silero-vad +options: + - threshold:0.55 + - min_silence_duration_ms:50 + - speech_pad_ms:450 +``` + +Supported options: + +| Option | Type | Default | Description | +|--------|------|---------|-------------| +| `threshold` | float | `0.5` | Speech probability threshold | +| `min_silence_duration_ms` | int | `100` | Minimum silence before ending a speech segment | +| `speech_pad_ms` | int | `30` | Padding added around each speech segment | + +Thresholds must be greater than 0 and less than 1. Durations must be nonnegative integers. +Malformed values, negative durations, and NaN thresholds are ignored; the default or last valid value remains in use. + +Reload the model (or restart LocalAI) after changing these options. + ## Detection Parameters -The Silero VAD backend uses the following internal defaults: +The Silero VAD backend uses the following internal defaults (overridable via `options` above): - **Sample rate:** 16kHz - **Threshold:** 0.5 diff --git a/examples/silero-vad-custom-options.yaml b/examples/silero-vad-custom-options.yaml new file mode 100644 index 000000000..fb9f8ccaf --- /dev/null +++ b/examples/silero-vad-custom-options.yaml @@ -0,0 +1,6 @@ +name: silero-vad +backend: silero-vad +options: + - threshold:0.55 + - min_silence_duration_ms:50 + - speech_pad_ms:450