diff --git a/bindings/go/params.go b/bindings/go/params.go index 95c5bfaf9..8f669e343 100644 --- a/bindings/go/params.go +++ b/bindings/go/params.go @@ -47,6 +47,44 @@ func (p *Params) SetPrintTimestamps(v bool) { p.print_timestamps = toBool(v) } +// Enable tinydiarize speaker turn detection +func (p *Params) SetDiarize(v bool) { + p.tdrz_enable = toBool(v) +} + +// Voice Activity Detection (VAD) +func (p *Params) SetVAD(v bool) { + p.vad = toBool(v) +} + +func (p *Params) SetVADModelPath(path string) { + p.vad_model_path = C.CString(path) +} + +func (p *Params) SetVADThreshold(t float32) { + p.vad_params.threshold = C.float(t) +} + +func (p *Params) SetVADMinSpeechMs(ms int) { + p.vad_params.min_speech_duration_ms = C.int(ms) +} + +func (p *Params) SetVADMinSilenceMs(ms int) { + p.vad_params.min_silence_duration_ms = C.int(ms) +} + +func (p *Params) SetVADMaxSpeechSec(s float32) { + p.vad_params.max_speech_duration_s = C.float(s) +} + +func (p *Params) SetVADSpeechPadMs(ms int) { + p.vad_params.speech_pad_ms = C.int(ms) +} + +func (p *Params) SetVADSamplesOverlap(sec float32) { + p.vad_params.samples_overlap = C.float(sec) +} + // Set language id func (p *Params) SetLanguage(lang int) error { if lang == -1 { diff --git a/bindings/go/pkg/whisper/consts.go b/bindings/go/pkg/whisper/consts.go index 0af45ee8c..ee002cff0 100644 --- a/bindings/go/pkg/whisper/consts.go +++ b/bindings/go/pkg/whisper/consts.go @@ -28,3 +28,10 @@ const SampleRate = whisper.SampleRate // SampleBits is the number of bytes per sample. const SampleBits = whisper.SampleBits + +type SamplingStrategy whisper.SamplingStrategy + +const ( + SAMPLING_GREEDY SamplingStrategy = SamplingStrategy(whisper.SAMPLING_GREEDY) + SAMPLING_BEAM_SEARCH SamplingStrategy = SamplingStrategy(whisper.SAMPLING_BEAM_SEARCH) +) diff --git a/bindings/go/pkg/whisper/context.go b/bindings/go/pkg/whisper/context.go index 01a510fe7..09b35be68 100644 --- a/bindings/go/pkg/whisper/context.go +++ b/bindings/go/pkg/whisper/context.go @@ -19,11 +19,11 @@ type context struct { Parameters } -func newContext(model Model, params whisper.Params) (Context, error) { +func newContext(model Model, params Parameters) (Context, error) { c := new(context) c.model = model - c.params = newParameters(¶ms) + c.params = params c.Parameters = c.params // allocate isolated state per context @@ -132,7 +132,7 @@ func (context *context) Process( context.params.SetSingleSegment(true) } - lowLevelParams := context.params.WhisperParams() + lowLevelParams := context.params.UnsafeParams() if lowLevelParams == nil { return fmt.Errorf("lowLevelParams is nil: %w", ErrInternalAppError) } @@ -249,11 +249,12 @@ func (context *context) IsLANG(t Token, lang string) bool { // State-backed helper functions func toSegmentFromState(ctx *whisper.Context, st *whisper.State, n int) Segment { return Segment{ - Num: n, - Text: strings.TrimSpace(ctx.Whisper_full_get_segment_text_from_state(st, n)), - Start: time.Duration(ctx.Whisper_full_get_segment_t0_from_state(st, n)) * time.Millisecond * 10, - End: time.Duration(ctx.Whisper_full_get_segment_t1_from_state(st, n)) * time.Millisecond * 10, - Tokens: toTokensFromState(ctx, st, n), + Num: n, + Text: strings.TrimSpace(ctx.Whisper_full_get_segment_text_from_state(st, n)), + Start: time.Duration(ctx.Whisper_full_get_segment_t0_from_state(st, n)) * time.Millisecond * 10, + End: time.Duration(ctx.Whisper_full_get_segment_t1_from_state(st, n)) * time.Millisecond * 10, + Tokens: toTokensFromState(ctx, st, n), + SpeakerTurnNext: ctx.Whisper_full_get_segment_speaker_turn_next_from_state(st, n), } } diff --git a/bindings/go/pkg/whisper/context_test.go b/bindings/go/pkg/whisper/context_test.go index 305446624..f94c9139d 100644 --- a/bindings/go/pkg/whisper/context_test.go +++ b/bindings/go/pkg/whisper/context_test.go @@ -17,7 +17,7 @@ func TestSetLanguage(t *testing.T) { model, err := whisper.New(ModelPath) assert.NoError(err) assert.NotNil(model) - defer model.Close() + defer func() { _ = model.Close() }() context, err := model.NewContext() assert.NoError(err) @@ -35,7 +35,7 @@ func TestContextModelIsMultilingual(t *testing.T) { model, err := whisper.New(ModelPath) assert.NoError(err) assert.NotNil(model) - defer model.Close() + defer func() { _ = model.Close() }() context, err := model.NewContext() assert.NoError(err) @@ -54,7 +54,7 @@ func TestLanguage(t *testing.T) { model, err := whisper.New(ModelPath) assert.NoError(err) assert.NotNil(model) - defer model.Close() + defer func() { _ = model.Close() }() context, err := model.NewContext() assert.NoError(err) @@ -72,7 +72,7 @@ func TestProcess(t *testing.T) { fh, err := os.Open(SamplePath) assert.NoError(err) - defer fh.Close() + defer func() { _ = fh.Close() }() // Decode the WAV file - load the full buffer dec := wav.NewDecoder(fh) @@ -85,7 +85,7 @@ func TestProcess(t *testing.T) { model, err := whisper.New(ModelPath) assert.NoError(err) assert.NotNil(model) - defer model.Close() + defer func() { _ = model.Close() }() context, err := model.NewContext() assert.NoError(err) @@ -99,7 +99,7 @@ func TestDetectedLanguage(t *testing.T) { fh, err := os.Open(SamplePath) assert.NoError(err) - defer fh.Close() + defer func() { _ = fh.Close() }() // Decode the WAV file - load the full buffer dec := wav.NewDecoder(fh) @@ -112,7 +112,7 @@ func TestDetectedLanguage(t *testing.T) { model, err := whisper.New(ModelPath) assert.NoError(err) assert.NotNil(model) - defer model.Close() + defer func() { _ = model.Close() }() context, err := model.NewContext() assert.NoError(err) @@ -139,7 +139,7 @@ func TestContext_ConcurrentProcessing(t *testing.T) { fh, err := os.Open(SamplePath) assert.NoError(err) - defer fh.Close() + defer func() { _ = fh.Close() }() dec := wav.NewDecoder(fh) buf, err := dec.FullPCMBuffer() @@ -150,12 +150,12 @@ func TestContext_ConcurrentProcessing(t *testing.T) { model, err := whisper.New(ModelPath) assert.NoError(err) assert.NotNil(model) - defer model.Close() + defer func() { _ = model.Close() }() ctx, err := model.NewContext() assert.NoError(err) assert.NotNil(ctx) - defer ctx.Close() + defer func() { _ = ctx.Close() }() err = ctx.Process(data, nil, nil, nil) assert.NoError(err) @@ -179,7 +179,7 @@ func TestContext_Parallel_DifferentInputs(t *testing.T) { fh, err := os.Open(SamplePath) assert.NoError(err) - defer fh.Close() + defer func() { _ = fh.Close() }() dec := wav.NewDecoder(fh) buf, err := dec.FullPCMBuffer() @@ -195,14 +195,14 @@ func TestContext_Parallel_DifferentInputs(t *testing.T) { model, err := whisper.New(ModelPath) assert.NoError(err) assert.NotNil(model) - defer model.Close() + defer func() { _ = model.Close() }() ctx1, err := model.NewContext() assert.NoError(err) - defer ctx1.Close() + defer func() { _ = ctx1.Close() }() ctx2, err := model.NewContext() assert.NoError(err) - defer ctx2.Close() + defer func() { _ = ctx2.Close() }() // Run in parallel - each context has isolated whisper_state var wg sync.WaitGroup @@ -258,7 +258,7 @@ func TestContext_Close(t *testing.T) { model, err := whisper.New(ModelPath) assert.NoError(err) assert.NotNil(model) - defer model.Close() + defer func() { _ = model.Close() }() ctx, err := model.NewContext() assert.NoError(err) @@ -294,3 +294,82 @@ func Test_Close_Context_of_Closed_Model(t *testing.T) { require.NoError(t, model.Close()) require.NoError(t, ctx.Close()) } + +func TestContext_VAD_And_Diarization_Params_DoNotPanic(t *testing.T) { + assert := assert.New(t) + + if _, err := os.Stat(ModelPath); os.IsNotExist(err) { + t.Skip("Skipping test, model not found:", ModelPath) + } + if _, err := os.Stat(SamplePath); os.IsNotExist(err) { + t.Skip("Skipping test, sample not found:", SamplePath) + } + + fh, err := os.Open(SamplePath) + assert.NoError(err) + defer func() { _ = fh.Close() }() + + dec := wav.NewDecoder(fh) + buf, err := dec.FullPCMBuffer() + assert.NoError(err) + assert.Equal(uint16(1), dec.NumChans) + data := buf.AsFloat32Buffer().Data + + model, err := whisper.New(ModelPath) + assert.NoError(err) + defer func() { _ = model.Close() }() + + ctx, err := model.NewContext() + assert.NoError(err) + defer func() { _ = ctx.Close() }() + + p := ctx.Params() + p.SetDiarize(true) + p.SetVAD(true) + p.SetVADThreshold(0.5) + p.SetVADMinSpeechMs(200) + p.SetVADMinSilenceMs(100) + p.SetVADMaxSpeechSec(10) + p.SetVADSpeechPadMs(30) + p.SetVADSamplesOverlap(0.02) + + err = ctx.Process(data, nil, nil, nil) + assert.NoError(err) +} + +func TestContext_SpeakerTurnNext_Field_Present(t *testing.T) { + assert := assert.New(t) + + if _, err := os.Stat(ModelPath); os.IsNotExist(err) { + t.Skip("Skipping test, model not found:", ModelPath) + } + if _, err := os.Stat(SamplePath); os.IsNotExist(err) { + t.Skip("Skipping test, sample not found:", SamplePath) + } + + fh, err := os.Open(SamplePath) + assert.NoError(err) + defer func() { _ = fh.Close() }() + + dec := wav.NewDecoder(fh) + buf, err := dec.FullPCMBuffer() + assert.NoError(err) + assert.Equal(uint16(1), dec.NumChans) + data := buf.AsFloat32Buffer().Data + + model, err := whisper.New(ModelPath) + assert.NoError(err) + defer func() { _ = model.Close() }() + + ctx, err := model.NewContext() + assert.NoError(err) + defer func() { _ = ctx.Close() }() + + err = ctx.Process(data, nil, nil, nil) + assert.NoError(err) + + seg, err := ctx.NextSegment() + assert.NoError(err) + t.Logf("SpeakerTurnNext: %v", seg.SpeakerTurnNext) + _ = seg.SpeakerTurnNext // ensure field exists and is readable +} diff --git a/bindings/go/pkg/whisper/interface.go b/bindings/go/pkg/whisper/interface.go index 30a41db3e..4aa679c66 100644 --- a/bindings/go/pkg/whisper/interface.go +++ b/bindings/go/pkg/whisper/interface.go @@ -48,6 +48,8 @@ type TokenIdentifier interface { IsText(Token) (bool, error) } +type ParamsConfigure func(Parameters) + // Model is the interface to a whisper model. Create a new model with the // function whisper.New(string) type Model interface { @@ -56,6 +58,11 @@ type Model interface { // Return a new speech-to-text context. NewContext() (Context, error) + NewParams( + sampling SamplingStrategy, + configure ParamsConfigure, + ) (Parameters, error) + // Return true if the model is multilingual. IsMultilingual() bool @@ -94,6 +101,25 @@ type Parameters interface { SetEntropyThold(t float32) SetInitialPrompt(prompt string) + SetNoContext(bool) + SetPrintSpecial(bool) + SetPrintProgress(bool) + SetPrintRealtime(bool) + SetPrintTimestamps(bool) + + // Diarization (tinydiarize) + SetDiarize(bool) + + // Voice Activity Detection (VAD) + SetVAD(bool) + SetVADModelPath(string) + SetVADThreshold(float32) + SetVADMinSpeechMs(int) + SetVADMinSilenceMs(int) + SetVADMaxSpeechSec(float32) + SetVADSpeechPadMs(int) + SetVADSamplesOverlap(float32) + // Set the temperature SetTemperature(t float32) @@ -108,7 +134,8 @@ type Parameters interface { // Getter methods Language() string Threads() int - WhisperParams() *whisper.Params + + UnsafeParams() *whisper.Params } // Context is the speech recognition context. @@ -231,6 +258,9 @@ type Segment struct { // The tokens of the segment. Tokens []Token + + // True if the next segment is predicted as a speaker turn (tinydiarize) + SpeakerTurnNext bool } // Token is a text or special token diff --git a/bindings/go/pkg/whisper/model.go b/bindings/go/pkg/whisper/model.go index 752c7cf02..203ec7b20 100644 --- a/bindings/go/pkg/whisper/model.go +++ b/bindings/go/pkg/whisper/model.go @@ -3,7 +3,6 @@ package whisper import ( "fmt" "os" - "runtime" // Bindings whisper "github.com/ggerganov/whisper.cpp/bindings/go" @@ -88,24 +87,72 @@ func (model *model) Languages() []string { // NewContext creates a new speech-to-text context. // Each context is backed by an isolated whisper_state for safe concurrent processing. func (model *model) NewContext() (Context, error) { - ctx, err := model.ctx.UnsafeContext() + // Create new context with default params + params, err := model.newParams(SAMPLING_GREEDY, nil) if err != nil { - return nil, ErrModelClosed + return nil, err } - // Create new context with default params - params := ctx.Whisper_full_default_params(whisper.SAMPLING_GREEDY) + // Return new context (now state-backed) + return newContext( + model, + params, + ) +} +func (model *model) NewParams( + sampling SamplingStrategy, + configure ParamsConfigure, +) (Parameters, error) { + return model.newParams(sampling, nil) +} + +// NewContextWithParams creates a new speech-to-text context and allows +// callers to customize the decoding parameters before the state is used. +// The resulting Context is backed by an isolated whisper_state for safe +// concurrent processing. +func (model *model) NewContextWithParams( + sampling SamplingStrategy, + configure ParamsConfigure, +) (Context, error) { + params, err := model.newParams(sampling, configure) + if err != nil { + return nil, err + } + + return newContext( + model, + params, + ) +} + +func defaultParamsConfigure(params Parameters) { params.SetTranslate(false) params.SetPrintSpecial(false) params.SetPrintProgress(false) params.SetPrintRealtime(false) params.SetPrintTimestamps(false) - params.SetThreads(runtime.NumCPU()) - params.SetNoContext(true) +} - // Return new context (now state-backed) - return newContext(model, params) +func (m *model) newParams( + sampling SamplingStrategy, + configure ParamsConfigure, +) (Parameters, error) { + ctx, err := m.ctx.UnsafeContext() + if err != nil { + return nil, ErrModelClosed + } + + p := ctx.Whisper_full_default_params(whisper.SamplingStrategy(sampling)) + safeParams := newParameters(&p) + + defaultParamsConfigure(safeParams) + + if configure != nil { + configure(safeParams) + } + + return safeParams, nil } // PrintTimings prints the model performance timings to stdout. diff --git a/bindings/go/pkg/whisper/params_wrap.go b/bindings/go/pkg/whisper/params_wrap.go index 55a414537..ac3a1ccc1 100644 --- a/bindings/go/pkg/whisper/params_wrap.go +++ b/bindings/go/pkg/whisper/params_wrap.go @@ -13,7 +13,11 @@ type parameters struct { p *whisper.Params } -func newParameters(p *whisper.Params) Parameters { return ¶meters{p: p} } +func newParameters(whisperParams *whisper.Params) Parameters { + return ¶meters{ + p: whisperParams, + } +} func (w *parameters) SetTranslate(v bool) { w.p.SetTranslate(v) } func (w *parameters) SetSplitOnWord(v bool) { w.p.SetSplitOnWord(v) } @@ -32,6 +36,24 @@ func (w *parameters) SetEntropyThold(t float32) { w.p.SetEntropyThold(t) func (w *parameters) SetInitialPrompt(prompt string) { w.p.SetInitialPrompt(prompt) } func (w *parameters) SetTemperature(t float32) { w.p.SetTemperature(t) } func (w *parameters) SetTemperatureFallback(t float32) { w.p.SetTemperatureFallback(t) } +func (w *parameters) SetNoContext(v bool) { w.p.SetNoContext(v) } +func (w *parameters) SetPrintSpecial(v bool) { w.p.SetPrintSpecial(v) } +func (w *parameters) SetPrintProgress(v bool) { w.p.SetPrintProgress(v) } +func (w *parameters) SetPrintRealtime(v bool) { w.p.SetPrintRealtime(v) } +func (w *parameters) SetPrintTimestamps(v bool) { w.p.SetPrintTimestamps(v) } + +// Diarization (tinydiarize) +func (w *parameters) SetDiarize(v bool) { w.p.SetDiarize(v) } + +// Voice Activity Detection (VAD) +func (w *parameters) SetVAD(v bool) { w.p.SetVAD(v) } +func (w *parameters) SetVADModelPath(p string) { w.p.SetVADModelPath(p) } +func (w *parameters) SetVADThreshold(t float32) { w.p.SetVADThreshold(t) } +func (w *parameters) SetVADMinSpeechMs(ms int) { w.p.SetVADMinSpeechMs(ms) } +func (w *parameters) SetVADMinSilenceMs(ms int) { w.p.SetVADMinSilenceMs(ms) } +func (w *parameters) SetVADMaxSpeechSec(s float32) { w.p.SetVADMaxSpeechSec(s) } +func (w *parameters) SetVADSpeechPadMs(ms int) { w.p.SetVADSpeechPadMs(ms) } +func (w *parameters) SetVADSamplesOverlap(sec float32) { w.p.SetVADSamplesOverlap(sec) } func (w *parameters) SetLanguage(lang string) error { if lang == "auto" { @@ -62,7 +84,7 @@ func (w *parameters) Threads() int { return w.p.Threads() } -func (w *parameters) WhisperParams() *whisper.Params { +func (w *parameters) UnsafeParams() *whisper.Params { return w.p } diff --git a/bindings/go/whisper.go b/bindings/go/whisper.go index b6ef48a85..83089a26f 100644 --- a/bindings/go/whisper.go +++ b/bindings/go/whisper.go @@ -533,6 +533,11 @@ func (ctx *Context) Whisper_get_logits_from_state(state *State) []float32 { return (*[1 << 30]float32)(unsafe.Pointer(C.whisper_get_logits_from_state((*C.struct_whisper_state)(state))))[:ctx.Whisper_n_vocab()] } +// Get whether the next segment is predicted as a speaker turn (tinydiarize) +func (ctx *Context) Whisper_full_get_segment_speaker_turn_next_from_state(state *State, segment int) bool { + return bool(C.whisper_full_get_segment_speaker_turn_next_from_state((*C.struct_whisper_state)(state), C.int(segment))) +} + /////////////////////////////////////////////////////////////////////////////// // CALLBACKS