Files
Captioneer/src/models/sherpa.go
T

145 lines
3.7 KiB
Go
Raw Normal View History

// Package models owns Sherpa-Onnx model setup and inference resources.
package models
import (
"fmt"
"os"
"path/filepath"
"strings"
sherpa "github.com/k2-fsa/sherpa-onnx-go-linux"
"tea.chunkbyte.com/kato/captioneer/src/audio"
"tea.chunkbyte.com/kato/captioneer/src/config"
)
type modelPaths struct {
encoder string
decoder string
joiner string
tokens string
vad string
}
type Resources struct {
recognizer *sherpa.OfflineRecognizer
vad *sherpa.VoiceActivityDetector
}
func New(settings config.Settings) (*Resources, error) {
paths, err := validatePaths(settings.ModelsDir, settings.Mode)
if err != nil {
return nil, err
}
recognizer := sherpa.NewOfflineRecognizer(&sherpa.OfflineRecognizerConfig{
FeatConfig: sherpa.FeatureConfig{SampleRate: audio.SampleRate, FeatureDim: 80},
ModelConfig: sherpa.OfflineModelConfig{
Transducer: sherpa.OfflineTransducerModelConfig{
Encoder: paths.encoder,
Decoder: paths.decoder,
Joiner: paths.joiner,
},
Tokens: paths.tokens,
NumThreads: settings.Threads,
Provider: "cpu",
ModelType: "nemo_transducer",
},
DecodingMethod: "greedy_search",
})
if recognizer == nil {
return nil, fmt.Errorf("create Parakeet recognizer: Sherpa returned nil")
}
resources := &Resources{recognizer: recognizer}
if settings.Mode == config.ModeVAD {
resources.vad = newVAD(paths.vad, settings.Threads)
if resources.vad == nil {
resources.Close()
return nil, fmt.Errorf("create Silero VAD: Sherpa returned nil")
}
}
return resources, nil
}
func (r *Resources) Close() {
if r.vad != nil {
sherpa.DeleteVoiceActivityDetector(r.vad)
r.vad = nil
}
if r.recognizer != nil {
sherpa.DeleteOfflineRecognizer(r.recognizer)
r.recognizer = nil
}
}
func (r *Resources) Decode(samples []float32) string {
if len(samples) == 0 {
return ""
}
stream := sherpa.NewOfflineStream(r.recognizer)
defer sherpa.DeleteOfflineStream(stream)
stream.AcceptWaveform(audio.SampleRate, samples)
r.recognizer.Decode(stream)
return strings.TrimSpace(stream.GetResult().Text)
}
func (r *Resources) AcceptVAD(samples []float32) {
r.vad.AcceptWaveform(samples)
}
func (r *Resources) FlushVAD() {
r.vad.Flush()
}
func (r *Resources) NextSpeechSegment() (int, []float32, bool) {
if r.vad.IsEmpty() {
return 0, nil, false
}
segment := r.vad.Front()
r.vad.Pop()
return segment.Start, segment.Samples, true
}
func validatePaths(modelsDir string, mode config.Mode) (modelPaths, error) {
parakeetDir := filepath.Join(modelsDir, "parakeet-tdt-v2")
paths := modelPaths{
encoder: filepath.Join(parakeetDir, "encoder.int8.onnx"),
decoder: filepath.Join(parakeetDir, "decoder.int8.onnx"),
joiner: filepath.Join(parakeetDir, "joiner.int8.onnx"),
tokens: filepath.Join(parakeetDir, "tokens.txt"),
vad: filepath.Join(modelsDir, "silero_vad_v5.onnx"),
}
required := []string{paths.encoder, paths.decoder, paths.joiner, paths.tokens}
if mode == config.ModeVAD {
required = append(required, paths.vad)
}
var missing []string
for _, path := range required {
info, err := os.Stat(path)
if err != nil || info.IsDir() {
missing = append(missing, path)
}
}
if len(missing) > 0 {
return modelPaths{}, fmt.Errorf("missing required model files:\n %s\nRun ./scripts/download-models.sh", strings.Join(missing, "\n "))
}
return paths, nil
}
func newVAD(model string, threads int) *sherpa.VoiceActivityDetector {
return sherpa.NewVoiceActivityDetector(&sherpa.VadModelConfig{
SileroVad: sherpa.SileroVadModelConfig{
Model: model,
Threshold: 0.5,
MinSilenceDuration: 0.5,
MinSpeechDuration: 0.25,
WindowSize: 512,
MaxSpeechDuration: 30,
},
SampleRate: audio.SampleRate,
NumThreads: threads,
Provider: "cpu",
}, 60)
}