Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,8 @@ object SpeakerDiarizationObject {
val config = OfflineSpeakerDiarizationConfig(
segmentation = OfflineSpeakerSegmentationModelConfig(
pyannote = OfflineSpeakerSegmentationPyannoteModelConfig(
segmentationModel
model = segmentationModel,
windowShiftRatio = 0.1f,
),
debug = true,
),
Expand Down
2 changes: 2 additions & 0 deletions c-api-examples/offline-speaker-diarization-c-api.c
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,8 @@ int main() {
memset(&config, 0, sizeof(config));

config.segmentation.pyannote.model = segmentation_model;
// A value of 0 also uses the default of 0.1.
config.segmentation.pyannote.window_shift_ratio = 0.1f;
config.embedding.model = embedding_extractor_model;

// the test wave ./0-four-speakers-zh.wav has 4 speakers, so
Expand Down
1 change: 1 addition & 0 deletions cxx-api-examples/offline-speaker-diarization-cxx-api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ int32_t main() {
OfflineSpeakerDiarizationConfig config;
config.segmentation.pyannote.model =
"./sherpa-onnx-pyannote-segmentation-3-0/model.onnx";
config.segmentation.pyannote.window_shift_ratio = 0.1f;
config.embedding.model =
"./3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx";
config.clustering.num_clusters = 4;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ void main(List<String> arguments) async {

final segmentationConfig = sherpa_onnx.OfflineSpeakerSegmentationModelConfig(
pyannote: sherpa_onnx.OfflineSpeakerSegmentationPyannoteModelConfig(
model: segmentationModel),
model: segmentationModel, windowShiftRatio: 0.1),
);

final embeddingConfig =
Expand Down
7 changes: 4 additions & 3 deletions dotnet-examples/offline-speaker-diarization/Program.cs
Original file line number Diff line number Diff line change
Expand Up @@ -39,9 +39,10 @@ class OfflineSpeakerDiarizationDemo
{
static void Main(string[] args)
{
var config = new OfflineSpeakerDiarizationConfig();
config.Segmentation.Pyannote.Model = "./sherpa-onnx-pyannote-segmentation-3-0/model.onnx";
config.Embedding.Model = "./3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx";
var config = new OfflineSpeakerDiarizationConfig();
config.Segmentation.Pyannote.Model = "./sherpa-onnx-pyannote-segmentation-3-0/model.onnx";
config.Segmentation.Pyannote.WindowShiftRatio = 0.1f;
config.Embedding.Model = "./3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx";

// the test wave ./0-four-speakers-zh.wav has 4 speakers, so
// we set num_clusters to 4
Expand Down
2 changes: 2 additions & 0 deletions flutter/sherpa_onnx/lib/src/offline_speaker_diarization.dart
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,8 @@ class OfflineSpeakerDiarization {

c.ref.segmentation.pyannote.model =
config.segmentation.pyannote.model.toNativeUtf8();
c.ref.segmentation.pyannote.windowShiftRatio =
config.segmentation.pyannote.windowShiftRatio;
c.ref.segmentation.numThreads = config.segmentation.numThreads;
c.ref.segmentation.debug = config.segmentation.debug ? 1 : 0;
c.ref.segmentation.provider = config.segmentation.provider.toNativeUtf8();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -43,25 +43,31 @@ class OfflineSpeakerDiarizationSegment {
class OfflineSpeakerSegmentationPyannoteModelConfig {
const OfflineSpeakerSegmentationPyannoteModelConfig({
this.model = '',
this.windowShiftRatio = 0.1,
});

factory OfflineSpeakerSegmentationPyannoteModelConfig.fromJson(
Map<String, dynamic> json) {
return OfflineSpeakerSegmentationPyannoteModelConfig(
model: json['model'] as String? ?? '',
windowShiftRatio:
(json['windowShiftRatio'] as num?)?.toDouble() ?? 0.1,
);
}

@override
String toString() {
return 'OfflineSpeakerSegmentationPyannoteModelConfig(model: $model)';
return 'OfflineSpeakerSegmentationPyannoteModelConfig(model: $model, '
'windowShiftRatio: $windowShiftRatio)';
}

Map<String, dynamic> toJson() => {
'model': model,
'windowShiftRatio': windowShiftRatio,
};

final String model;
final double windowShiftRatio;
}

/// Segmentation model configuration for speaker diarization.
Expand Down
3 changes: 3 additions & 0 deletions flutter/sherpa_onnx/lib/src/sherpa_onnx_bindings.dart
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,9 @@ final class SherpaOnnxOfflineSpeakerDiarizationSegment extends Struct {
final class SherpaOnnxOfflineSpeakerSegmentationPyannoteModelConfig
extends Struct {
external Pointer<Utf8> model;

@Float()
external double windowShiftRatio;
}

final class SherpaOnnxOfflineSpeakerSegmentationModelConfig extends Struct {
Expand Down
1 change: 1 addition & 0 deletions go-api-examples/non-streaming-speaker-diarization/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ func initSpeakerDiarization() *sherpa.OfflineSpeakerDiarization {
config := sherpa.OfflineSpeakerDiarizationConfig{}

config.Segmentation.Pyannote.Model = "./sherpa-onnx-pyannote-segmentation-3-0/model.onnx"
config.Segmentation.Pyannote.WindowShiftRatio = 0.1
config.Embedding.Model = "./3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx"

// The test wave file contains 4 speakers, so we use 4 here
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ GetOfflineSpeakerSegmentationPyannoteModelConfig(Napi::Object obj) {

Napi::Object o = obj.Get("pyannote").As<Napi::Object>();
SHERPA_ONNX_ASSIGN_ATTR_STR(model, model);
SHERPA_ONNX_ASSIGN_ATTR_FLOAT(window_shift_ratio, windowShiftRatio);

return c;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import { SpeakerEmbeddingExtractorConfig } from './SpeakerIdentification';

export class OfflineSpeakerSegmentationPyannoteModelConfig {
public model: string = '';
public windowShiftRatio: number = 0.1;
}

export class OfflineSpeakerSegmentationModelConfig {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@ function initOfflineSpeakerDiarization(context: Context): OfflineSpeakerDiarizat
//
// Also, please delete unused files to reduce the size of the app
config.segmentation.pyannote.model = 'sherpa-onnx-pyannote-segmentation-3-0/model.int8.onnx';
config.segmentation.pyannote.windowShiftRatio = 0.1;
config.segmentation.numThreads = 2;
config.segmentation.debug = true;

Expand Down Expand Up @@ -186,4 +187,4 @@ workerPort.onmessageerror = (e: MessageEvents) => {
* @param e error message
*/
workerPort.onerror = (e: ErrorEvent) => {
}
}
5 changes: 4 additions & 1 deletion java-api-examples/OfflineSpeakerDiarizationDemo.java
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,10 @@ public static void main(String[] args) {
WaveReader reader = new WaveReader(waveFilename);

OfflineSpeakerSegmentationPyannoteModelConfig pyannote =
OfflineSpeakerSegmentationPyannoteModelConfig.builder().setModel(segmentationModel).build();
OfflineSpeakerSegmentationPyannoteModelConfig.builder()
.setModel(segmentationModel)
.setWindowShiftRatio(0.1f)
.build();

OfflineSpeakerSegmentationModelConfig segmentation =
OfflineSpeakerSegmentationModelConfig.builder()
Expand Down
5 changes: 4 additions & 1 deletion kotlin-api-examples/test_offline_speaker_diarization.kt
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,10 @@ fun callback(numProcessedChunks: Int, numTotalChunks: Int, arg: Long): Int {
fun testOfflineSpeakerDiarization() {
var config = OfflineSpeakerDiarizationConfig(
segmentation=OfflineSpeakerSegmentationModelConfig(
pyannote=OfflineSpeakerSegmentationPyannoteModelConfig("./sherpa-onnx-pyannote-segmentation-3-0/model.onnx"),
pyannote=OfflineSpeakerSegmentationPyannoteModelConfig(
model="./sherpa-onnx-pyannote-segmentation-3-0/model.onnx",
windowShiftRatio=0.1f,
),
),
embedding=SpeakerEmbeddingExtractorConfig(
model="./3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx",
Expand Down
1 change: 1 addition & 0 deletions nodejs-addon-examples/test_offline_speaker_diarization.js
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ const config = {
segmentation: {
pyannote: {
model: './sherpa-onnx-pyannote-segmentation-3-0/model.onnx',
windowShiftRatio: 0.1,
},
},
embedding: {
Expand Down
1 change: 1 addition & 0 deletions nodejs-examples/test-offline-speaker-diarization.js
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ const config = {
segmentation: {
pyannote: {
model: './sherpa-onnx-pyannote-segmentation-3-0/model.onnx',
windowShiftRatio: 0.1,
debug: 1,
},
},
Expand Down
1 change: 1 addition & 0 deletions pascal-api-examples/speaker-diarization/main.pas
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@ function ProgressCallback(
Wave := SherpaOnnxReadWave('./0-four-speakers-zh.wav');

Config.Segmentation.Pyannote.Model := './sherpa-onnx-pyannote-segmentation-3-0/model.onnx';
Config.Segmentation.Pyannote.WindowShiftRatio := 0.1;
Config.Embedding.Model := './3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx';

{
Expand Down
2 changes: 1 addition & 1 deletion python-api-examples/offline-speaker-diarization.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ def init_speaker_diarization(num_speakers: int = -1, cluster_threshold: float =
config = sherpa_onnx.OfflineSpeakerDiarizationConfig(
segmentation=sherpa_onnx.OfflineSpeakerSegmentationModelConfig(
pyannote=sherpa_onnx.OfflineSpeakerSegmentationPyannoteModelConfig(
model=segmentation_model
model=segmentation_model, window_shift_ratio=0.1
),
),
embedding=sherpa_onnx.SpeakerEmbeddingExtractorConfig(
Expand Down
1 change: 1 addition & 0 deletions rust-api-examples/examples/offline_speaker_diarization.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ fn main() {
segmentation: OfflineSpeakerSegmentationModelConfig {
pyannote: OfflineSpeakerSegmentationPyannoteModelConfig {
model: Some("./sherpa-onnx-pyannote-segmentation-3-0/model.onnx".into()),
window_shift_ratio: 0.1,
},
..Default::default()
},
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,10 +11,11 @@ public struct OfflineSpeakerSegmentationPyannoteModelConfig
public OfflineSpeakerSegmentationPyannoteModelConfig()
{
Model = "";
WindowShiftRatio = 0.1f;
}

[MarshalAs(UnmanagedType.LPStr)]
public string Model;
public float WindowShiftRatio;
}
}

4 changes: 3 additions & 1 deletion scripts/go/sherpa_onnx.go
Original file line number Diff line number Diff line change
Expand Up @@ -2080,7 +2080,8 @@ func ReadWave(filename string) *Wave {
// For offline speaker diarization
// ============================================================
type OfflineSpeakerSegmentationPyannoteModelConfig struct {
Model string
Model string
WindowShiftRatio float32
}

type OfflineSpeakerSegmentationModelConfig struct {
Expand Down Expand Up @@ -2116,6 +2117,7 @@ func NewOfflineSpeakerDiarization(config *OfflineSpeakerDiarizationConfig) *Offl
c := C.struct_SherpaOnnxOfflineSpeakerDiarizationConfig{}
c.segmentation.pyannote.model = C.CString(config.Segmentation.Pyannote.Model)
defer C.free(unsafe.Pointer(c.segmentation.pyannote.model))
c.segmentation.pyannote.window_shift_ratio = C.float(config.Segmentation.Pyannote.WindowShiftRatio)

c.segmentation.num_threads = C.int(config.Segmentation.NumThreads)

Expand Down
9 changes: 5 additions & 4 deletions scripts/node-addon-api/lib/types.js
Original file line number Diff line number Diff line change
Expand Up @@ -645,10 +645,11 @@
*/

/**
* Offline speaker segmentation (pyannote) model config
* @typedef {Object} OfflineSpeakerSegmentationPyannoteModelConfig
* @property {string} [model]
*/
* Offline speaker segmentation (pyannote) model config
* @typedef {Object} OfflineSpeakerSegmentationPyannoteModelConfig
* @property {string} [model]
* @property {number} [windowShiftRatio=0.1]
*/

/**
* Offline speaker segmentation model config
Expand Down
4 changes: 4 additions & 0 deletions sherpa-onnx/c-api/c-api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -3149,6 +3149,10 @@ GetOfflineSpeakerDiarizationConfig(

sd_config.segmentation.pyannote.model =
SHERPA_ONNX_OR(config->segmentation.pyannote.model, "");
sd_config.segmentation.pyannote.window_shift_ratio =
config->segmentation.pyannote.window_shift_ratio <= 0
? 0.1f
: config->segmentation.pyannote.window_shift_ratio;
sd_config.segmentation.num_threads =
GetNumThreads(config->segmentation.num_threads);
sd_config.segmentation.debug = config->segmentation.debug;
Expand Down
7 changes: 7 additions & 0 deletions sherpa-onnx/c-api/c-api.h
Original file line number Diff line number Diff line change
Expand Up @@ -3836,6 +3836,13 @@ SHERPA_ONNX_API int32_t SherpaOnnxLinearResamplerResampleGetOutputSampleRate(
typedef struct SherpaOnnxOfflineSpeakerSegmentationPyannoteModelConfig {
/** Segmentation model filename. */
const char *model;
/**
* Sliding-window shift as a fraction of the model window size.
*
* Must be in (0, 1]. Set to 0 (or a negative value) to use the default of
* 0.1.
*/
float window_shift_ratio;
} SherpaOnnxOfflineSpeakerSegmentationPyannoteModelConfig;

/**
Expand Down
2 changes: 2 additions & 0 deletions sherpa-onnx/c-api/cxx-api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1464,6 +1464,8 @@ OfflineSpeakerDiarization OfflineSpeakerDiarization::Create(
memset(&c, 0, sizeof(c));
c.segmentation.pyannote.model =
config.segmentation.pyannote.model.c_str();
c.segmentation.pyannote.window_shift_ratio =
config.segmentation.pyannote.window_shift_ratio;
c.segmentation.num_threads = config.segmentation.num_threads;
c.segmentation.debug = config.segmentation.debug;
c.segmentation.provider = config.segmentation.provider.c_str();
Expand Down
2 changes: 2 additions & 0 deletions sherpa-onnx/c-api/cxx-api.h
Original file line number Diff line number Diff line change
Expand Up @@ -1912,6 +1912,8 @@ class SHERPA_ONNX_API SpeakerEmbeddingManager
struct OfflineSpeakerSegmentationPyannoteModelConfig {
/** Segmentation model filename. */
std::string model;
/** Sliding-window shift as a fraction of the model window size. */
float window_shift_ratio = 0.1f;
};

/** @brief Segmentation model configuration for offline speaker diarization. */
Expand Down
4 changes: 4 additions & 0 deletions sherpa-onnx/c-api/docs/speaker-diarization.dox
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,10 @@ SherpaOnnxOfflineSpeakerDiarizationConfig config;
memset(&config, 0, sizeof(config));
config.segmentation.pyannote.model =
"./sherpa-onnx-pyannote-segmentation-3-0/model.onnx";
// Zero uses the default window shift ratio of 0.1. Set a value in (0, 1] to
// change the sliding-window shift; smaller values mean more overlap and more
// compute.
config.segmentation.pyannote.window_shift_ratio = 0.1f;
config.embedding.model =
"./3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx";
config.clustering.num_clusters = 4;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,11 @@ class OfflineSpeakerSegmentationPyannoteModel::Impl {
meta_data_.window_shift = static_cast<int32_t>(window_shift);
}

if (config_.debug) {
SHERPA_ONNX_LOGE("Pyannote window shift: %d samples",
meta_data_.window_shift);
}

SHERPA_ONNX_READ_META_DATA(meta_data_.receptive_field_size,
"receptive_field_size");
SHERPA_ONNX_READ_META_DATA(meta_data_.receptive_field_shift,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,11 @@

public class OfflineSpeakerSegmentationPyannoteModelConfig {
private final String model;
private final float windowShiftRatio;

private OfflineSpeakerSegmentationPyannoteModelConfig(Builder builder) {
this.model = builder.model;
this.windowShiftRatio = builder.windowShiftRatio;
}

public static Builder builder() {
Expand All @@ -17,8 +19,13 @@ public String getModel() {
return model;
}

public float getWindowShiftRatio() {
return windowShiftRatio;
}

public static class Builder {
private String model = "";
private float windowShiftRatio = 0.1f;

public OfflineSpeakerSegmentationPyannoteModelConfig build() {
return new OfflineSpeakerSegmentationPyannoteModelConfig(this);
Expand All @@ -28,5 +35,10 @@ public Builder setModel(String model) {
this.model = model;
return this;
}

public Builder setWindowShiftRatio(float value) {
this.windowShiftRatio = value;
return this;
}
}
}
3 changes: 3 additions & 0 deletions sherpa-onnx/jni/offline-speaker-diarization.cc
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,9 @@ static OfflineSpeakerDiarizationConfig GetOfflineSpeakerDiarizationConfig(

SHERPA_ONNX_JNI_READ_STRING(ans.segmentation.pyannote.model, model,
pyannote_config_cls, pyannote_config);
SHERPA_ONNX_JNI_READ_FLOAT(ans.segmentation.pyannote.window_shift_ratio,
windowShiftRatio, pyannote_config_cls,
pyannote_config);

SHERPA_ONNX_JNI_READ_INT(ans.segmentation.num_threads, numThreads,
segmentation_config_cls, segmentation_config);
Expand Down
1 change: 1 addition & 0 deletions sherpa-onnx/kotlin-api/OfflineSpeakerDiarization.kt
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import android.content.res.AssetManager

data class OfflineSpeakerSegmentationPyannoteModelConfig(
var model: String = "",
var windowShiftRatio: Float = 0.1f,
)

data class OfflineSpeakerSegmentationModelConfig(
Expand Down
Loading
Loading