@@ -29,6 +29,12 @@ type SherpaBackend struct {
2929 vadWindowSize int
3030 ttsSpeed float32
3131 onlineChunkSamples int
32+
33+ // Speaker diarization (offline pyannote + embedding extractor + clustering).
34+ // diarSampleRate is reported by sherpa at create time; we cache it so
35+ // runDiarization can resample only when the input doesn't already match.
36+ diarizer uintptr
37+ diarSampleRate int
3238}
3339
3440var onnxProvider = "cpu"
@@ -128,6 +134,25 @@ var (
128134
129135 // TTS streaming callback trampoline
130136 shimTtsGenerateWithCallback func (tts uintptr , text string , sid int32 , speed float32 , cb uintptr , ud uintptr ) uintptr
137+
138+ // Diarization config + result accessors (see csrc/shim.h).
139+ shimDiarizeConfigNew func () uintptr
140+ shimDiarizeConfigFree func (uintptr )
141+ shimDiarizeConfigSetSegmentationModel func (uintptr , string )
142+ shimDiarizeConfigSetSegmentationNumThreads func (uintptr , int32 )
143+ shimDiarizeConfigSetSegmentationProvider func (uintptr , string )
144+ shimDiarizeConfigSetSegmentationDebug func (uintptr , int32 )
145+ shimDiarizeConfigSetEmbeddingModel func (uintptr , string )
146+ shimDiarizeConfigSetEmbeddingNumThreads func (uintptr , int32 )
147+ shimDiarizeConfigSetEmbeddingProvider func (uintptr , string )
148+ shimDiarizeConfigSetEmbeddingDebug func (uintptr , int32 )
149+ shimDiarizeConfigSetClusteringNumClusters func (uintptr , int32 )
150+ shimDiarizeConfigSetClusteringThreshold func (uintptr , float32 )
151+ shimDiarizeConfigSetMinDurationOn func (uintptr , float32 )
152+ shimDiarizeConfigSetMinDurationOff func (uintptr , float32 )
153+ shimCreateOfflineSpeakerDiarization func (uintptr ) uintptr
154+ shimDiarizeSetClustering func (uintptr , int32 , float32 )
155+ shimDiarizeSegmentAt func (segs uintptr , i int32 , outStart unsafe.Pointer , outEnd unsafe.Pointer , outSpeaker unsafe.Pointer )
131156)
132157
133158// libsherpa-onnx-c-api pass-throughs — called directly from Go via purego.
@@ -172,6 +197,18 @@ var (
172197 sherpaOfflineTtsGenerate func (tts uintptr , text string , sid int32 , speed float32 ) uintptr
173198 sherpaDestroyOfflineTtsGeneratedAudio func (audio uintptr )
174199 sherpaOfflineTtsSampleRate func (tts uintptr ) int32
200+
201+ // Offline speaker diarization. Result handle owns the segment-array
202+ // pointer returned by ResultSortByStartTime; destroy the segment
203+ // array first, then the result, then (at backend Free()) the diarizer.
204+ sherpaDestroyOfflineSpeakerDiarization func (sd uintptr )
205+ sherpaOfflineSpeakerDiarizationGetSampleRate func (sd uintptr ) int32
206+ sherpaOfflineSpeakerDiarizationProcess func (sd uintptr , samples unsafe.Pointer , n int32 ) uintptr
207+ sherpaOfflineSpeakerDiarizationResultGetNumSegments func (result uintptr ) int32
208+ sherpaOfflineSpeakerDiarizationResultGetNumSpeakers func (result uintptr ) int32
209+ sherpaOfflineSpeakerDiarizationResultSortByStartTime func (result uintptr ) uintptr
210+ sherpaOfflineSpeakerDiarizationDestroySegment func (segs uintptr )
211+ sherpaDestroyOfflineSpeakerDiarizationResult func (result uintptr )
175212)
176213
177214var (
@@ -292,6 +329,24 @@ func loadSherpaLibsOnce() error {
292329 {& shimSpeechSegmentStart , "sherpa_shim_speech_segment_start" },
293330 {& shimSpeechSegmentN , "sherpa_shim_speech_segment_n" },
294331 {& shimTtsGenerateWithCallback , "sherpa_shim_tts_generate_with_callback" },
332+
333+ {& shimDiarizeConfigNew , "sherpa_shim_diarize_config_new" },
334+ {& shimDiarizeConfigFree , "sherpa_shim_diarize_config_free" },
335+ {& shimDiarizeConfigSetSegmentationModel , "sherpa_shim_diarize_config_set_segmentation_model" },
336+ {& shimDiarizeConfigSetSegmentationNumThreads , "sherpa_shim_diarize_config_set_segmentation_num_threads" },
337+ {& shimDiarizeConfigSetSegmentationProvider , "sherpa_shim_diarize_config_set_segmentation_provider" },
338+ {& shimDiarizeConfigSetSegmentationDebug , "sherpa_shim_diarize_config_set_segmentation_debug" },
339+ {& shimDiarizeConfigSetEmbeddingModel , "sherpa_shim_diarize_config_set_embedding_model" },
340+ {& shimDiarizeConfigSetEmbeddingNumThreads , "sherpa_shim_diarize_config_set_embedding_num_threads" },
341+ {& shimDiarizeConfigSetEmbeddingProvider , "sherpa_shim_diarize_config_set_embedding_provider" },
342+ {& shimDiarizeConfigSetEmbeddingDebug , "sherpa_shim_diarize_config_set_embedding_debug" },
343+ {& shimDiarizeConfigSetClusteringNumClusters , "sherpa_shim_diarize_config_set_clustering_num_clusters" },
344+ {& shimDiarizeConfigSetClusteringThreshold , "sherpa_shim_diarize_config_set_clustering_threshold" },
345+ {& shimDiarizeConfigSetMinDurationOn , "sherpa_shim_diarize_config_set_min_duration_on" },
346+ {& shimDiarizeConfigSetMinDurationOff , "sherpa_shim_diarize_config_set_min_duration_off" },
347+ {& shimCreateOfflineSpeakerDiarization , "sherpa_shim_create_offline_speaker_diarization" },
348+ {& shimDiarizeSetClustering , "sherpa_shim_diarize_set_clustering" },
349+ {& shimDiarizeSegmentAt , "sherpa_shim_diarize_segment_at" },
295350 } {
296351 purego .RegisterLibFunc (r .ptr , shim , r .name )
297352 }
@@ -334,6 +389,15 @@ func loadSherpaLibsOnce() error {
334389 {& sherpaOfflineTtsGenerate , "SherpaOnnxOfflineTtsGenerate" },
335390 {& sherpaDestroyOfflineTtsGeneratedAudio , "SherpaOnnxDestroyOfflineTtsGeneratedAudio" },
336391 {& sherpaOfflineTtsSampleRate , "SherpaOnnxOfflineTtsSampleRate" },
392+
393+ {& sherpaDestroyOfflineSpeakerDiarization , "SherpaOnnxDestroyOfflineSpeakerDiarization" },
394+ {& sherpaOfflineSpeakerDiarizationGetSampleRate , "SherpaOnnxOfflineSpeakerDiarizationGetSampleRate" },
395+ {& sherpaOfflineSpeakerDiarizationProcess , "SherpaOnnxOfflineSpeakerDiarizationProcess" },
396+ {& sherpaOfflineSpeakerDiarizationResultGetNumSegments , "SherpaOnnxOfflineSpeakerDiarizationResultGetNumSegments" },
397+ {& sherpaOfflineSpeakerDiarizationResultGetNumSpeakers , "SherpaOnnxOfflineSpeakerDiarizationResultGetNumSpeakers" },
398+ {& sherpaOfflineSpeakerDiarizationResultSortByStartTime , "SherpaOnnxOfflineSpeakerDiarizationResultSortByStartTime" },
399+ {& sherpaOfflineSpeakerDiarizationDestroySegment , "SherpaOnnxOfflineSpeakerDiarizationDestroySegment" },
400+ {& sherpaDestroyOfflineSpeakerDiarizationResult , "SherpaOnnxDestroyOfflineSpeakerDiarizationResult" },
337401 } {
338402 purego .RegisterLibFunc (r .ptr , capi , r .name )
339403 }
@@ -383,6 +447,11 @@ func isVADType(t string) bool {
383447 return t == "vad"
384448}
385449
450+ func isDiarizationType (t string ) bool {
451+ t = strings .ToLower (t )
452+ return t == "diarization" || t == "diarize" || t == "speaker-diarization"
453+ }
454+
386455// Model-options prefixes recognised by this backend. Kept as typed
387456// constants so the asrFamily / loadWhisperASR / loadGenericASR paths
388457// can all speak the same vocabulary.
@@ -423,6 +492,19 @@ const (
423492 optionOnlineRule2 = "online.rule2_min_trailing_silence="
424493 optionOnlineRule3 = "online.rule3_min_utterance_length="
425494 optionOnlineChunkSamples = "online.chunk_samples="
495+
496+ // Speaker diarization (offline pyannote + speaker-embedding extractor).
497+ // `diarize.segmentation_model` overrides the auto-detected pyannote
498+ // segmentation .onnx in modelDir; `diarize.embedding_model` does the
499+ // same for the speaker-embedding extractor. `diarize.num_clusters`
500+ // pins a known speaker count at load time; per-call DiarizeRequest
501+ // fields take precedence at process time.
502+ optionDiarizeSegmentationModel = "diarize.segmentation_model="
503+ optionDiarizeEmbeddingModel = "diarize.embedding_model="
504+ optionDiarizeNumClusters = "diarize.num_clusters="
505+ optionDiarizeThreshold = "diarize.threshold="
506+ optionDiarizeMinDurationOn = "diarize.min_duration_on="
507+ optionDiarizeMinDurationOff = "diarize.min_duration_off="
426508)
427509
428510func hasOption (opts * pb.ModelOptions , prefix string ) bool {
@@ -493,6 +575,9 @@ func (s *SherpaBackend) Load(opts *pb.ModelOptions) error {
493575 if isVADType (opts .Type ) {
494576 return s .loadVAD (opts )
495577 }
578+ if isDiarizationType (opts .Type ) {
579+ return s .loadDiarization (opts )
580+ }
496581 // An explicit `subtype=...` option routes to ASR even when Type is
497582 // unset — handy for the e2e-backends harness, which doesn't know
498583 // about ModelOptions.Type.
@@ -1247,3 +1332,176 @@ func (s *SherpaBackend) TTSStream(req *pb.TTSRequest, results chan []byte) error
12471332 }
12481333 return nil
12491334}
1335+
1336+ // =============================================================
1337+ // Speaker diarization (offline)
1338+ // =============================================================
1339+ //
1340+ // Conventions:
1341+ // - opts.ModelFile is the pyannote segmentation .onnx (e.g. model.onnx
1342+ // under sherpa-onnx-pyannote-segmentation-3-0/). Override with
1343+ // `diarize.segmentation_model=` if the gallery layout differs.
1344+ // - The speaker-embedding extractor must be provided via
1345+ // `diarize.embedding_model=`. There's no reliable filename heuristic
1346+ // we can rely on (3dspeaker, NeMo, WeSpeaker all ship with
1347+ // model-specific names), so we require it to be explicit.
1348+ // - Both paths are resolved relative to opts.ModelPath if not absolute.
1349+
1350+ func (s * SherpaBackend ) loadDiarization (opts * pb.ModelOptions ) error {
1351+ if s .diarizer != 0 {
1352+ return nil
1353+ }
1354+
1355+ modelDir := filepath .Dir (opts .ModelFile )
1356+ segModel := findOptionValue (opts , optionDiarizeSegmentationModel , opts .ModelFile )
1357+ if segModel != "" && ! filepath .IsAbs (segModel ) && opts .ModelPath != "" {
1358+ segModel = filepath .Join (opts .ModelPath , segModel )
1359+ }
1360+ if ! fileExists (segModel ) {
1361+ return fmt .Errorf ("sherpa-onnx diarization: pyannote segmentation model not found at %q (set diarize.segmentation_model=...)" , segModel )
1362+ }
1363+
1364+ embModel := findOptionValue (opts , optionDiarizeEmbeddingModel , "" )
1365+ if embModel == "" {
1366+ return fmt .Errorf ("sherpa-onnx diarization: speaker-embedding model is required — pass options: [diarize.embedding_model=<path>] (e.g. 3dspeaker_speech_campplus_sv_zh-cn_16k-common.onnx)" )
1367+ }
1368+ if ! filepath .IsAbs (embModel ) {
1369+ base := opts .ModelPath
1370+ if base == "" {
1371+ base = modelDir
1372+ }
1373+ embModel = filepath .Join (base , embModel )
1374+ }
1375+ if ! fileExists (embModel ) {
1376+ return fmt .Errorf ("sherpa-onnx diarization: speaker-embedding model not found at %q" , embModel )
1377+ }
1378+
1379+ threads := int32 (1 )
1380+ if opts .Threads != 0 {
1381+ threads = opts .Threads
1382+ }
1383+
1384+ cfg := shimDiarizeConfigNew ()
1385+ defer shimDiarizeConfigFree (cfg )
1386+
1387+ shimDiarizeConfigSetSegmentationModel (cfg , segModel )
1388+ shimDiarizeConfigSetSegmentationNumThreads (cfg , threads )
1389+ shimDiarizeConfigSetSegmentationProvider (cfg , onnxProvider )
1390+ shimDiarizeConfigSetSegmentationDebug (cfg , 0 )
1391+
1392+ shimDiarizeConfigSetEmbeddingModel (cfg , embModel )
1393+ shimDiarizeConfigSetEmbeddingNumThreads (cfg , threads )
1394+ shimDiarizeConfigSetEmbeddingProvider (cfg , onnxProvider )
1395+ shimDiarizeConfigSetEmbeddingDebug (cfg , 0 )
1396+
1397+ shimDiarizeConfigSetClusteringNumClusters (cfg , findOptionInt (opts , optionDiarizeNumClusters , - 1 ))
1398+ shimDiarizeConfigSetClusteringThreshold (cfg , findOptionFloat (opts , optionDiarizeThreshold , 0.5 ))
1399+ shimDiarizeConfigSetMinDurationOn (cfg , findOptionFloat (opts , optionDiarizeMinDurationOn , 0.3 ))
1400+ shimDiarizeConfigSetMinDurationOff (cfg , findOptionFloat (opts , optionDiarizeMinDurationOff , 0.5 ))
1401+
1402+ sd := shimCreateOfflineSpeakerDiarization (cfg )
1403+ if sd == 0 {
1404+ return fmt .Errorf ("sherpa-onnx diarization: failed to create diarizer (segmentation=%s embedding=%s)" , segModel , embModel )
1405+ }
1406+ s .diarizer = sd
1407+ s .diarSampleRate = int (sherpaOfflineSpeakerDiarizationGetSampleRate (sd ))
1408+ return nil
1409+ }
1410+
1411+ // applyDiarizeOverrides re-applies clustering knobs onto an existing
1412+ // diarizer when per-call DiarizeRequest fields are set. Both -1/0 sentinels
1413+ // follow sherpa's convention: num_clusters<=0 → use threshold-based
1414+ // clustering, threshold<=0 → keep load-time default.
1415+ func (s * SherpaBackend ) applyDiarizeOverrides (req * pb.DiarizeRequest ) {
1416+ num := int32 (- 1 )
1417+ if req .NumSpeakers > 0 {
1418+ num = req .NumSpeakers
1419+ }
1420+ threshold := float32 (0 )
1421+ if req .ClusteringThreshold > 0 {
1422+ threshold = req .ClusteringThreshold
1423+ }
1424+ if num > 0 || threshold > 0 {
1425+ shimDiarizeSetClustering (s .diarizer , num , threshold )
1426+ }
1427+ }
1428+
1429+ func (s * SherpaBackend ) Diarize (req * pb.DiarizeRequest ) (pb.DiarizeResponse , error ) {
1430+ if s .diarizer == 0 {
1431+ return pb.DiarizeResponse {}, fmt .Errorf ("sherpa-onnx diarization not loaded (model must be loaded with type=diarization)" )
1432+ }
1433+ if req .Dst == "" {
1434+ return pb.DiarizeResponse {}, fmt .Errorf ("sherpa-onnx diarization: DiarizeRequest.dst (audio path) is required" )
1435+ }
1436+
1437+ dir , err := os .MkdirTemp ("" , "sherpa-diarize" )
1438+ if err != nil {
1439+ return pb.DiarizeResponse {}, fmt .Errorf ("failed to create temp dir: %w" , err )
1440+ }
1441+ defer os .RemoveAll (dir )
1442+
1443+ wavPath := filepath .Join (dir , "input.wav" )
1444+ if err := utils .AudioToWav (req .Dst , wavPath ); err != nil {
1445+ return pb.DiarizeResponse {}, fmt .Errorf ("failed to convert audio to wav: %w" , err )
1446+ }
1447+
1448+ wave := sherpaReadWave (wavPath )
1449+ if wave == 0 {
1450+ return pb.DiarizeResponse {}, fmt .Errorf ("failed to read wav %s" , wavPath )
1451+ }
1452+ defer sherpaFreeWave (wave )
1453+
1454+ sr := int (shimWaveSampleRate (wave ))
1455+ nSamples := shimWaveNumSamples (wave )
1456+ samples := shimWaveSamples (wave )
1457+ duration := float32 (nSamples ) / float32 (sr )
1458+ if sr != s .diarSampleRate {
1459+ // AudioToWav already targets 16 kHz; pyannote-3.0 also wants 16 kHz, so
1460+ // this branch should be unreachable. Fail loudly instead of silently
1461+ // passing mismatched audio to the model.
1462+ return pb.DiarizeResponse {}, fmt .Errorf ("sherpa-onnx diarization: input sample rate %d Hz does not match model %d Hz" , sr , s .diarSampleRate )
1463+ }
1464+
1465+ s .applyDiarizeOverrides (req )
1466+
1467+ result := sherpaOfflineSpeakerDiarizationProcess (s .diarizer , samples , nSamples )
1468+ if result == 0 {
1469+ return pb.DiarizeResponse {}, fmt .Errorf ("sherpa-onnx diarization: process failed" )
1470+ }
1471+ defer sherpaDestroyOfflineSpeakerDiarizationResult (result )
1472+
1473+ numSegments := sherpaOfflineSpeakerDiarizationResultGetNumSegments (result )
1474+ numSpeakers := sherpaOfflineSpeakerDiarizationResultGetNumSpeakers (result )
1475+ if numSegments <= 0 {
1476+ return pb.DiarizeResponse {
1477+ Segments : []* pb.DiarizeSegment {},
1478+ NumSpeakers : numSpeakers ,
1479+ Duration : duration ,
1480+ }, nil
1481+ }
1482+
1483+ segs := sherpaOfflineSpeakerDiarizationResultSortByStartTime (result )
1484+ if segs == 0 {
1485+ return pb.DiarizeResponse {}, fmt .Errorf ("sherpa-onnx diarization: failed to retrieve segments" )
1486+ }
1487+ defer sherpaOfflineSpeakerDiarizationDestroySegment (segs )
1488+
1489+ out := make ([]* pb.DiarizeSegment , 0 , numSegments )
1490+ for i := int32 (0 ); i < numSegments ; i ++ {
1491+ var start , end float32
1492+ var spk int32
1493+ shimDiarizeSegmentAt (segs , i ,
1494+ unsafe .Pointer (& start ), unsafe .Pointer (& end ), unsafe .Pointer (& spk ))
1495+ out = append (out , & pb.DiarizeSegment {
1496+ Id : i ,
1497+ Start : start ,
1498+ End : end ,
1499+ Speaker : strconv .FormatInt (int64 (spk ), 10 ),
1500+ })
1501+ }
1502+ return pb.DiarizeResponse {
1503+ Segments : out ,
1504+ NumSpeakers : numSpeakers ,
1505+ Duration : duration ,
1506+ }, nil
1507+ }
0 commit comments