//go:build darwin
|
|
package engine
|
|
import (
|
"fmt"
|
"strings"
|
"sync"
|
"time"
|
"voicesnap/internal/logger"
|
"voicesnap/internal/model"
|
|
_ "github.com/k2-fsa/sherpa-onnx-go-macos"
|
sherpa "github.com/k2-fsa/sherpa-onnx-go/sherpa_onnx"
|
)
|
|
const (
|
asrSampleRate = 16000
|
xasrHeadPaddingSamples = asrSampleRate / 4
|
xasrTailPaddingSamples = asrSampleRate + asrSampleRate/2
|
xasrReleaseTailDelay = 300 * time.Millisecond
|
)
|
|
var (
|
xasrHeadPadding = make([]float32, xasrHeadPaddingSamples)
|
xasrTailPadding = make([]float32, xasrTailPaddingSamples)
|
)
|
|
type sherpaEngine struct {
|
recognizer *sherpa.OfflineRecognizer
|
backendKind string
|
provider string
|
hwInfo string
|
}
|
|
type xasrStreamingEngine struct {
|
recognizer *sherpa.OnlineRecognizer
|
provider string
|
hwInfo string
|
mu sync.Mutex
|
}
|
|
type xasrStreamingSession struct {
|
engine *xasrStreamingEngine
|
stream *sherpa.OnlineStream
|
lastText string
|
hasAcceptedSamples bool
|
finished bool
|
}
|
|
func newPlatformEngine(resolved model.ResolvedModel) (Engine, error) {
|
if !isSupportedBackend(resolved.BackendKind) {
|
return nil, fmt.Errorf("unsupported backend kind: %s", resolved.BackendKind)
|
}
|
if resolved.BackendKind == model.BackendXASRStreaming {
|
return newXASRStreamingEngine(resolved)
|
}
|
|
providers := darwinProviders(resolved.ProviderOrder)
|
|
for _, p := range providers {
|
config, err := offlineConfigForResolvedModel(resolved, p.provider)
|
if err != nil {
|
return nil, err
|
}
|
|
recognizer := sherpa.NewOfflineRecognizer(&config)
|
if recognizer != nil {
|
info := fmt.Sprintf("%s · %s", resolved.Profile.DisplayName, p.name)
|
logger.Info("Engine initialized: %s", info)
|
return &sherpaEngine{
|
recognizer: recognizer,
|
backendKind: resolved.BackendKind,
|
provider: p.provider,
|
hwInfo: info,
|
}, nil
|
}
|
logger.Info("Failed to init with %s, trying next provider", p.name)
|
}
|
|
return nil, fmt.Errorf("failed to initialize sherpa-onnx with any provider")
|
}
|
|
func newXASRStreamingEngine(resolved model.ResolvedModel) (Engine, error) {
|
providers := darwinProviders(resolved.ProviderOrder)
|
|
for _, p := range providers {
|
config, err := onlineConfigForResolvedModel(resolved, p.provider)
|
if err != nil {
|
return nil, err
|
}
|
|
recognizer := sherpa.NewOnlineRecognizer(&config)
|
if recognizer != nil {
|
info := fmt.Sprintf("%s · %s", resolved.Profile.DisplayName, p.name)
|
logger.Info("Engine initialized: %s", info)
|
return &xasrStreamingEngine{
|
recognizer: recognizer,
|
provider: p.provider,
|
hwInfo: info,
|
}, nil
|
}
|
logger.Info("Failed to init X-ASR with %s, trying next provider", p.name)
|
}
|
|
return nil, fmt.Errorf("failed to initialize X-ASR streaming recognizer with any provider")
|
}
|
|
func offlineConfigForResolvedModel(resolved model.ResolvedModel, provider string) (sherpa.OfflineRecognizerConfig, error) {
|
config := sherpa.OfflineRecognizerConfig{}
|
config.FeatConfig.SampleRate = 16000
|
config.FeatConfig.FeatureDim = 80
|
config.ModelConfig.Tokens = resolved.Files["tokens"]
|
config.ModelConfig.NumThreads = resolved.Profile.NumThreads
|
config.ModelConfig.Provider = provider
|
config.DecodingMethod = "greedy_search"
|
|
switch resolved.BackendKind {
|
case model.BackendSenseVoice:
|
config.ModelConfig.SenseVoice.Model = resolved.Files["model"]
|
config.ModelConfig.SenseVoice.Language = resolved.Profile.LanguageParam
|
config.ModelConfig.SenseVoice.UseInverseTextNormalization = 1
|
case model.BackendMoonshine:
|
config.ModelConfig.Moonshine.Preprocessor = resolved.Files["preprocessor"]
|
config.ModelConfig.Moonshine.Encoder = resolved.Files["encoder"]
|
config.ModelConfig.Moonshine.UncachedDecoder = resolved.Files["uncached_decoder"]
|
config.ModelConfig.Moonshine.CachedDecoder = resolved.Files["cached_decoder"]
|
case model.BackendTransducer:
|
config.ModelConfig.Transducer.Encoder = resolved.Files["encoder"]
|
config.ModelConfig.Transducer.Decoder = resolved.Files["decoder"]
|
config.ModelConfig.Transducer.Joiner = resolved.Files["joiner"]
|
case model.BackendNemoTransducer:
|
config.ModelConfig.Transducer.Encoder = resolved.Files["encoder"]
|
config.ModelConfig.Transducer.Decoder = resolved.Files["decoder"]
|
config.ModelConfig.Transducer.Joiner = resolved.Files["joiner"]
|
config.ModelConfig.ModelType = model.BackendNemoTransducer
|
case model.BackendQwen3ASR:
|
if err := applyQwen3ASRConfig(&config, resolved); err != nil {
|
return sherpa.OfflineRecognizerConfig{}, err
|
}
|
default:
|
return sherpa.OfflineRecognizerConfig{}, fmt.Errorf("unsupported backend kind: %s", resolved.BackendKind)
|
}
|
|
return config, nil
|
}
|
|
func onlineConfigForResolvedModel(resolved model.ResolvedModel, provider string) (sherpa.OnlineRecognizerConfig, error) {
|
config := sherpa.OnlineRecognizerConfig{}
|
config.FeatConfig.SampleRate = 16000
|
config.FeatConfig.FeatureDim = 80
|
config.ModelConfig.Tokens = resolved.Files["tokens"]
|
config.ModelConfig.NumThreads = resolved.Profile.NumThreads
|
config.ModelConfig.Provider = provider
|
config.ModelConfig.ModelType = "zipformer2"
|
config.DecodingMethod = "greedy_search"
|
config.EnableEndpoint = 0
|
|
switch resolved.BackendKind {
|
case model.BackendXASRStreaming:
|
config.ModelConfig.Transducer.Encoder = resolved.Files["encoder"]
|
config.ModelConfig.Transducer.Decoder = resolved.Files["decoder"]
|
config.ModelConfig.Transducer.Joiner = resolved.Files["joiner"]
|
default:
|
return sherpa.OnlineRecognizerConfig{}, fmt.Errorf("unsupported online backend kind: %s", resolved.BackendKind)
|
}
|
|
return config, nil
|
}
|
|
func darwinProviders(providerOrder []string) []struct {
|
name string
|
provider string
|
} {
|
names := map[string]string{
|
"coreml": "CoreML (Apple Neural Engine)",
|
"cpu": "CPU",
|
}
|
providers := make([]struct {
|
name string
|
provider string
|
}, 0, len(providerOrder))
|
for _, provider := range providerOrder {
|
name, ok := names[provider]
|
if !ok {
|
continue
|
}
|
providers = append(providers, struct {
|
name string
|
provider string
|
}{name: name, provider: provider})
|
}
|
if len(providers) == 0 {
|
providers = append(providers, struct {
|
name string
|
provider string
|
}{name: "CPU", provider: "cpu"})
|
}
|
return providers
|
}
|
|
func (e *sherpaEngine) Recognize(samples []float32) (string, error) {
|
totalStart := time.Now()
|
streamStart := time.Now()
|
stream := sherpa.NewOfflineStream(e.recognizer)
|
streamMS := durationMS(time.Since(streamStart))
|
defer sherpa.DeleteOfflineStream(stream)
|
|
acceptStart := time.Now()
|
stream.AcceptWaveform(asrSampleRate, samples)
|
acceptMS := durationMS(time.Since(acceptStart))
|
|
decodeStart := time.Now()
|
e.recognizer.Decode(stream)
|
decodeMS := durationMS(time.Since(decodeStart))
|
resultStart := time.Now()
|
result := stream.GetResult()
|
resultMS := durationMS(time.Since(resultStart))
|
logger.Info(
|
"PERF engine_recognize_detail backend=%q provider=%q engine=%q samples=%d audio_ms=%d create_stream_ms=%d accept_ms=%d decode_ms=%d result_ms=%d total_ms=%d",
|
e.backendKind,
|
e.provider,
|
e.hwInfo,
|
len(samples),
|
audioDurationMS(samples),
|
streamMS,
|
acceptMS,
|
decodeMS,
|
resultMS,
|
durationMS(time.Since(totalStart)),
|
)
|
|
return result.Text, nil
|
}
|
|
func (e *xasrStreamingEngine) Recognize(samples []float32) (string, error) {
|
totalStart := time.Now()
|
sessionStart := time.Now()
|
session, err := e.NewStreamingSession()
|
if err != nil {
|
return "", err
|
}
|
sessionMS := durationMS(time.Since(sessionStart))
|
defer session.Close()
|
|
acceptStart := time.Now()
|
if _, err := session.Accept(samples); err != nil {
|
return "", err
|
}
|
acceptMS := durationMS(time.Since(acceptStart))
|
finishStart := time.Now()
|
text, err := session.Finish()
|
finishMS := durationMS(time.Since(finishStart))
|
if err == nil {
|
logger.Info(
|
"PERF engine_recognize_detail backend=%q provider=%q engine=%q samples=%d audio_ms=%d create_session_ms=%d accept_ms=%d finish_ms=%d total_ms=%d",
|
model.BackendXASRStreaming,
|
e.provider,
|
e.hwInfo,
|
len(samples),
|
audioDurationMS(samples),
|
sessionMS,
|
acceptMS,
|
finishMS,
|
durationMS(time.Since(totalStart)),
|
)
|
}
|
return text, err
|
}
|
|
func (e *xasrStreamingEngine) NewStreamingSession() (StreamingSession, error) {
|
if e.recognizer == nil {
|
return nil, fmt.Errorf("X-ASR recognizer is not initialized")
|
}
|
stream := sherpa.NewOnlineStream(e.recognizer)
|
if stream == nil {
|
return nil, fmt.Errorf("failed to create X-ASR streaming session")
|
}
|
return &xasrStreamingSession{
|
engine: e,
|
stream: stream,
|
}, nil
|
}
|
|
func (s *xasrStreamingSession) Accept(samples []float32) (string, error) {
|
if len(samples) == 0 {
|
return s.lastText, nil
|
}
|
if s.stream == nil || s.finished {
|
return s.lastText, nil
|
}
|
if !s.hasAcceptedSamples {
|
samples = xasrSamplesWithHeadPadding(samples)
|
s.hasAcceptedSamples = true
|
}
|
|
s.engine.mu.Lock()
|
defer s.engine.mu.Unlock()
|
|
s.stream.AcceptWaveform(asrSampleRate, samples)
|
s.lastText = s.engine.decodeReadyLocked(s.stream, s.lastText)
|
return s.lastText, nil
|
}
|
|
func (s *xasrStreamingSession) Finish() (string, error) {
|
if s.stream == nil || s.finished {
|
return s.lastText, nil
|
}
|
|
s.engine.mu.Lock()
|
defer s.engine.mu.Unlock()
|
|
s.stream.AcceptWaveform(asrSampleRate, xasrTailPadding)
|
s.lastText = s.engine.decodeReadyLocked(s.stream, s.lastText)
|
s.stream.InputFinished()
|
s.lastText = s.engine.decodeReadyLocked(s.stream, s.lastText)
|
s.finished = true
|
return s.lastText, nil
|
}
|
|
func (s *xasrStreamingSession) Close() {
|
if s.stream == nil {
|
return
|
}
|
|
s.engine.mu.Lock()
|
defer s.engine.mu.Unlock()
|
|
sherpa.DeleteOnlineStream(s.stream)
|
s.stream = nil
|
}
|
|
func (e *xasrStreamingEngine) decodeReadyLocked(stream *sherpa.OnlineStream, lastText string) string {
|
for e.recognizer.IsReady(stream) {
|
e.recognizer.Decode(stream)
|
lastText = rememberNonEmptyText(lastText, onlineResultText(e.recognizer, stream))
|
}
|
return rememberNonEmptyText(lastText, onlineResultText(e.recognizer, stream))
|
}
|
|
func onlineResultText(recognizer *sherpa.OnlineRecognizer, stream *sherpa.OnlineStream) string {
|
result := recognizer.GetResult(stream)
|
if result == nil {
|
return ""
|
}
|
return result.Text
|
}
|
|
func rememberNonEmptyText(lastText, nextText string) string {
|
if strings.TrimSpace(nextText) == "" {
|
return lastText
|
}
|
return nextText
|
}
|
|
func xasrSamplesWithHeadPadding(samples []float32) []float32 {
|
padded := make([]float32, 0, len(xasrHeadPadding)+len(samples))
|
padded = append(padded, xasrHeadPadding...)
|
padded = append(padded, samples...)
|
return padded
|
}
|
|
func (e *sherpaEngine) HardwareInfo() string {
|
return e.hwInfo
|
}
|
|
func (e *sherpaEngine) Provider() string {
|
return e.provider
|
}
|
|
func (e *sherpaEngine) ReleaseTailCaptureDelay() time.Duration {
|
if e.backendKind == model.BackendSenseVoice {
|
return senseVoiceReleaseTailDelay
|
}
|
return 0
|
}
|
|
func (e *xasrStreamingEngine) HardwareInfo() string {
|
return e.hwInfo
|
}
|
|
func (e *xasrStreamingEngine) Provider() string {
|
return e.provider
|
}
|
|
func (e *xasrStreamingEngine) ReleaseTailCaptureDelay() time.Duration {
|
return xasrReleaseTailDelay
|
}
|
|
func (e *xasrStreamingEngine) HoldPreCaptureEnabled() bool {
|
return true
|
}
|
|
func (e *sherpaEngine) Close() {
|
if e.recognizer != nil {
|
sherpa.DeleteOfflineRecognizer(e.recognizer)
|
e.recognizer = nil
|
}
|
}
|
|
func (e *xasrStreamingEngine) Close() {
|
e.mu.Lock()
|
defer e.mu.Unlock()
|
|
if e.recognizer != nil {
|
sherpa.DeleteOnlineRecognizer(e.recognizer)
|
e.recognizer = nil
|
}
|
}
|