feat(go bindings): add VAD and Diarization parameters

This commit is contained in:
ciricc 2025-09-13 22:28:14 +03:00
parent ebbcf3f17f
commit 97e6ce2bc4
8 changed files with 264 additions and 35 deletions

View File

@ -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 {

View File

@ -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)
)

View File

@ -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(&params)
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),
}
}

View File

@ -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
}

View File

@ -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

View File

@ -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.

View File

@ -13,7 +13,11 @@ type parameters struct {
p *whisper.Params
}
func newParameters(p *whisper.Params) Parameters { return &parameters{p: p} }
func newParameters(whisperParams *whisper.Params) Parameters {
return &parameters{
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
}

View File

@ -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