// 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 Validate(settings config.Settings) error { _, err := validatePaths(settings.ModelsDir, settings.Mode) return err } 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, // A 500 ms pause ends the current caption line and starts the next one. MinSilenceDuration: 0.5, MinSpeechDuration: 0.25, WindowSize: 512, MaxSpeechDuration: 30, }, SampleRate: audio.SampleRate, NumThreads: threads, Provider: "cpu", }, 60) }