155 lines
3.6 KiB
Go
155 lines
3.6 KiB
Go
package main
|
|
|
|
import (
|
|
"fmt"
|
|
"io"
|
|
"math"
|
|
"strings"
|
|
"time"
|
|
|
|
sherpa "github.com/k2-fsa/sherpa-onnx-go-linux"
|
|
)
|
|
|
|
type captionProcessor struct {
|
|
cfg settings
|
|
recognizer *sherpa.OfflineRecognizer
|
|
vad *sherpa.VoiceActivityDetector
|
|
output io.Writer
|
|
|
|
fixedSamples []float32
|
|
previewSamples []float32
|
|
previewActive bool
|
|
previewVisible bool
|
|
nextPreviewAt int
|
|
totalSamples int
|
|
}
|
|
|
|
func newCaptionProcessor(cfg settings, recognizer *sherpa.OfflineRecognizer, vad *sherpa.VoiceActivityDetector, output io.Writer) *captionProcessor {
|
|
return &captionProcessor{
|
|
cfg: cfg,
|
|
recognizer: recognizer,
|
|
vad: vad,
|
|
output: output,
|
|
nextPreviewAt: durationSamples(cfg.chunkDuration),
|
|
}
|
|
}
|
|
|
|
func (p *captionProcessor) accept(samples []float32) {
|
|
if len(samples) == 0 {
|
|
return
|
|
}
|
|
if p.cfg.mode == "fixed" {
|
|
p.acceptFixed(samples)
|
|
return
|
|
}
|
|
p.acceptVAD(samples)
|
|
}
|
|
|
|
func (p *captionProcessor) acceptFixed(samples []float32) {
|
|
p.fixedSamples = append(p.fixedSamples, samples...)
|
|
chunkSize := durationSamples(p.cfg.chunkDuration)
|
|
for len(p.fixedSamples) >= chunkSize {
|
|
p.emitCommitted(p.totalSamples, p.totalSamples+chunkSize, p.fixedSamples[:chunkSize])
|
|
p.fixedSamples = p.fixedSamples[chunkSize:]
|
|
p.totalSamples += chunkSize
|
|
}
|
|
}
|
|
|
|
func (p *captionProcessor) acceptVAD(samples []float32) {
|
|
p.totalSamples += len(samples)
|
|
p.vad.AcceptWaveform(samples)
|
|
|
|
if p.previewActive || rmsDBFS(samples) >= p.cfg.previewThresholdDB {
|
|
p.previewActive = true
|
|
p.previewSamples = append(p.previewSamples, samples...)
|
|
if len(p.previewSamples) >= p.nextPreviewAt {
|
|
p.emitPreview(p.previewSamples)
|
|
p.nextPreviewAt += durationSamples(p.cfg.chunkDuration)
|
|
}
|
|
}
|
|
p.drainVAD()
|
|
}
|
|
|
|
func (p *captionProcessor) drainVAD() {
|
|
for !p.vad.IsEmpty() {
|
|
segment := p.vad.Front()
|
|
p.vad.Pop()
|
|
p.clearPreview()
|
|
p.emitCommitted(segment.Start, segment.Start+len(segment.Samples), segment.Samples)
|
|
p.previewSamples = nil
|
|
p.previewActive = false
|
|
p.nextPreviewAt = durationSamples(p.cfg.chunkDuration)
|
|
}
|
|
}
|
|
|
|
func (p *captionProcessor) flush() {
|
|
if p.cfg.mode == "fixed" {
|
|
if len(p.fixedSamples) > 0 {
|
|
p.emitCommitted(p.totalSamples, p.totalSamples+len(p.fixedSamples), p.fixedSamples)
|
|
}
|
|
return
|
|
}
|
|
p.vad.Flush()
|
|
p.drainVAD()
|
|
p.clearPreview()
|
|
}
|
|
|
|
func (p *captionProcessor) emitPreview(samples []float32) {
|
|
text := p.decode(samples)
|
|
if text == "" {
|
|
return
|
|
}
|
|
fmt.Fprintf(p.output, "\r\033[2K… %s", text)
|
|
p.previewVisible = true
|
|
}
|
|
|
|
func (p *captionProcessor) emitCommitted(start, end int, samples []float32) {
|
|
text := p.decode(samples)
|
|
if text == "" {
|
|
return
|
|
}
|
|
p.clearPreview()
|
|
fmt.Fprintf(p.output, "[%s-%s] %s\n", formatAudioTime(start), formatAudioTime(end), text)
|
|
}
|
|
|
|
func (p *captionProcessor) clearPreview() {
|
|
if p.previewVisible {
|
|
fmt.Fprint(p.output, "\r\033[2K")
|
|
p.previewVisible = false
|
|
}
|
|
}
|
|
|
|
func (p *captionProcessor) decode(samples []float32) string {
|
|
if len(samples) == 0 {
|
|
return ""
|
|
}
|
|
stream := sherpa.NewOfflineStream(p.recognizer)
|
|
defer sherpa.DeleteOfflineStream(stream)
|
|
stream.AcceptWaveform(sampleRate, samples)
|
|
p.recognizer.Decode(stream)
|
|
return strings.TrimSpace(stream.GetResult().Text)
|
|
}
|
|
|
|
func durationSamples(duration time.Duration) int {
|
|
return int(duration * sampleRate / time.Second)
|
|
}
|
|
|
|
func rmsDBFS(samples []float32) float64 {
|
|
if len(samples) == 0 {
|
|
return math.Inf(-1)
|
|
}
|
|
var sum float64
|
|
for _, sample := range samples {
|
|
sum += float64(sample * sample)
|
|
}
|
|
rms := math.Sqrt(sum / float64(len(samples)))
|
|
if rms == 0 {
|
|
return math.Inf(-1)
|
|
}
|
|
return 20 * math.Log10(rms)
|
|
}
|
|
|
|
func formatAudioTime(samples int) string {
|
|
return fmt.Sprintf("%.2fs", float64(samples)/sampleRate)
|
|
}
|