package main import ( "context" "encoding/binary" "errors" "flag" "fmt" "io" "log" "os" "os/exec" "os/signal" "path/filepath" "strings" "sync" "syscall" "time" sherpa "github.com/k2-fsa/sherpa-onnx-go-linux" ) const ( sampleRate = 16000 channels = 1 bytesPerSample = 2 packetDuration = 100 * time.Millisecond defaultModelsDir = "models" ) type settings struct { mode string chunkDuration time.Duration modelsDir string threads int previewThresholdDB float64 } func main() { cfg := parseFlags() if err := validateSettings(cfg); err != nil { log.Fatal(err) } if err := validatePrograms(); err != nil { log.Fatal(err) } paths, err := validateModels(cfg.modelsDir, cfg.mode) if err != nil { log.Fatal(err) } ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) defer stop() recognizer := newRecognizer(paths, cfg.threads) defer sherpa.DeleteOfflineRecognizer(recognizer) var vad *sherpa.VoiceActivityDetector if cfg.mode == "vad" { vad = newVAD(paths.vad, cfg.threads) defer sherpa.DeleteVoiceActivityDetector(vad) } monitorSource, err := defaultMonitorSource(ctx) if err != nil { log.Fatalf("find default output monitor: %v", err) } fmt.Fprintf(os.Stderr, "Capturing from: %s\n", monitorSource) fmt.Fprintln(os.Stderr, "Press Ctrl+C to stop.") packets, waitCapture, err := capturePackets(ctx, monitorSource) if err != nil { log.Fatalf("start system-audio capture: %v", err) } processor := newCaptionProcessor(cfg, recognizer, vad, os.Stdout) for packet := range packets { processor.accept(packet) } processor.flush() if err := waitCapture(); err != nil && ctx.Err() == nil { log.Fatalf("audio capture stopped unexpectedly: %v", err) } } func parseFlags() settings { var cfg settings flag.StringVar(&cfg.mode, "mode", "", "caption mode: vad or fixed (required)") flag.DurationVar(&cfg.chunkDuration, "chunk-duration", time.Second, "fixed chunk size or VAD preview refresh interval") flag.StringVar(&cfg.modelsDir, "models-dir", defaultModelsDir, "directory containing downloaded Sherpa models") flag.IntVar(&cfg.threads, "threads", 2, "CPU threads used by recognition and VAD") flag.Float64Var(&cfg.previewThresholdDB, "preview-threshold-dbfs", -45, "RMS dBFS threshold used to begin VAD previews") flag.Parse() return cfg } func validateSettings(cfg settings) error { if cfg.mode != "fixed" && cfg.mode != "vad" { return errors.New("--mode is required and must be either fixed or vad") } if cfg.chunkDuration <= 0 { return errors.New("--chunk-duration must be positive") } if cfg.threads <= 0 { return errors.New("--threads must be positive") } return nil } func validatePrograms() error { for _, program := range []string{"pactl", "parec"} { if _, err := exec.LookPath(program); err != nil { return fmt.Errorf("%s was not found in PATH: %w", program, err) } } return nil } type modelPaths struct { encoder string decoder string joiner string tokens string vad string } func validateModels(modelsDir, mode string) (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 == "vad" { 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 newRecognizer(paths modelPaths, threads int) *sherpa.OfflineRecognizer { recognizer := sherpa.NewOfflineRecognizer(&sherpa.OfflineRecognizerConfig{ FeatConfig: sherpa.FeatureConfig{SampleRate: sampleRate, FeatureDim: 80}, ModelConfig: sherpa.OfflineModelConfig{ Transducer: sherpa.OfflineTransducerModelConfig{ Encoder: paths.encoder, Decoder: paths.decoder, Joiner: paths.joiner, }, Tokens: paths.tokens, NumThreads: threads, Provider: "cpu", ModelType: "nemo_transducer", }, DecodingMethod: "greedy_search", }) if recognizer == nil { log.Fatal("create Parakeet recognizer: Sherpa returned nil") } return recognizer } func newVAD(model string, threads int) *sherpa.VoiceActivityDetector { vad := sherpa.NewVoiceActivityDetector(&sherpa.VadModelConfig{ SileroVad: sherpa.SileroVadModelConfig{ Model: model, Threshold: 0.5, MinSilenceDuration: 0.5, MinSpeechDuration: 0.25, WindowSize: 512, MaxSpeechDuration: 30, }, SampleRate: sampleRate, NumThreads: threads, Provider: "cpu", }, 60) if vad == nil { log.Fatal("create Silero VAD: Sherpa returned nil") } return vad } func defaultMonitorSource(ctx context.Context) (string, error) { output, err := exec.CommandContext(ctx, "pactl", "get-default-sink").Output() if err != nil { return "", err } return strings.TrimSpace(string(output)) + ".monitor", nil } func capturePackets(ctx context.Context, monitorSource string) (<-chan []float32, func() error, error) { cmd := exec.CommandContext(ctx, "parec", "--device="+monitorSource, "--format=s16le", fmt.Sprintf("--rate=%d", sampleRate), fmt.Sprintf("--channels=%d", channels), "--raw", "--latency-msec=50", ) cmd.Stderr = os.Stderr stdout, err := cmd.StdoutPipe() if err != nil { return nil, nil, err } if err := cmd.Start(); err != nil { return nil, nil, err } packets := make(chan []float32, 50) // Five seconds at 100 ms per packet. var waitOnce sync.Once var waitErr error wait := func() error { waitOnce.Do(func() { waitErr = cmd.Wait() }) return waitErr } go func() { defer close(packets) packetBytes := make([]byte, samplesPerPacket()*bytesPerSample) for { n, readErr := io.ReadFull(stdout, packetBytes) if n > 0 { packet := pcm16LEToFloat32(packetBytes[:n]) select { case packets <- packet: default: // Keep the newest audio so captions recover quickly after overload. select { case <-packets: default: } select { case packets <- packet: default: } fmt.Fprintln(os.Stderr, "warning: caption processing overloaded; dropped 100 ms of oldest audio") } } if readErr != nil { if ctx.Err() == nil && !errors.Is(readErr, io.EOF) && !errors.Is(readErr, io.ErrUnexpectedEOF) { fmt.Fprintf(os.Stderr, "audio capture read error: %v\n", readErr) } return } } }() return packets, wait, nil } func samplesPerPacket() int { return int(sampleRate * packetDuration / time.Second) } func pcm16LEToFloat32(pcm []byte) []float32 { samples := make([]float32, len(pcm)/bytesPerSample) for i := range samples { samples[i] = float32(int16(binary.LittleEndian.Uint16(pcm[i*2:]))) / 32768 } return samples }