271 lines
7.1 KiB
Go
271 lines
7.1 KiB
Go
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
|
|
}
|