150 lines
3.8 KiB
Go
150 lines
3.8 KiB
Go
// 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 SpeechSegment struct {
|
||
|
|
Start int
|
||
|
|
Samples []float32
|
||
|
|
}
|
||
|
|
|
||
|
|
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() (SpeechSegment, bool) {
|
||
|
|
if r.vad.IsEmpty() {
|
||
|
|
return SpeechSegment{}, false
|
||
|
|
}
|
||
|
|
segment := r.vad.Front()
|
||
|
|
r.vad.Pop()
|
||
|
|
return SpeechSegment{Start: segment.Start, Samples: 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)
|
||
|
|
}
|