wfloat 2.0.0__py3-none-win_amd64.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
wfloat/_core.py ADDED
@@ -0,0 +1,1795 @@
1
+ import ctypes
2
+ import os
3
+ import platform
4
+ import sys
5
+ import warnings
6
+ from pathlib import Path
7
+ from typing import Dict, List, Optional, Sequence
8
+
9
+ from ._constants import (
10
+ DEFAULT_INTENSITY,
11
+ DEFAULT_PROVIDER,
12
+ DEFAULT_SPEED,
13
+ )
14
+ from ._results import (
15
+ Audio,
16
+ GenerationResult,
17
+ LlmGenerationResult,
18
+ StreamingTranscriptionResult,
19
+ Timeline,
20
+ TimelineChunk,
21
+ TranscriptionResult,
22
+ TranscriptionSegment,
23
+ TranscriptionToken,
24
+ )
25
+
26
+ WFLOAT_TTS_FAMILY_WFLOAT_EXPRESSIVE = 1
27
+ WFLOAT_STT_FAMILY_WHISPER = 1
28
+ WFLOAT_STT_FAMILY_MOONSHINE = 2
29
+ WFLOAT_STT_FAMILY_PARAKEET_CTC = 3
30
+ WFLOAT_STT_FAMILY_PARAKEET_TDT = 4
31
+ WFLOAT_STT_FAMILY_ZIPFORMER_TRANSDUCER = 5
32
+ WFLOAT_VAD_FAMILY_SILERO = 1
33
+ WFLOAT_VAD_FAMILY_TEN = 2
34
+ WFLOAT_LLM_FAMILY_LLAMA = 1
35
+ WFLOAT_LLM_FAMILY_QWEN = 2
36
+ WFLOAT_LLM_FAMILY_SMOLLM = 3
37
+ WFLOAT_LLM_FAMILY_GEMMA = 4
38
+ WFLOAT_LLM_FAMILY_MISTRAL = 5
39
+ WFLOAT_LLM_FAMILY_PHI = 6
40
+ WFLOAT_LLM_FAMILY_LIQUID = 7
41
+ WFLOAT_STATUS_OK = 0
42
+
43
+
44
+ class _WfloatStringMapEntry(ctypes.Structure):
45
+ _fields_ = [
46
+ ("key", ctypes.c_char_p),
47
+ ("value", ctypes.c_char_p),
48
+ ]
49
+
50
+
51
+ class _WfloatAudioResult(ctypes.Structure):
52
+ _fields_ = [
53
+ ("samples", ctypes.POINTER(ctypes.c_float)),
54
+ ("sample_count", ctypes.c_size_t),
55
+ ("sample_rate", ctypes.c_int32),
56
+ ("duration_sec", ctypes.c_float),
57
+ ]
58
+
59
+
60
+ class _WfloatTimelineChunk(ctypes.Structure):
61
+ _fields_ = [
62
+ ("index", ctypes.c_int32),
63
+ ("text", ctypes.c_char_p),
64
+ ("highlight_start", ctypes.c_int32),
65
+ ("highlight_end", ctypes.c_int32),
66
+ ("start_sec", ctypes.c_float),
67
+ ("end_sec", ctypes.c_float),
68
+ ("duration_sec", ctypes.c_float),
69
+ ("progress", ctypes.c_float),
70
+ ("voice", ctypes.c_char_p),
71
+ ("sid", ctypes.c_int32),
72
+ ("segment_index", ctypes.c_int32),
73
+ ]
74
+
75
+
76
+ class _WfloatTimeline(ctypes.Structure):
77
+ _fields_ = [
78
+ ("chunks", ctypes.POINTER(_WfloatTimelineChunk)),
79
+ ("chunk_count", ctypes.c_size_t),
80
+ ("duration_sec", ctypes.c_float),
81
+ ]
82
+
83
+
84
+ class _WfloatTtsSynthesisResult(ctypes.Structure):
85
+ _fields_ = [
86
+ ("audio", _WfloatAudioResult),
87
+ ("timeline", _WfloatTimeline),
88
+ ("model_id", ctypes.c_char_p),
89
+ ("text", ctypes.c_char_p),
90
+ ]
91
+
92
+
93
+ class _WfloatTtsSynthesizeOptions(ctypes.Structure):
94
+ _fields_ = [
95
+ ("text", ctypes.c_char_p),
96
+ ("voice", ctypes.c_char_p),
97
+ ("sid", ctypes.c_int32),
98
+ ("speed", ctypes.c_float),
99
+ ("silence_padding_sec", ctypes.c_float),
100
+ ("reference_audio", ctypes.POINTER(ctypes.c_float)),
101
+ ("reference_audio_sample_count", ctypes.c_size_t),
102
+ ("reference_audio_sample_rate", ctypes.c_int32),
103
+ ("reference_text", ctypes.c_char_p),
104
+ ("num_steps", ctypes.c_int32),
105
+ ("extra_entries", ctypes.POINTER(_WfloatStringMapEntry)),
106
+ ("extra_entry_count", ctypes.c_size_t),
107
+ ]
108
+
109
+
110
+ class _WfloatTtsDialogueSegment(ctypes.Structure):
111
+ _fields_ = [
112
+ ("text", ctypes.c_char_p),
113
+ ("voice", ctypes.c_char_p),
114
+ ("sid", ctypes.c_int32),
115
+ ("speed", ctypes.c_float),
116
+ ("silence_padding_sec", ctypes.c_float),
117
+ ("extra_entries", ctypes.POINTER(_WfloatStringMapEntry)),
118
+ ("extra_entry_count", ctypes.c_size_t),
119
+ ]
120
+
121
+
122
+ class _WfloatTtsDialogueOptions(ctypes.Structure):
123
+ _fields_ = [
124
+ ("segments", ctypes.POINTER(_WfloatTtsDialogueSegment)),
125
+ ("segment_count", ctypes.c_size_t),
126
+ ("silence_between_segments_sec", ctypes.c_float),
127
+ ]
128
+
129
+
130
+ class _WfloatTtsModelInfo(ctypes.Structure):
131
+ _fields_ = [
132
+ ("model_id", ctypes.c_char_p),
133
+ ("backend", ctypes.c_char_p),
134
+ ("family", ctypes.c_char_p),
135
+ ("feature_flags", ctypes.c_uint64),
136
+ ("sample_rate", ctypes.c_int32),
137
+ ("num_speakers", ctypes.c_int32),
138
+ ]
139
+
140
+
141
+ class _WfloatTtsModelConfig(ctypes.Structure):
142
+ _fields_ = [
143
+ ("model_id", ctypes.c_char_p),
144
+ ("family", ctypes.c_int32),
145
+ ("model_path", ctypes.c_char_p),
146
+ ("tokens_path", ctypes.c_char_p),
147
+ ("data_dir", ctypes.c_char_p),
148
+ ("lexicon_path", ctypes.c_char_p),
149
+ ("voices_path", ctypes.c_char_p),
150
+ ("lang", ctypes.c_char_p),
151
+ ("acoustic_model_path", ctypes.c_char_p),
152
+ ("vocoder_path", ctypes.c_char_p),
153
+ ("encoder_path", ctypes.c_char_p),
154
+ ("decoder_path", ctypes.c_char_p),
155
+ ("text_conditioner_path", ctypes.c_char_p),
156
+ ("lm_flow_path", ctypes.c_char_p),
157
+ ("lm_main_path", ctypes.c_char_p),
158
+ ("vocab_json_path", ctypes.c_char_p),
159
+ ("token_scores_json_path", ctypes.c_char_p),
160
+ ("num_threads", ctypes.c_int32),
161
+ ("debug", ctypes.c_int32),
162
+ ("provider", ctypes.c_char_p),
163
+ ("rule_fsts", ctypes.c_char_p),
164
+ ("rule_fars", ctypes.c_char_p),
165
+ ("max_num_sentences", ctypes.c_int32),
166
+ ("silence_scale", ctypes.c_float),
167
+ ("noise_scale", ctypes.c_float),
168
+ ("noise_scale_w", ctypes.c_float),
169
+ ("length_scale", ctypes.c_float),
170
+ ("feat_scale", ctypes.c_float),
171
+ ("t_shift", ctypes.c_float),
172
+ ("target_rms", ctypes.c_float),
173
+ ("guidance_scale", ctypes.c_float),
174
+ ]
175
+
176
+
177
+ class _WfloatSttToken(ctypes.Structure):
178
+ _fields_ = [
179
+ ("text", ctypes.c_char_p),
180
+ ("start_sec", ctypes.c_float),
181
+ ("duration_sec", ctypes.c_float),
182
+ ("confidence", ctypes.c_float),
183
+ ]
184
+
185
+
186
+ class _WfloatSttSegment(ctypes.Structure):
187
+ _fields_ = [
188
+ ("text", ctypes.c_char_p),
189
+ ("start_sec", ctypes.c_float),
190
+ ("duration_sec", ctypes.c_float),
191
+ ]
192
+
193
+
194
+ class _WfloatSttTranscriptionResult(ctypes.Structure):
195
+ _fields_ = [
196
+ ("model_id", ctypes.c_char_p),
197
+ ("text", ctypes.c_char_p),
198
+ ("language", ctypes.c_char_p),
199
+ ("emotion", ctypes.c_char_p),
200
+ ("event", ctypes.c_char_p),
201
+ ("json", ctypes.c_char_p),
202
+ ("tokens", ctypes.POINTER(_WfloatSttToken)),
203
+ ("token_count", ctypes.c_size_t),
204
+ ("segments", ctypes.POINTER(_WfloatSttSegment)),
205
+ ("segment_count", ctypes.c_size_t),
206
+ ]
207
+
208
+
209
+ class _WfloatSttSessionResult(ctypes.Structure):
210
+ _fields_ = [
211
+ ("model_id", ctypes.c_char_p),
212
+ ("text", ctypes.c_char_p),
213
+ ("json", ctypes.c_char_p),
214
+ ("is_endpoint", ctypes.c_int32),
215
+ ]
216
+
217
+
218
+ class _WfloatSttModelInfo(ctypes.Structure):
219
+ _fields_ = [
220
+ ("model_id", ctypes.c_char_p),
221
+ ("backend", ctypes.c_char_p),
222
+ ("family", ctypes.c_char_p),
223
+ ("feature_flags", ctypes.c_uint64),
224
+ ("sample_rate", ctypes.c_int32),
225
+ ("supports_language_override", ctypes.c_int32),
226
+ ]
227
+
228
+
229
+ class _WfloatSttModelConfig(ctypes.Structure):
230
+ _fields_ = [
231
+ ("model_id", ctypes.c_char_p),
232
+ ("family", ctypes.c_int32),
233
+ ("model_path", ctypes.c_char_p),
234
+ ("tokens_path", ctypes.c_char_p),
235
+ ("preprocessor_path", ctypes.c_char_p),
236
+ ("encoder_path", ctypes.c_char_p),
237
+ ("decoder_path", ctypes.c_char_p),
238
+ ("joiner_path", ctypes.c_char_p),
239
+ ("uncached_decoder_path", ctypes.c_char_p),
240
+ ("cached_decoder_path", ctypes.c_char_p),
241
+ ("provider", ctypes.c_char_p),
242
+ ("language", ctypes.c_char_p),
243
+ ("task", ctypes.c_char_p),
244
+ ("hotwords_file", ctypes.c_char_p),
245
+ ("rule_fsts", ctypes.c_char_p),
246
+ ("rule_fars", ctypes.c_char_p),
247
+ ("sample_rate", ctypes.c_int32),
248
+ ("feat_dim", ctypes.c_int32),
249
+ ("num_threads", ctypes.c_int32),
250
+ ("debug", ctypes.c_int32),
251
+ ("max_active_paths", ctypes.c_int32),
252
+ ("tail_paddings", ctypes.c_int32),
253
+ ("enable_token_timestamps", ctypes.c_int32),
254
+ ("enable_segment_timestamps", ctypes.c_int32),
255
+ ("hotwords_score", ctypes.c_float),
256
+ ("blank_penalty", ctypes.c_float),
257
+ ]
258
+
259
+
260
+ class _WfloatSttTranscribeOptions(ctypes.Structure):
261
+ _fields_ = [
262
+ ("samples", ctypes.POINTER(ctypes.c_float)),
263
+ ("sample_count", ctypes.c_size_t),
264
+ ("sample_rate", ctypes.c_int32),
265
+ ("language", ctypes.c_char_p),
266
+ ("task", ctypes.c_char_p),
267
+ ("hotwords", ctypes.c_char_p),
268
+ ]
269
+
270
+
271
+ class _WfloatVadModelConfig(ctypes.Structure):
272
+ _fields_ = [
273
+ ("model_id", ctypes.c_char_p),
274
+ ("family", ctypes.c_int32),
275
+ ("model_path", ctypes.c_char_p),
276
+ ("threshold", ctypes.c_float),
277
+ ("min_silence_duration_sec", ctypes.c_float),
278
+ ("min_speech_duration_sec", ctypes.c_float),
279
+ ("max_speech_duration_sec", ctypes.c_float),
280
+ ("sample_rate", ctypes.c_int32),
281
+ ("window_size", ctypes.c_int32),
282
+ ("num_threads", ctypes.c_int32),
283
+ ("provider", ctypes.c_char_p),
284
+ ("debug", ctypes.c_int32),
285
+ ("buffer_size_in_seconds", ctypes.c_float),
286
+ ]
287
+
288
+
289
+ class _WfloatVadModelInfo(ctypes.Structure):
290
+ _fields_ = [
291
+ ("model_id", ctypes.c_char_p),
292
+ ("backend", ctypes.c_char_p),
293
+ ("family", ctypes.c_char_p),
294
+ ("feature_flags", ctypes.c_uint64),
295
+ ("sample_rate", ctypes.c_int32),
296
+ ("window_size", ctypes.c_int32),
297
+ ]
298
+
299
+
300
+ class _WfloatVadSegment(ctypes.Structure):
301
+ _fields_ = [
302
+ ("start_sample", ctypes.c_int32),
303
+ ("samples", ctypes.POINTER(ctypes.c_float)),
304
+ ("sample_count", ctypes.c_size_t),
305
+ ]
306
+
307
+
308
+ class _WfloatLlmModelConfig(ctypes.Structure):
309
+ _fields_ = [
310
+ ("model_id", ctypes.c_char_p),
311
+ ("family", ctypes.c_int32),
312
+ ("model_path", ctypes.c_char_p),
313
+ ("chat_template", ctypes.c_char_p),
314
+ ("provider", ctypes.c_char_p),
315
+ ("context_size", ctypes.c_int32),
316
+ ("num_threads", ctypes.c_int32),
317
+ ("gpu_layer_count", ctypes.c_int32),
318
+ ("seed", ctypes.c_int32),
319
+ ]
320
+
321
+
322
+ class _WfloatLlmModelInfo(ctypes.Structure):
323
+ _fields_ = [
324
+ ("model_id", ctypes.c_char_p),
325
+ ("backend", ctypes.c_char_p),
326
+ ("family", ctypes.c_char_p),
327
+ ("feature_flags", ctypes.c_uint64),
328
+ ("context_size", ctypes.c_int32),
329
+ ]
330
+
331
+
332
+ class _WfloatLlmGenerateOptions(ctypes.Structure):
333
+ _fields_ = [
334
+ ("prompt", ctypes.c_char_p),
335
+ ("max_tokens", ctypes.c_int32),
336
+ ("temperature", ctypes.c_float),
337
+ ("top_p", ctypes.c_float),
338
+ ("top_k", ctypes.c_int32),
339
+ ("repeat_penalty", ctypes.c_float),
340
+ ("seed", ctypes.c_int32),
341
+ ]
342
+
343
+
344
+ class _WfloatLlmTokenEvent(ctypes.Structure):
345
+ _fields_ = [
346
+ ("text", ctypes.c_char_p),
347
+ ("token_index", ctypes.c_int32),
348
+ ("token_id", ctypes.c_int32),
349
+ ("is_done", ctypes.c_int32),
350
+ ]
351
+
352
+
353
+ class _WfloatLlmGenerateResult(ctypes.Structure):
354
+ _fields_ = [
355
+ ("model_id", ctypes.c_char_p),
356
+ ("text", ctypes.c_char_p),
357
+ ("finish_reason", ctypes.c_char_p),
358
+ ("json", ctypes.c_char_p),
359
+ ("prompt_token_count", ctypes.c_int32),
360
+ ("completion_token_count", ctypes.c_int32),
361
+ ]
362
+
363
+
364
+ class _WfloatLlmChatMessage(ctypes.Structure):
365
+ _fields_ = [
366
+ ("role", ctypes.c_char_p),
367
+ ("content", ctypes.c_char_p),
368
+ ]
369
+
370
+
371
+ class _WfloatLlmChatTemplateOptions(ctypes.Structure):
372
+ _fields_ = [
373
+ ("messages", ctypes.POINTER(_WfloatLlmChatMessage)),
374
+ ("message_count", ctypes.c_size_t),
375
+ ("add_generation_prompt", ctypes.c_int32),
376
+ ]
377
+
378
+
379
+ class _WfloatLlmChatTemplateResult(ctypes.Structure):
380
+ _fields_ = [
381
+ ("prompt", ctypes.c_char_p),
382
+ ("chat_template", ctypes.c_char_p),
383
+ ("json", ctypes.c_char_p),
384
+ ("used_fallback", ctypes.c_int32),
385
+ ]
386
+
387
+
388
+ _WfloatLlmTokenCallback = ctypes.CFUNCTYPE(
389
+ ctypes.c_int32,
390
+ ctypes.POINTER(_WfloatLlmTokenEvent),
391
+ ctypes.c_void_p,
392
+ )
393
+
394
+
395
+ class _CoreLibraryError(ImportError):
396
+ pass
397
+
398
+
399
+ _DLL_DIRECTORY_HANDLES = []
400
+
401
+
402
+ def _decode(value: Optional[bytes]) -> str:
403
+ return value.decode("utf-8") if value else ""
404
+
405
+
406
+ def _native_dir() -> Path:
407
+ return Path(__file__).resolve().parent / "native"
408
+
409
+
410
+ def _library_names() -> tuple[str, ...]:
411
+ if sys.platform == "win32":
412
+ return ("wfloat-core.dll", "libwfloat-core.dll")
413
+ if sys.platform == "darwin":
414
+ return ("libwfloat-core.dylib",)
415
+ return ("libwfloat-core.so",)
416
+
417
+
418
+ def _iter_packaged_library_paths() -> Sequence[Path]:
419
+ native_dir = _native_dir()
420
+ return [native_dir / name for name in _library_names()]
421
+
422
+
423
+ def _iter_candidate_library_paths() -> Sequence[Path]:
424
+ candidates: List[Path] = []
425
+
426
+ env_path = os.environ.get("WFLOAT_CORE_LIBRARY")
427
+ if env_path:
428
+ candidates.append(Path(env_path))
429
+
430
+ candidates.extend(
431
+ candidate for candidate in _iter_packaged_library_paths() if candidate.exists()
432
+ )
433
+
434
+ try:
435
+ import wfloat_core
436
+
437
+ candidates.append(Path(wfloat_core.get_library_path()))
438
+ except ImportError:
439
+ pass
440
+
441
+ repo_root = Path(__file__).resolve().parents[4]
442
+ for pattern in (
443
+ "out/**/libwfloat-core.so",
444
+ "out/**/libwfloat-core.dylib",
445
+ "out/**/wfloat-core.dll",
446
+ "build/**/libwfloat-core.so",
447
+ "build/**/libwfloat-core.dylib",
448
+ "build/**/wfloat-core.dll",
449
+ ):
450
+ candidates.extend(repo_root.glob(pattern))
451
+
452
+ return candidates
453
+
454
+
455
+ def _prepare_dll_directory(candidate: Path) -> None:
456
+ if sys.platform != "win32" or not hasattr(os, "add_dll_directory"):
457
+ return
458
+
459
+ native_dir = candidate.parent
460
+ if not native_dir.exists():
461
+ return
462
+
463
+ # Keep the handle alive so dependent DLL lookup stays enabled.
464
+ _DLL_DIRECTORY_HANDLES.append(os.add_dll_directory(str(native_dir)))
465
+
466
+
467
+ def _load_core_library() -> ctypes.CDLL:
468
+ errors: List[str] = []
469
+
470
+ for candidate in _iter_candidate_library_paths():
471
+ try:
472
+ _prepare_dll_directory(candidate)
473
+ return ctypes.CDLL(str(candidate))
474
+ except OSError as exc:
475
+ errors.append(f"{candidate}: {exc}")
476
+
477
+ if errors:
478
+ raise _CoreLibraryError(
479
+ "Failed to load wfloat-core shared library. "
480
+ + " ".join(errors)
481
+ )
482
+
483
+ raise _CoreLibraryError(
484
+ "Could not find a built wfloat-core shared library. "
485
+ f"Looked in {_native_dir()} for {', '.join(_library_names())} on "
486
+ f"{sys.platform}/{platform.machine() or 'unknown'}. "
487
+ "Set WFLOAT_CORE_LIBRARY or build wfloat-core as a shared library."
488
+ )
489
+
490
+
491
+ def _prepare_library(lib: ctypes.CDLL) -> ctypes.CDLL:
492
+ lib.wfloat_tts_model_create.argtypes = [
493
+ ctypes.POINTER(_WfloatTtsModelConfig),
494
+ ctypes.POINTER(ctypes.c_void_p),
495
+ ]
496
+ lib.wfloat_tts_model_create.restype = ctypes.c_int32
497
+
498
+ lib.wfloat_tts_model_destroy.argtypes = [ctypes.c_void_p]
499
+ lib.wfloat_tts_model_destroy.restype = None
500
+
501
+ lib.wfloat_tts_model_get_info.argtypes = [
502
+ ctypes.c_void_p,
503
+ ctypes.POINTER(_WfloatTtsModelInfo),
504
+ ]
505
+ lib.wfloat_tts_model_get_info.restype = ctypes.c_int32
506
+
507
+ lib.wfloat_tts_model_synthesize.argtypes = [
508
+ ctypes.c_void_p,
509
+ ctypes.POINTER(_WfloatTtsSynthesizeOptions),
510
+ ctypes.c_void_p,
511
+ ctypes.c_void_p,
512
+ ctypes.POINTER(ctypes.POINTER(_WfloatTtsSynthesisResult)),
513
+ ]
514
+ lib.wfloat_tts_model_synthesize.restype = ctypes.c_int32
515
+
516
+ lib.wfloat_tts_model_synthesize_dialogue.argtypes = [
517
+ ctypes.c_void_p,
518
+ ctypes.POINTER(_WfloatTtsDialogueOptions),
519
+ ctypes.c_void_p,
520
+ ctypes.c_void_p,
521
+ ctypes.POINTER(ctypes.POINTER(_WfloatTtsSynthesisResult)),
522
+ ]
523
+ lib.wfloat_tts_model_synthesize_dialogue.restype = ctypes.c_int32
524
+
525
+ lib.wfloat_tts_synthesis_result_destroy.argtypes = [
526
+ ctypes.POINTER(_WfloatTtsSynthesisResult)
527
+ ]
528
+ lib.wfloat_tts_synthesis_result_destroy.restype = None
529
+
530
+ lib.wfloat_stt_model_create.argtypes = [
531
+ ctypes.POINTER(_WfloatSttModelConfig),
532
+ ctypes.POINTER(ctypes.c_void_p),
533
+ ]
534
+ lib.wfloat_stt_model_create.restype = ctypes.c_int32
535
+
536
+ lib.wfloat_stt_model_destroy.argtypes = [ctypes.c_void_p]
537
+ lib.wfloat_stt_model_destroy.restype = None
538
+
539
+ lib.wfloat_stt_model_get_info.argtypes = [
540
+ ctypes.c_void_p,
541
+ ctypes.POINTER(_WfloatSttModelInfo),
542
+ ]
543
+ lib.wfloat_stt_model_get_info.restype = ctypes.c_int32
544
+
545
+ lib.wfloat_stt_model_transcribe.argtypes = [
546
+ ctypes.c_void_p,
547
+ ctypes.POINTER(_WfloatSttTranscribeOptions),
548
+ ctypes.POINTER(ctypes.POINTER(_WfloatSttTranscriptionResult)),
549
+ ]
550
+ lib.wfloat_stt_model_transcribe.restype = ctypes.c_int32
551
+
552
+ lib.wfloat_stt_model_create_session.argtypes = [
553
+ ctypes.c_void_p,
554
+ ctypes.POINTER(ctypes.c_void_p),
555
+ ]
556
+ lib.wfloat_stt_model_create_session.restype = ctypes.c_int32
557
+
558
+ lib.wfloat_stt_session_push_audio.argtypes = [
559
+ ctypes.c_void_p,
560
+ ctypes.POINTER(ctypes.c_float),
561
+ ctypes.c_size_t,
562
+ ctypes.c_int32,
563
+ ]
564
+ lib.wfloat_stt_session_push_audio.restype = ctypes.c_int32
565
+
566
+ lib.wfloat_stt_session_get_result.argtypes = [
567
+ ctypes.c_void_p,
568
+ ctypes.POINTER(ctypes.POINTER(_WfloatSttSessionResult)),
569
+ ]
570
+ lib.wfloat_stt_session_get_result.restype = ctypes.c_int32
571
+
572
+ lib.wfloat_stt_session_finish.argtypes = [
573
+ ctypes.c_void_p,
574
+ ctypes.POINTER(ctypes.POINTER(_WfloatSttSessionResult)),
575
+ ]
576
+ lib.wfloat_stt_session_finish.restype = ctypes.c_int32
577
+
578
+ lib.wfloat_stt_session_reset.argtypes = [ctypes.c_void_p]
579
+ lib.wfloat_stt_session_reset.restype = ctypes.c_int32
580
+
581
+ lib.wfloat_stt_session_destroy.argtypes = [ctypes.c_void_p]
582
+ lib.wfloat_stt_session_destroy.restype = None
583
+
584
+ lib.wfloat_stt_session_result_destroy.argtypes = [
585
+ ctypes.POINTER(_WfloatSttSessionResult)
586
+ ]
587
+ lib.wfloat_stt_session_result_destroy.restype = None
588
+
589
+ lib.wfloat_stt_transcription_result_destroy.argtypes = [
590
+ ctypes.POINTER(_WfloatSttTranscriptionResult)
591
+ ]
592
+ lib.wfloat_stt_transcription_result_destroy.restype = None
593
+
594
+ lib.wfloat_vad_model_create.argtypes = [
595
+ ctypes.POINTER(_WfloatVadModelConfig),
596
+ ctypes.POINTER(ctypes.c_void_p),
597
+ ]
598
+ lib.wfloat_vad_model_create.restype = ctypes.c_int32
599
+
600
+ lib.wfloat_vad_model_destroy.argtypes = [ctypes.c_void_p]
601
+ lib.wfloat_vad_model_destroy.restype = None
602
+
603
+ lib.wfloat_vad_model_get_info.argtypes = [
604
+ ctypes.c_void_p,
605
+ ctypes.POINTER(_WfloatVadModelInfo),
606
+ ]
607
+ lib.wfloat_vad_model_get_info.restype = ctypes.c_int32
608
+
609
+ lib.wfloat_vad_model_accept_waveform.argtypes = [
610
+ ctypes.c_void_p,
611
+ ctypes.POINTER(ctypes.c_float),
612
+ ctypes.c_size_t,
613
+ ]
614
+ lib.wfloat_vad_model_accept_waveform.restype = ctypes.c_int32
615
+
616
+ lib.wfloat_vad_model_reset.argtypes = [ctypes.c_void_p]
617
+ lib.wfloat_vad_model_reset.restype = ctypes.c_int32
618
+
619
+ lib.wfloat_vad_model_flush.argtypes = [ctypes.c_void_p]
620
+ lib.wfloat_vad_model_flush.restype = ctypes.c_int32
621
+
622
+ lib.wfloat_vad_model_empty.argtypes = [
623
+ ctypes.c_void_p,
624
+ ctypes.POINTER(ctypes.c_int32),
625
+ ]
626
+ lib.wfloat_vad_model_empty.restype = ctypes.c_int32
627
+
628
+ lib.wfloat_vad_model_detected.argtypes = [
629
+ ctypes.c_void_p,
630
+ ctypes.POINTER(ctypes.c_int32),
631
+ ]
632
+ lib.wfloat_vad_model_detected.restype = ctypes.c_int32
633
+
634
+ lib.wfloat_vad_model_front.argtypes = [
635
+ ctypes.c_void_p,
636
+ ctypes.POINTER(ctypes.POINTER(_WfloatVadSegment)),
637
+ ]
638
+ lib.wfloat_vad_model_front.restype = ctypes.c_int32
639
+
640
+ lib.wfloat_vad_model_pop.argtypes = [ctypes.c_void_p]
641
+ lib.wfloat_vad_model_pop.restype = ctypes.c_int32
642
+
643
+ lib.wfloat_vad_model_clear.argtypes = [ctypes.c_void_p]
644
+ lib.wfloat_vad_model_clear.restype = ctypes.c_int32
645
+
646
+ lib.wfloat_vad_segment_destroy.argtypes = [
647
+ ctypes.POINTER(_WfloatVadSegment)
648
+ ]
649
+ lib.wfloat_vad_segment_destroy.restype = None
650
+
651
+ lib.wfloat_llm_model_create.argtypes = [
652
+ ctypes.POINTER(_WfloatLlmModelConfig),
653
+ ctypes.POINTER(ctypes.c_void_p),
654
+ ]
655
+ lib.wfloat_llm_model_create.restype = ctypes.c_int32
656
+
657
+ lib.wfloat_llm_model_destroy.argtypes = [ctypes.c_void_p]
658
+ lib.wfloat_llm_model_destroy.restype = None
659
+
660
+ lib.wfloat_llm_model_get_info.argtypes = [
661
+ ctypes.c_void_p,
662
+ ctypes.POINTER(_WfloatLlmModelInfo),
663
+ ]
664
+ lib.wfloat_llm_model_get_info.restype = ctypes.c_int32
665
+
666
+ lib.wfloat_llm_model_generate.argtypes = [
667
+ ctypes.c_void_p,
668
+ ctypes.POINTER(_WfloatLlmGenerateOptions),
669
+ _WfloatLlmTokenCallback,
670
+ ctypes.c_void_p,
671
+ ctypes.POINTER(ctypes.POINTER(_WfloatLlmGenerateResult)),
672
+ ]
673
+ lib.wfloat_llm_model_generate.restype = ctypes.c_int32
674
+
675
+ lib.wfloat_llm_generate_result_destroy.argtypes = [
676
+ ctypes.POINTER(_WfloatLlmGenerateResult)
677
+ ]
678
+ lib.wfloat_llm_generate_result_destroy.restype = None
679
+
680
+ lib.wfloat_llm_model_format_chat.argtypes = [
681
+ ctypes.c_void_p,
682
+ ctypes.POINTER(_WfloatLlmChatTemplateOptions),
683
+ ctypes.POINTER(ctypes.POINTER(_WfloatLlmChatTemplateResult)),
684
+ ]
685
+ lib.wfloat_llm_model_format_chat.restype = ctypes.c_int32
686
+
687
+ lib.wfloat_llm_chat_template_result_destroy.argtypes = [
688
+ ctypes.POINTER(_WfloatLlmChatTemplateResult)
689
+ ]
690
+ lib.wfloat_llm_chat_template_result_destroy.restype = None
691
+
692
+ return lib
693
+
694
+
695
+ class CoreTts:
696
+ def __init__(
697
+ self,
698
+ model_id: str,
699
+ model_path: Path,
700
+ tokens_path: Path,
701
+ espeak_data_dir: Path,
702
+ ) -> None:
703
+ self._lib = _prepare_library(_load_core_library())
704
+ self._model = ctypes.c_void_p()
705
+
706
+ self._config_bytes = {
707
+ "model_id": model_id.encode("utf-8"),
708
+ "model_path": str(model_path).encode("utf-8"),
709
+ "tokens_path": str(tokens_path).encode("utf-8"),
710
+ "data_dir": str(espeak_data_dir).encode("utf-8"),
711
+ "provider": DEFAULT_PROVIDER.encode("utf-8"),
712
+ }
713
+
714
+ config = _WfloatTtsModelConfig(
715
+ model_id=self._config_bytes["model_id"],
716
+ family=WFLOAT_TTS_FAMILY_WFLOAT_EXPRESSIVE,
717
+ model_path=self._config_bytes["model_path"],
718
+ tokens_path=self._config_bytes["tokens_path"],
719
+ data_dir=self._config_bytes["data_dir"],
720
+ lexicon_path=None,
721
+ voices_path=None,
722
+ lang=None,
723
+ acoustic_model_path=None,
724
+ vocoder_path=None,
725
+ encoder_path=None,
726
+ decoder_path=None,
727
+ text_conditioner_path=None,
728
+ lm_flow_path=None,
729
+ lm_main_path=None,
730
+ vocab_json_path=None,
731
+ token_scores_json_path=None,
732
+ num_threads=1,
733
+ debug=0,
734
+ provider=self._config_bytes["provider"],
735
+ rule_fsts=None,
736
+ rule_fars=None,
737
+ max_num_sentences=1,
738
+ silence_scale=0.2,
739
+ noise_scale=0.667,
740
+ noise_scale_w=0.8,
741
+ length_scale=1.0,
742
+ feat_scale=0.0,
743
+ t_shift=0.0,
744
+ target_rms=0.0,
745
+ guidance_scale=0.0,
746
+ )
747
+
748
+ status = self._lib.wfloat_tts_model_create(
749
+ ctypes.byref(config),
750
+ ctypes.byref(self._model),
751
+ )
752
+ if status != WFLOAT_STATUS_OK:
753
+ raise RuntimeError(f"wfloat-core model creation failed with status {status}.")
754
+
755
+ info = _WfloatTtsModelInfo()
756
+ status = self._lib.wfloat_tts_model_get_info(self._model, ctypes.byref(info))
757
+ if status != WFLOAT_STATUS_OK:
758
+ self.close()
759
+ raise RuntimeError(f"wfloat-core model info failed with status {status}.")
760
+
761
+ self.sample_rate = int(info.sample_rate)
762
+ self.num_speakers = int(info.num_speakers)
763
+
764
+ def close(self) -> None:
765
+ if self._model and self._model.value:
766
+ self._lib.wfloat_tts_model_destroy(self._model)
767
+ self._model = ctypes.c_void_p()
768
+
769
+ def __del__(self) -> None:
770
+ try:
771
+ self.close()
772
+ except Exception:
773
+ pass
774
+
775
+ def synthesize_result(
776
+ self,
777
+ *,
778
+ model_id: str,
779
+ text: str,
780
+ voice: Optional[object],
781
+ sid: int,
782
+ emotion: str,
783
+ intensity: float,
784
+ speed: float,
785
+ silence_padding_sec: float,
786
+ ) -> GenerationResult:
787
+ text_bytes = text.encode("utf-8")
788
+ voice_bytes = None if voice is None else str(voice).encode("utf-8")
789
+ extra_entries_storage = [
790
+ _WfloatStringMapEntry(b"emotion", emotion.encode("utf-8")),
791
+ _WfloatStringMapEntry(b"intensity", str(float(intensity)).encode("utf-8")),
792
+ ]
793
+ extra_entries = (_WfloatStringMapEntry * len(extra_entries_storage))(
794
+ *extra_entries_storage
795
+ )
796
+ options = _WfloatTtsSynthesizeOptions(
797
+ text=text_bytes,
798
+ voice=voice_bytes,
799
+ sid=sid,
800
+ speed=float(speed),
801
+ silence_padding_sec=float(silence_padding_sec),
802
+ reference_audio=None,
803
+ reference_audio_sample_count=0,
804
+ reference_audio_sample_rate=0,
805
+ reference_text=None,
806
+ num_steps=0,
807
+ extra_entries=extra_entries,
808
+ extra_entry_count=len(extra_entries_storage),
809
+ )
810
+ return self._run_synthesize(
811
+ model_id=model_id,
812
+ text=text,
813
+ emotion=emotion,
814
+ intensity=float(intensity),
815
+ speed=float(speed),
816
+ voice_by_segment={None: voice},
817
+ defaults_by_segment={None: {"sid": sid, "emotion": emotion, "intensity": intensity, "speed": speed}},
818
+ options=options,
819
+ )
820
+
821
+ def synthesize_dialogue_result(
822
+ self,
823
+ *,
824
+ model_id: str,
825
+ segments: Sequence[Dict[str, object]],
826
+ silence_between_segments_sec: float,
827
+ ) -> GenerationResult:
828
+ segment_structs = []
829
+ segment_buffers = []
830
+ defaults_by_segment: Dict[Optional[int], Dict[str, object]] = {}
831
+ voice_by_segment: Dict[Optional[int], Optional[object]] = {}
832
+
833
+ for index, segment in enumerate(segments):
834
+ text = str(segment["text"])
835
+ voice = segment.get("voice_id")
836
+ sid = int(segment["sid"])
837
+ emotion = str(segment["emotion"])
838
+ intensity = float(segment["intensity"])
839
+ speed = float(segment["speed"])
840
+ silence_padding_sec = float(segment["sentence_silence_padding_sec"])
841
+
842
+ text_bytes = text.encode("utf-8")
843
+ voice_bytes = None if voice is None else str(voice).encode("utf-8")
844
+ extras_storage = [
845
+ _WfloatStringMapEntry(b"emotion", emotion.encode("utf-8")),
846
+ _WfloatStringMapEntry(b"intensity", str(intensity).encode("utf-8")),
847
+ ]
848
+ extras = (_WfloatStringMapEntry * len(extras_storage))(*extras_storage)
849
+
850
+ segment_structs.append(
851
+ _WfloatTtsDialogueSegment(
852
+ text=text_bytes,
853
+ voice=voice_bytes,
854
+ sid=sid,
855
+ speed=speed,
856
+ silence_padding_sec=silence_padding_sec,
857
+ extra_entries=extras,
858
+ extra_entry_count=len(extras_storage),
859
+ )
860
+ )
861
+ segment_buffers.append((text_bytes, voice_bytes, extras_storage, extras))
862
+ defaults_by_segment[index] = {
863
+ "sid": sid,
864
+ "emotion": emotion,
865
+ "intensity": intensity,
866
+ "speed": speed,
867
+ }
868
+ voice_by_segment[index] = voice
869
+
870
+ segments_array = (_WfloatTtsDialogueSegment * len(segment_structs))(*segment_structs)
871
+ options = _WfloatTtsDialogueOptions(
872
+ segments=segments_array,
873
+ segment_count=len(segment_structs),
874
+ silence_between_segments_sec=float(silence_between_segments_sec),
875
+ )
876
+
877
+ return self._run_synthesize_dialogue(
878
+ model_id=model_id,
879
+ text="\n".join(str(segment["text"]) for segment in segments),
880
+ voice_by_segment=voice_by_segment,
881
+ defaults_by_segment=defaults_by_segment,
882
+ options=options,
883
+ )
884
+
885
+ def _run_synthesize(
886
+ self,
887
+ *,
888
+ model_id: str,
889
+ text: str,
890
+ emotion: str,
891
+ intensity: float,
892
+ speed: float,
893
+ voice_by_segment: Dict[Optional[int], Optional[object]],
894
+ defaults_by_segment: Dict[Optional[int], Dict[str, object]],
895
+ options: _WfloatTtsSynthesizeOptions,
896
+ ) -> GenerationResult:
897
+ result_ptr = ctypes.POINTER(_WfloatTtsSynthesisResult)()
898
+ status = self._lib.wfloat_tts_model_synthesize(
899
+ self._model,
900
+ ctypes.byref(options),
901
+ None,
902
+ None,
903
+ ctypes.byref(result_ptr),
904
+ )
905
+ if status != WFLOAT_STATUS_OK:
906
+ raise RuntimeError(f"wfloat-core synthesize failed with status {status}.")
907
+
908
+ try:
909
+ return self._convert_result(
910
+ model_id=model_id,
911
+ fallback_text=text,
912
+ voice_by_segment=voice_by_segment,
913
+ defaults_by_segment=defaults_by_segment,
914
+ result_ptr=result_ptr,
915
+ )
916
+ finally:
917
+ self._lib.wfloat_tts_synthesis_result_destroy(result_ptr)
918
+
919
+ def _run_synthesize_dialogue(
920
+ self,
921
+ *,
922
+ model_id: str,
923
+ text: str,
924
+ voice_by_segment: Dict[Optional[int], Optional[object]],
925
+ defaults_by_segment: Dict[Optional[int], Dict[str, object]],
926
+ options: _WfloatTtsDialogueOptions,
927
+ ) -> GenerationResult:
928
+ result_ptr = ctypes.POINTER(_WfloatTtsSynthesisResult)()
929
+ status = self._lib.wfloat_tts_model_synthesize_dialogue(
930
+ self._model,
931
+ ctypes.byref(options),
932
+ None,
933
+ None,
934
+ ctypes.byref(result_ptr),
935
+ )
936
+ if status != WFLOAT_STATUS_OK:
937
+ raise RuntimeError(
938
+ f"wfloat-core synthesize_dialogue failed with status {status}."
939
+ )
940
+
941
+ try:
942
+ return self._convert_result(
943
+ model_id=model_id,
944
+ fallback_text=text,
945
+ voice_by_segment=voice_by_segment,
946
+ defaults_by_segment=defaults_by_segment,
947
+ result_ptr=result_ptr,
948
+ )
949
+ finally:
950
+ self._lib.wfloat_tts_synthesis_result_destroy(result_ptr)
951
+
952
+ def _convert_result(
953
+ self,
954
+ *,
955
+ model_id: str,
956
+ fallback_text: str,
957
+ voice_by_segment: Dict[Optional[int], Optional[object]],
958
+ defaults_by_segment: Dict[Optional[int], Dict[str, object]],
959
+ result_ptr: ctypes.POINTER(_WfloatTtsSynthesisResult),
960
+ ) -> GenerationResult:
961
+ result = result_ptr.contents
962
+ audio_samples = [
963
+ float(result.audio.samples[index]) for index in range(int(result.audio.sample_count))
964
+ ]
965
+ audio = Audio(
966
+ samples=audio_samples,
967
+ sample_rate=int(result.audio.sample_rate),
968
+ )
969
+
970
+ timeline_chunks: List[TimelineChunk] = []
971
+ for index in range(int(result.timeline.chunk_count)):
972
+ chunk = result.timeline.chunks[index]
973
+ segment_index = int(chunk.segment_index)
974
+ if segment_index < 0:
975
+ segment_index = None
976
+
977
+ defaults = defaults_by_segment.get(segment_index, defaults_by_segment.get(None, {}))
978
+ timeline_chunks.append(
979
+ TimelineChunk(
980
+ index=int(chunk.index),
981
+ text=_decode(chunk.text),
982
+ highlight_start=int(chunk.highlight_start),
983
+ highlight_end=int(chunk.highlight_end),
984
+ start_sec=float(chunk.start_sec),
985
+ end_sec=float(chunk.end_sec),
986
+ duration_sec=float(chunk.duration_sec),
987
+ progress=float(chunk.progress),
988
+ voice_id=voice_by_segment.get(segment_index, voice_by_segment.get(None)),
989
+ sid=int(defaults.get("sid", chunk.sid)),
990
+ emotion=str(defaults.get("emotion", "neutral")),
991
+ intensity=float(defaults.get("intensity", DEFAULT_INTENSITY)),
992
+ speed=float(defaults.get("speed", DEFAULT_SPEED)),
993
+ segment_index=segment_index,
994
+ )
995
+ )
996
+
997
+ timeline = Timeline(
998
+ chunks=timeline_chunks,
999
+ duration_sec=float(result.timeline.duration_sec or audio.duration_sec),
1000
+ )
1001
+ return GenerationResult(
1002
+ audio=audio,
1003
+ timeline=timeline,
1004
+ text=_decode(result.text) or fallback_text,
1005
+ model_name=_decode(result.model_id) or model_id,
1006
+ )
1007
+
1008
+
1009
+ def create_core_tts(
1010
+ model_name: str,
1011
+ model_path: Path,
1012
+ tokens_path: Path,
1013
+ espeak_data_dir: Path,
1014
+ ):
1015
+ return CoreTts(
1016
+ model_id=model_name,
1017
+ model_path=model_path,
1018
+ tokens_path=tokens_path,
1019
+ espeak_data_dir=espeak_data_dir,
1020
+ )
1021
+
1022
+
1023
+ class CoreStt:
1024
+ def __init__(
1025
+ self,
1026
+ *,
1027
+ model_id: str,
1028
+ family: int,
1029
+ model_path: Optional[Path],
1030
+ tokens_path: Path,
1031
+ preprocessor_path: Optional[Path] = None,
1032
+ encoder_path: Optional[Path] = None,
1033
+ decoder_path: Optional[Path] = None,
1034
+ joiner_path: Optional[Path] = None,
1035
+ uncached_decoder_path: Optional[Path] = None,
1036
+ cached_decoder_path: Optional[Path] = None,
1037
+ language: Optional[str] = None,
1038
+ task: Optional[str] = None,
1039
+ enable_token_timestamps: bool = False,
1040
+ enable_segment_timestamps: bool = False,
1041
+ ) -> None:
1042
+ self._lib = _prepare_library(_load_core_library())
1043
+ self._model = ctypes.c_void_p()
1044
+
1045
+ self._config_bytes = {
1046
+ "model_id": model_id.encode("utf-8"),
1047
+ "model_path": None if model_path is None else str(model_path).encode("utf-8"),
1048
+ "tokens_path": str(tokens_path).encode("utf-8"),
1049
+ "preprocessor_path": None
1050
+ if preprocessor_path is None
1051
+ else str(preprocessor_path).encode("utf-8"),
1052
+ "encoder_path": None if encoder_path is None else str(encoder_path).encode("utf-8"),
1053
+ "decoder_path": None if decoder_path is None else str(decoder_path).encode("utf-8"),
1054
+ "joiner_path": None if joiner_path is None else str(joiner_path).encode("utf-8"),
1055
+ "uncached_decoder_path": None
1056
+ if uncached_decoder_path is None
1057
+ else str(uncached_decoder_path).encode("utf-8"),
1058
+ "cached_decoder_path": None
1059
+ if cached_decoder_path is None
1060
+ else str(cached_decoder_path).encode("utf-8"),
1061
+ "provider": DEFAULT_PROVIDER.encode("utf-8"),
1062
+ "language": None if language is None else language.encode("utf-8"),
1063
+ "task": None if task is None else task.encode("utf-8"),
1064
+ }
1065
+
1066
+ config = _WfloatSttModelConfig(
1067
+ model_id=self._config_bytes["model_id"],
1068
+ family=family,
1069
+ model_path=self._config_bytes["model_path"],
1070
+ tokens_path=self._config_bytes["tokens_path"],
1071
+ preprocessor_path=self._config_bytes["preprocessor_path"],
1072
+ encoder_path=self._config_bytes["encoder_path"],
1073
+ decoder_path=self._config_bytes["decoder_path"],
1074
+ joiner_path=self._config_bytes["joiner_path"],
1075
+ uncached_decoder_path=self._config_bytes["uncached_decoder_path"],
1076
+ cached_decoder_path=self._config_bytes["cached_decoder_path"],
1077
+ provider=self._config_bytes["provider"],
1078
+ language=self._config_bytes["language"],
1079
+ task=self._config_bytes["task"],
1080
+ hotwords_file=None,
1081
+ rule_fsts=None,
1082
+ rule_fars=None,
1083
+ sample_rate=16000,
1084
+ feat_dim=80,
1085
+ num_threads=1,
1086
+ debug=0,
1087
+ max_active_paths=4,
1088
+ tail_paddings=0,
1089
+ enable_token_timestamps=1 if enable_token_timestamps else 0,
1090
+ enable_segment_timestamps=1 if enable_segment_timestamps else 0,
1091
+ hotwords_score=1.5,
1092
+ blank_penalty=0.0,
1093
+ )
1094
+
1095
+ status = self._lib.wfloat_stt_model_create(
1096
+ ctypes.byref(config),
1097
+ ctypes.byref(self._model),
1098
+ )
1099
+ if status != WFLOAT_STATUS_OK:
1100
+ raise RuntimeError(f"wfloat-core STT model creation failed with status {status}.")
1101
+
1102
+ info = _WfloatSttModelInfo()
1103
+ status = self._lib.wfloat_stt_model_get_info(self._model, ctypes.byref(info))
1104
+ if status != WFLOAT_STATUS_OK:
1105
+ self.close()
1106
+ raise RuntimeError(f"wfloat-core STT model info failed with status {status}.")
1107
+
1108
+ self.sample_rate = int(info.sample_rate)
1109
+ self.supports_language_override = bool(info.supports_language_override)
1110
+
1111
+ def close(self) -> None:
1112
+ if self._model and self._model.value:
1113
+ self._lib.wfloat_stt_model_destroy(self._model)
1114
+ self._model = ctypes.c_void_p()
1115
+
1116
+ def __del__(self) -> None:
1117
+ try:
1118
+ self.close()
1119
+ except Exception:
1120
+ pass
1121
+
1122
+ def transcribe_result(
1123
+ self,
1124
+ *,
1125
+ model_id: str,
1126
+ samples: Sequence[float],
1127
+ sample_rate: int,
1128
+ language: Optional[str] = None,
1129
+ task: Optional[str] = None,
1130
+ hotwords: Optional[str] = None,
1131
+ ) -> TranscriptionResult:
1132
+ sample_values = [float(sample) for sample in samples]
1133
+ sample_array = (ctypes.c_float * len(sample_values))(*sample_values)
1134
+ language_bytes = None if language is None else language.encode("utf-8")
1135
+ task_bytes = None if task is None else task.encode("utf-8")
1136
+ hotwords_bytes = None if hotwords is None else hotwords.encode("utf-8")
1137
+
1138
+ options = _WfloatSttTranscribeOptions(
1139
+ samples=sample_array,
1140
+ sample_count=len(sample_values),
1141
+ sample_rate=int(sample_rate),
1142
+ language=language_bytes,
1143
+ task=task_bytes,
1144
+ hotwords=hotwords_bytes,
1145
+ )
1146
+
1147
+ result_ptr = ctypes.POINTER(_WfloatSttTranscriptionResult)()
1148
+ status = self._lib.wfloat_stt_model_transcribe(
1149
+ self._model,
1150
+ ctypes.byref(options),
1151
+ ctypes.byref(result_ptr),
1152
+ )
1153
+ if status != WFLOAT_STATUS_OK:
1154
+ raise RuntimeError(f"wfloat-core transcribe failed with status {status}.")
1155
+
1156
+ try:
1157
+ result = result_ptr.contents
1158
+ tokens = [
1159
+ TranscriptionToken(
1160
+ text=_decode(result.tokens[index].text),
1161
+ start_sec=float(result.tokens[index].start_sec),
1162
+ duration_sec=float(result.tokens[index].duration_sec),
1163
+ confidence=float(result.tokens[index].confidence),
1164
+ )
1165
+ for index in range(int(result.token_count))
1166
+ ]
1167
+ segments = [
1168
+ TranscriptionSegment(
1169
+ text=_decode(result.segments[index].text),
1170
+ start_sec=float(result.segments[index].start_sec),
1171
+ duration_sec=float(result.segments[index].duration_sec),
1172
+ )
1173
+ for index in range(int(result.segment_count))
1174
+ ]
1175
+
1176
+ return TranscriptionResult(
1177
+ text=_decode(result.text),
1178
+ model_id=_decode(result.model_id) or model_id,
1179
+ language=_decode(result.language),
1180
+ emotion=_decode(result.emotion),
1181
+ event=_decode(result.event),
1182
+ json=_decode(result.json),
1183
+ tokens=tokens or None,
1184
+ segments=segments or None,
1185
+ )
1186
+ finally:
1187
+ self._lib.wfloat_stt_transcription_result_destroy(result_ptr)
1188
+
1189
+ def create_session(self):
1190
+ session = ctypes.c_void_p()
1191
+ status = self._lib.wfloat_stt_model_create_session(
1192
+ self._model,
1193
+ ctypes.byref(session),
1194
+ )
1195
+ if status != WFLOAT_STATUS_OK:
1196
+ raise RuntimeError(f"wfloat-core create_session failed with status {status}.")
1197
+
1198
+ return CoreSttSession(
1199
+ lib=self._lib,
1200
+ model_id=_decode(self._config_bytes["model_id"]),
1201
+ session=session,
1202
+ sample_rate=self.sample_rate,
1203
+ )
1204
+
1205
+
1206
+ class CoreSttSession:
1207
+ def __init__(
1208
+ self,
1209
+ *,
1210
+ lib: ctypes.CDLL,
1211
+ model_id: str,
1212
+ session: ctypes.c_void_p,
1213
+ sample_rate: int,
1214
+ ) -> None:
1215
+ self._lib = lib
1216
+ self._model_id = model_id
1217
+ self._session = session
1218
+ self.sample_rate = int(sample_rate)
1219
+
1220
+ def close(self) -> None:
1221
+ if self._session and self._session.value:
1222
+ self._lib.wfloat_stt_session_destroy(self._session)
1223
+ self._session = ctypes.c_void_p()
1224
+
1225
+ def __del__(self) -> None:
1226
+ try:
1227
+ self.close()
1228
+ except Exception:
1229
+ pass
1230
+
1231
+ def push(self, samples: Sequence[float], sample_rate: Optional[int] = None) -> None:
1232
+ sample_values = [float(sample) for sample in samples]
1233
+ if not sample_values:
1234
+ raise ValueError("samples must not be empty.")
1235
+
1236
+ resolved_sample_rate = int(sample_rate or self.sample_rate)
1237
+ sample_array = (ctypes.c_float * len(sample_values))(*sample_values)
1238
+ status = self._lib.wfloat_stt_session_push_audio(
1239
+ self._session,
1240
+ sample_array,
1241
+ len(sample_values),
1242
+ resolved_sample_rate,
1243
+ )
1244
+ if status != WFLOAT_STATUS_OK:
1245
+ raise RuntimeError(f"wfloat-core session push failed with status {status}.")
1246
+
1247
+ def get_result(self) -> StreamingTranscriptionResult:
1248
+ result_ptr = ctypes.POINTER(_WfloatSttSessionResult)()
1249
+ status = self._lib.wfloat_stt_session_get_result(
1250
+ self._session,
1251
+ ctypes.byref(result_ptr),
1252
+ )
1253
+ if status != WFLOAT_STATUS_OK:
1254
+ raise RuntimeError(
1255
+ f"wfloat-core session get_result failed with status {status}."
1256
+ )
1257
+
1258
+ try:
1259
+ result = result_ptr.contents
1260
+ return StreamingTranscriptionResult(
1261
+ text=_decode(result.text),
1262
+ model_id=_decode(result.model_id) or self._model_id,
1263
+ is_endpoint=bool(result.is_endpoint),
1264
+ json=_decode(result.json),
1265
+ )
1266
+ finally:
1267
+ self._lib.wfloat_stt_session_result_destroy(result_ptr)
1268
+
1269
+ def finish(self) -> StreamingTranscriptionResult:
1270
+ result_ptr = ctypes.POINTER(_WfloatSttSessionResult)()
1271
+ status = self._lib.wfloat_stt_session_finish(
1272
+ self._session,
1273
+ ctypes.byref(result_ptr),
1274
+ )
1275
+ if status != WFLOAT_STATUS_OK:
1276
+ raise RuntimeError(f"wfloat-core session finish failed with status {status}.")
1277
+
1278
+ try:
1279
+ result = result_ptr.contents
1280
+ return StreamingTranscriptionResult(
1281
+ text=_decode(result.text),
1282
+ model_id=_decode(result.model_id) or self._model_id,
1283
+ is_endpoint=bool(result.is_endpoint),
1284
+ json=_decode(result.json),
1285
+ )
1286
+ finally:
1287
+ self._lib.wfloat_stt_session_result_destroy(result_ptr)
1288
+
1289
+ def reset(self) -> None:
1290
+ status = self._lib.wfloat_stt_session_reset(self._session)
1291
+ if status != WFLOAT_STATUS_OK:
1292
+ raise RuntimeError(f"wfloat-core session reset failed with status {status}.")
1293
+
1294
+
1295
+ def create_core_stt_whisper(
1296
+ *,
1297
+ model_name: str,
1298
+ encoder_path: Path,
1299
+ decoder_path: Path,
1300
+ tokens_path: Path,
1301
+ language: Optional[str] = None,
1302
+ task: Optional[str] = None,
1303
+ enable_token_timestamps: bool = False,
1304
+ enable_segment_timestamps: bool = False,
1305
+ ):
1306
+ return CoreStt(
1307
+ model_id=model_name,
1308
+ family=WFLOAT_STT_FAMILY_WHISPER,
1309
+ model_path=None,
1310
+ tokens_path=tokens_path,
1311
+ encoder_path=encoder_path,
1312
+ decoder_path=decoder_path,
1313
+ language=language,
1314
+ task=task,
1315
+ enable_token_timestamps=enable_token_timestamps,
1316
+ enable_segment_timestamps=enable_segment_timestamps,
1317
+ )
1318
+
1319
+
1320
+ def create_core_stt(
1321
+ *,
1322
+ model_name: str,
1323
+ family: str,
1324
+ model_path: Optional[Path],
1325
+ tokens_path: Path,
1326
+ preprocessor_path: Optional[Path] = None,
1327
+ encoder_path: Optional[Path] = None,
1328
+ decoder_path: Optional[Path] = None,
1329
+ joiner_path: Optional[Path] = None,
1330
+ uncached_decoder_path: Optional[Path] = None,
1331
+ cached_decoder_path: Optional[Path] = None,
1332
+ language: Optional[str] = None,
1333
+ task: Optional[str] = None,
1334
+ enable_token_timestamps: bool = False,
1335
+ enable_segment_timestamps: bool = False,
1336
+ ):
1337
+ normalized_family = family.strip().lower().replace("_", "-")
1338
+ family_map = {
1339
+ "whisper": WFLOAT_STT_FAMILY_WHISPER,
1340
+ "moonshine": WFLOAT_STT_FAMILY_MOONSHINE,
1341
+ "parakeet-ctc": WFLOAT_STT_FAMILY_PARAKEET_CTC,
1342
+ "parakeet-tdt": WFLOAT_STT_FAMILY_PARAKEET_TDT,
1343
+ "zipformer-transducer": WFLOAT_STT_FAMILY_ZIPFORMER_TRANSDUCER,
1344
+ }
1345
+ family_value = family_map.get(normalized_family)
1346
+ if family_value is None:
1347
+ raise ValueError(f"Unsupported STT family: {family}")
1348
+
1349
+ return CoreStt(
1350
+ model_id=model_name,
1351
+ family=family_value,
1352
+ model_path=model_path,
1353
+ tokens_path=tokens_path,
1354
+ preprocessor_path=preprocessor_path,
1355
+ encoder_path=encoder_path,
1356
+ decoder_path=decoder_path,
1357
+ joiner_path=joiner_path,
1358
+ uncached_decoder_path=uncached_decoder_path,
1359
+ cached_decoder_path=cached_decoder_path,
1360
+ language=language,
1361
+ task=task,
1362
+ enable_token_timestamps=enable_token_timestamps,
1363
+ enable_segment_timestamps=enable_segment_timestamps,
1364
+ )
1365
+
1366
+
1367
+ class _CoreVadNativeSegment:
1368
+ def __init__(self, *, start: int, samples: Sequence[float]) -> None:
1369
+ self.start = int(start)
1370
+ self.samples = list(samples)
1371
+
1372
+
1373
+ class CoreVad:
1374
+ def __init__(
1375
+ self,
1376
+ *,
1377
+ model_id: str,
1378
+ family: int,
1379
+ model_path: Path,
1380
+ threshold: float,
1381
+ min_silence_duration_sec: float,
1382
+ min_speech_duration_sec: float,
1383
+ max_speech_duration_sec: float,
1384
+ sample_rate: int,
1385
+ window_size: int,
1386
+ buffer_size_in_seconds: float,
1387
+ ) -> None:
1388
+ self._lib = _prepare_library(_load_core_library())
1389
+ self._model = ctypes.c_void_p()
1390
+ self._config_bytes = {
1391
+ "model_id": model_id.encode("utf-8"),
1392
+ "model_path": str(model_path).encode("utf-8"),
1393
+ "provider": DEFAULT_PROVIDER.encode("utf-8"),
1394
+ }
1395
+
1396
+ config = _WfloatVadModelConfig(
1397
+ model_id=self._config_bytes["model_id"],
1398
+ family=family,
1399
+ model_path=self._config_bytes["model_path"],
1400
+ threshold=float(threshold),
1401
+ min_silence_duration_sec=float(min_silence_duration_sec),
1402
+ min_speech_duration_sec=float(min_speech_duration_sec),
1403
+ max_speech_duration_sec=float(max_speech_duration_sec),
1404
+ sample_rate=int(sample_rate),
1405
+ window_size=int(window_size),
1406
+ num_threads=1,
1407
+ provider=self._config_bytes["provider"],
1408
+ debug=0,
1409
+ buffer_size_in_seconds=float(buffer_size_in_seconds),
1410
+ )
1411
+
1412
+ status = self._lib.wfloat_vad_model_create(
1413
+ ctypes.byref(config),
1414
+ ctypes.byref(self._model),
1415
+ )
1416
+ if status != WFLOAT_STATUS_OK:
1417
+ raise RuntimeError(f"wfloat-core VAD model creation failed with status {status}.")
1418
+
1419
+ info = _WfloatVadModelInfo()
1420
+ status = self._lib.wfloat_vad_model_get_info(self._model, ctypes.byref(info))
1421
+ if status != WFLOAT_STATUS_OK:
1422
+ self.close()
1423
+ raise RuntimeError(f"wfloat-core VAD model info failed with status {status}.")
1424
+
1425
+ self.sample_rate = int(info.sample_rate)
1426
+ self.window_size = int(info.window_size)
1427
+
1428
+ def close(self) -> None:
1429
+ if self._model and self._model.value:
1430
+ self._lib.wfloat_vad_model_destroy(self._model)
1431
+ self._model = ctypes.c_void_p()
1432
+
1433
+ def __del__(self) -> None:
1434
+ try:
1435
+ self.close()
1436
+ except Exception:
1437
+ pass
1438
+
1439
+ def reset(self) -> None:
1440
+ status = self._lib.wfloat_vad_model_reset(self._model)
1441
+ if status != WFLOAT_STATUS_OK:
1442
+ raise RuntimeError(f"wfloat-core VAD reset failed with status {status}.")
1443
+
1444
+ def accept_waveform(self, samples: Sequence[float]) -> None:
1445
+ sample_values = [float(sample) for sample in samples]
1446
+ if not sample_values:
1447
+ return
1448
+
1449
+ sample_array = (ctypes.c_float * len(sample_values))(*sample_values)
1450
+ status = self._lib.wfloat_vad_model_accept_waveform(
1451
+ self._model,
1452
+ sample_array,
1453
+ len(sample_values),
1454
+ )
1455
+ if status != WFLOAT_STATUS_OK:
1456
+ raise RuntimeError(
1457
+ f"wfloat-core VAD accept_waveform failed with status {status}."
1458
+ )
1459
+
1460
+ def flush(self) -> None:
1461
+ status = self._lib.wfloat_vad_model_flush(self._model)
1462
+ if status != WFLOAT_STATUS_OK:
1463
+ raise RuntimeError(f"wfloat-core VAD flush failed with status {status}.")
1464
+
1465
+ def empty(self) -> bool:
1466
+ value = ctypes.c_int32()
1467
+ status = self._lib.wfloat_vad_model_empty(self._model, ctypes.byref(value))
1468
+ if status != WFLOAT_STATUS_OK:
1469
+ raise RuntimeError(f"wfloat-core VAD empty failed with status {status}.")
1470
+ return bool(value.value)
1471
+
1472
+ def detected(self) -> bool:
1473
+ value = ctypes.c_int32()
1474
+ status = self._lib.wfloat_vad_model_detected(self._model, ctypes.byref(value))
1475
+ if status != WFLOAT_STATUS_OK:
1476
+ raise RuntimeError(f"wfloat-core VAD detected failed with status {status}.")
1477
+ return bool(value.value)
1478
+
1479
+ @property
1480
+ def front(self) -> _CoreVadNativeSegment:
1481
+ segment_ptr = ctypes.POINTER(_WfloatVadSegment)()
1482
+ status = self._lib.wfloat_vad_model_front(
1483
+ self._model,
1484
+ ctypes.byref(segment_ptr),
1485
+ )
1486
+ if status != WFLOAT_STATUS_OK:
1487
+ raise RuntimeError(f"wfloat-core VAD front failed with status {status}.")
1488
+
1489
+ try:
1490
+ segment = segment_ptr.contents
1491
+ samples = [
1492
+ float(segment.samples[index])
1493
+ for index in range(int(segment.sample_count))
1494
+ ]
1495
+ return _CoreVadNativeSegment(
1496
+ start=int(segment.start_sample),
1497
+ samples=samples,
1498
+ )
1499
+ finally:
1500
+ self._lib.wfloat_vad_segment_destroy(segment_ptr)
1501
+
1502
+ def pop(self) -> None:
1503
+ status = self._lib.wfloat_vad_model_pop(self._model)
1504
+ if status != WFLOAT_STATUS_OK:
1505
+ raise RuntimeError(f"wfloat-core VAD pop failed with status {status}.")
1506
+
1507
+ def clear(self) -> None:
1508
+ status = self._lib.wfloat_vad_model_clear(self._model)
1509
+ if status != WFLOAT_STATUS_OK:
1510
+ raise RuntimeError(f"wfloat-core VAD clear failed with status {status}.")
1511
+
1512
+
1513
+ def create_core_vad(
1514
+ *,
1515
+ model_name: str,
1516
+ family: str,
1517
+ model_path: Path,
1518
+ threshold: float,
1519
+ min_silence_duration_sec: float,
1520
+ min_speech_duration_sec: float,
1521
+ max_speech_duration_sec: float,
1522
+ sample_rate: int,
1523
+ buffer_size_in_seconds: float,
1524
+ ):
1525
+ normalized_family = family.strip().lower().replace("_", "-")
1526
+ family_map = {
1527
+ "silero": WFLOAT_VAD_FAMILY_SILERO,
1528
+ "silero-vad": WFLOAT_VAD_FAMILY_SILERO,
1529
+ "ten-vad": WFLOAT_VAD_FAMILY_TEN,
1530
+ "tenvad": WFLOAT_VAD_FAMILY_TEN,
1531
+ }
1532
+ family_value = family_map.get(normalized_family)
1533
+ if family_value is None:
1534
+ raise ValueError(f"Unsupported VAD family: {family}")
1535
+
1536
+ window_size = 256 if family_value == WFLOAT_VAD_FAMILY_TEN else 512
1537
+ return CoreVad(
1538
+ model_id=model_name,
1539
+ family=family_value,
1540
+ model_path=model_path,
1541
+ threshold=threshold,
1542
+ min_silence_duration_sec=min_silence_duration_sec,
1543
+ min_speech_duration_sec=min_speech_duration_sec,
1544
+ max_speech_duration_sec=max_speech_duration_sec,
1545
+ sample_rate=sample_rate,
1546
+ window_size=window_size,
1547
+ buffer_size_in_seconds=buffer_size_in_seconds,
1548
+ )
1549
+
1550
+
1551
+ class CoreLlm:
1552
+ def __init__(
1553
+ self,
1554
+ *,
1555
+ model_id: str,
1556
+ family: int,
1557
+ model_path: Path,
1558
+ context_size: int = 2048,
1559
+ num_threads: int = 1,
1560
+ gpu_layer_count: int = 0,
1561
+ chat_template: Optional[str] = None,
1562
+ ) -> None:
1563
+ self._lib = _prepare_library(_load_core_library())
1564
+ self._model = ctypes.c_void_p()
1565
+
1566
+ self._config_bytes = {
1567
+ "model_id": model_id.encode("utf-8"),
1568
+ "model_path": str(model_path).encode("utf-8"),
1569
+ "chat_template": None
1570
+ if chat_template is None
1571
+ else chat_template.encode("utf-8"),
1572
+ "provider": b"llama.cpp",
1573
+ }
1574
+
1575
+ config = _WfloatLlmModelConfig(
1576
+ model_id=self._config_bytes["model_id"],
1577
+ family=family,
1578
+ model_path=self._config_bytes["model_path"],
1579
+ chat_template=self._config_bytes["chat_template"],
1580
+ provider=self._config_bytes["provider"],
1581
+ context_size=int(context_size),
1582
+ num_threads=int(num_threads),
1583
+ gpu_layer_count=int(gpu_layer_count),
1584
+ seed=0,
1585
+ )
1586
+
1587
+ status = self._lib.wfloat_llm_model_create(
1588
+ ctypes.byref(config),
1589
+ ctypes.byref(self._model),
1590
+ )
1591
+ if status != WFLOAT_STATUS_OK:
1592
+ raise RuntimeError(f"wfloat-core LLM model creation failed with status {status}.")
1593
+
1594
+ info = _WfloatLlmModelInfo()
1595
+ status = self._lib.wfloat_llm_model_get_info(self._model, ctypes.byref(info))
1596
+ if status != WFLOAT_STATUS_OK:
1597
+ self.close()
1598
+ raise RuntimeError(f"wfloat-core LLM model info failed with status {status}.")
1599
+
1600
+ self.model_id = _decode(info.model_id) or model_id
1601
+ self.backend = _decode(info.backend)
1602
+ self.family = _decode(info.family)
1603
+ self.context_size = int(info.context_size)
1604
+
1605
+ def close(self) -> None:
1606
+ if self._model and self._model.value:
1607
+ self._lib.wfloat_llm_model_destroy(self._model)
1608
+ self._model = ctypes.c_void_p()
1609
+
1610
+ def __del__(self) -> None:
1611
+ try:
1612
+ self.close()
1613
+ except Exception:
1614
+ pass
1615
+
1616
+ def generate(
1617
+ self,
1618
+ prompt: str,
1619
+ *,
1620
+ max_tokens: int = 128,
1621
+ temperature: float = 0.8,
1622
+ top_p: float = 0.95,
1623
+ top_k: int = 40,
1624
+ repeat_penalty: float = 1.0,
1625
+ seed: int = 0,
1626
+ on_token=None,
1627
+ ) -> LlmGenerationResult:
1628
+ prompt_bytes = prompt.encode("utf-8")
1629
+ options = _WfloatLlmGenerateOptions(
1630
+ prompt=prompt_bytes,
1631
+ max_tokens=int(max_tokens),
1632
+ temperature=float(temperature),
1633
+ top_p=float(top_p),
1634
+ top_k=int(top_k),
1635
+ repeat_penalty=float(repeat_penalty),
1636
+ seed=int(seed),
1637
+ )
1638
+
1639
+ callback_ref = _WfloatLlmTokenCallback(
1640
+ lambda event, _user_data: self._handle_token(event, on_token)
1641
+ )
1642
+ result_ptr = ctypes.POINTER(_WfloatLlmGenerateResult)()
1643
+ status = self._lib.wfloat_llm_model_generate(
1644
+ self._model,
1645
+ ctypes.byref(options),
1646
+ callback_ref,
1647
+ None,
1648
+ ctypes.byref(result_ptr),
1649
+ )
1650
+ if status != WFLOAT_STATUS_OK:
1651
+ raise RuntimeError(f"wfloat-core LLM generate failed with status {status}.")
1652
+
1653
+ try:
1654
+ result = result_ptr.contents
1655
+ return LlmGenerationResult(
1656
+ text=_decode(result.text),
1657
+ model_id=_decode(result.model_id) or self.model_id,
1658
+ finish_reason=_decode(result.finish_reason),
1659
+ json=_decode(result.json),
1660
+ prompt_token_count=int(result.prompt_token_count),
1661
+ completion_token_count=int(result.completion_token_count),
1662
+ )
1663
+ finally:
1664
+ self._lib.wfloat_llm_generate_result_destroy(result_ptr)
1665
+
1666
+ def format_chat(
1667
+ self,
1668
+ messages: Sequence[Dict[str, str]],
1669
+ *,
1670
+ add_generation_prompt: bool = True,
1671
+ ) -> str:
1672
+ message_structs = []
1673
+ message_buffers = []
1674
+ for message in messages:
1675
+ role = str(message["role"])
1676
+ content = str(message["content"])
1677
+ role_bytes = role.encode("utf-8")
1678
+ content_bytes = content.encode("utf-8")
1679
+ message_buffers.append((role_bytes, content_bytes))
1680
+ message_structs.append(
1681
+ _WfloatLlmChatMessage(
1682
+ role=role_bytes,
1683
+ content=content_bytes,
1684
+ )
1685
+ )
1686
+
1687
+ if not message_structs:
1688
+ raise ValueError("LLM chat messages cannot be empty.")
1689
+
1690
+ message_array = (_WfloatLlmChatMessage * len(message_structs))(
1691
+ *message_structs
1692
+ )
1693
+ options = _WfloatLlmChatTemplateOptions(
1694
+ messages=message_array,
1695
+ message_count=len(message_structs),
1696
+ add_generation_prompt=1 if add_generation_prompt else 0,
1697
+ )
1698
+ result_ptr = ctypes.POINTER(_WfloatLlmChatTemplateResult)()
1699
+ status = self._lib.wfloat_llm_model_format_chat(
1700
+ self._model,
1701
+ ctypes.byref(options),
1702
+ ctypes.byref(result_ptr),
1703
+ )
1704
+ if status != WFLOAT_STATUS_OK:
1705
+ raise RuntimeError(f"wfloat-core LLM format_chat failed with status {status}.")
1706
+
1707
+ try:
1708
+ result = result_ptr.contents
1709
+ if result.used_fallback:
1710
+ warnings.warn(
1711
+ "wfloat-core could not apply this GGUF model's chat "
1712
+ "template with llama.cpp, so it used a generic fallback "
1713
+ "prompt format. Output quality may be degraded.",
1714
+ RuntimeWarning,
1715
+ stacklevel=2,
1716
+ )
1717
+ return _decode(result.prompt)
1718
+ finally:
1719
+ self._lib.wfloat_llm_chat_template_result_destroy(result_ptr)
1720
+
1721
+ def chat(
1722
+ self,
1723
+ messages: Sequence[Dict[str, str]],
1724
+ *,
1725
+ max_tokens: int = 128,
1726
+ temperature: float = 0.8,
1727
+ top_p: float = 0.95,
1728
+ top_k: int = 40,
1729
+ repeat_penalty: float = 1.0,
1730
+ seed: int = 0,
1731
+ on_token=None,
1732
+ ) -> LlmGenerationResult:
1733
+ prompt = self.format_chat(messages, add_generation_prompt=True)
1734
+ return self.generate(
1735
+ prompt,
1736
+ max_tokens=max_tokens,
1737
+ temperature=temperature,
1738
+ top_p=top_p,
1739
+ top_k=top_k,
1740
+ repeat_penalty=repeat_penalty,
1741
+ seed=seed,
1742
+ on_token=on_token,
1743
+ )
1744
+
1745
+ @staticmethod
1746
+ def _handle_token(event_ptr, on_token) -> int:
1747
+ if on_token is None:
1748
+ return 0
1749
+
1750
+ event = event_ptr.contents
1751
+ if event.is_done:
1752
+ return 0
1753
+
1754
+ on_token(_decode(event.text))
1755
+ return 0
1756
+
1757
+
1758
+ def create_core_llm(
1759
+ *,
1760
+ model_name: str,
1761
+ family: str,
1762
+ model_path: Path,
1763
+ context_size: int = 2048,
1764
+ num_threads: int = 1,
1765
+ gpu_layer_count: int = 0,
1766
+ chat_template: Optional[str] = None,
1767
+ ):
1768
+ normalized_family = family.strip().lower().replace("_", "-")
1769
+ family_map = {
1770
+ "llama": WFLOAT_LLM_FAMILY_LLAMA,
1771
+ "qwen": WFLOAT_LLM_FAMILY_QWEN,
1772
+ "qwen2": WFLOAT_LLM_FAMILY_QWEN,
1773
+ "qwen3": WFLOAT_LLM_FAMILY_QWEN,
1774
+ "smollm": WFLOAT_LLM_FAMILY_SMOLLM,
1775
+ "smollm2": WFLOAT_LLM_FAMILY_SMOLLM,
1776
+ "gemma": WFLOAT_LLM_FAMILY_GEMMA,
1777
+ "mistral": WFLOAT_LLM_FAMILY_MISTRAL,
1778
+ "phi": WFLOAT_LLM_FAMILY_PHI,
1779
+ "liquid": WFLOAT_LLM_FAMILY_LIQUID,
1780
+ "lfm": WFLOAT_LLM_FAMILY_LIQUID,
1781
+ "lfm2": WFLOAT_LLM_FAMILY_LIQUID,
1782
+ }
1783
+ family_value = family_map.get(normalized_family)
1784
+ if family_value is None:
1785
+ raise ValueError(f"Unsupported LLM family: {family}")
1786
+
1787
+ return CoreLlm(
1788
+ model_id=model_name,
1789
+ family=family_value,
1790
+ model_path=model_path,
1791
+ context_size=context_size,
1792
+ num_threads=num_threads,
1793
+ gpu_layer_count=gpu_layer_count,
1794
+ chat_template=chat_template,
1795
+ )