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
21 changes: 20 additions & 1 deletion bindings/python/src/transcribe_cpp/_generated.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
# Stable digest of the ABI surface below (structs, enums, macros, layout,
# prototypes). A native provider package echoes this back so the API
# package can reject an ABI-mismatched provider before dlopen.
PUBLIC_HEADER_HASH = "59b9a92b47074666"
PUBLIC_HEADER_HASH = "57b1af43650f195d"

# === enum constants ===
TRANSCRIBE_OK = 0
Expand Down Expand Up @@ -51,6 +51,7 @@
TRANSCRIBE_ABI_EXT = 12
TRANSCRIBE_ABI_DEVICE_INFO = 13
TRANSCRIBE_ABI_SPEAKER_SEGMENT = 14
TRANSCRIBE_ABI_BACKEND_INIT_PARAMS = 15
TRANSCRIBE_LOG_LEVEL_NONE = 0
TRANSCRIBE_LOG_LEVEL_INFO = 1
TRANSCRIBE_LOG_LEVEL_WARN = 2
Expand Down Expand Up @@ -116,6 +117,13 @@
TRANSCRIBE_WHISPER_PROMPT_ALL_SEGMENTS = 1

# === macro constants (integer object-like macros) ===
TRANSCRIBE_BACKEND_MASK_ALL = 4294967295
TRANSCRIBE_BACKEND_MASK_CPU = 1
TRANSCRIBE_BACKEND_MASK_CUDA = 8
TRANSCRIBE_BACKEND_MASK_METAL = 2
TRANSCRIBE_BACKEND_MASK_OTHER = 2147483648
TRANSCRIBE_BACKEND_MASK_ROCM = 16
TRANSCRIBE_BACKEND_MASK_VULKAN = 4
TRANSCRIBE_EXT_KIND_MOONSHINE_STREAMING_STREAM = 1414746957
TRANSCRIBE_EXT_KIND_PARAKEET_BUFFERED_STREAM = 1396853584
TRANSCRIBE_EXT_KIND_PARAKEET_STREAM = 1414744912
Expand All @@ -126,6 +134,8 @@
# === structs ===
class transcribe_ext(_c.Structure):
pass
class transcribe_backend_init_params(_c.Structure):
pass
class transcribe_device_info(_c.Structure):
pass
class transcribe_model_load_params(_c.Structure):
Expand Down Expand Up @@ -170,6 +180,7 @@ class transcribe_whisper_chunk_trace(_c.Structure):
pass

transcribe_ext._fields_ = [("size", _c.c_uint64), ("kind", _c.c_uint32)]
transcribe_backend_init_params._fields_ = [("struct_size", _c.c_uint64), ("artifact_dir", _c.c_char_p), ("allowed_backends", _c.c_uint32)]
transcribe_device_info._fields_ = [("struct_size", _c.c_uint64), ("name", _c.c_char_p), ("description", _c.c_char_p), ("kind", _c.c_char_p), ("device_id", _c.c_char_p), ("memory_total", _c.c_uint64), ("memory_free", _c.c_uint64), ("device_type", _c.c_int)]
transcribe_model_load_params._fields_ = [("struct_size", _c.c_uint64), ("backend", _c.c_int), ("device", _c.c_void_p)]
transcribe_session_params._fields_ = [("struct_size", _c.c_uint64), ("n_threads", _c.c_int), ("kv_type", _c.c_int), ("n_ctx", _c.c_int32)]
Expand All @@ -196,6 +207,7 @@ class transcribe_whisper_chunk_trace(_c.Structure):
# transcribe_abi_struct id per struct (for the native size/align check).
ABI_STRUCT_IDS = {
'transcribe_ext': 12,
'transcribe_backend_init_params': 15,
'transcribe_device_info': 13,
'transcribe_model_load_params': 0,
'transcribe_session_params': 1,
Expand All @@ -215,6 +227,7 @@ class transcribe_whisper_chunk_trace(_c.Structure):
# C-compiler layout captured at generation (for offset self-check).
STRUCT_LAYOUT = {
'transcribe_ext': {'size': 16, 'align': 8, 'offsets': {'size': 0, 'kind': 8}},
'transcribe_backend_init_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'artifact_dir': 8, 'allowed_backends': 16}},
'transcribe_device_info': {'size': 64, 'align': 8, 'offsets': {'struct_size': 0, 'name': 8, 'description': 16, 'kind': 24, 'device_id': 32, 'memory_total': 40, 'memory_free': 48, 'device_type': 56}},
'transcribe_model_load_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'backend': 8, 'device': 16}},
'transcribe_session_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'n_threads': 8, 'kv_type': 12, 'n_ctx': 16}},
Expand Down Expand Up @@ -245,8 +258,12 @@ def configure(lib):
lib.transcribe_abi_struct_align.argtypes = [_c.c_int]
lib.transcribe_abi_struct_size.restype = _c.c_size_t
lib.transcribe_abi_struct_size.argtypes = [_c.c_int]
lib.transcribe_allowed_backends.restype = _c.c_uint32
lib.transcribe_allowed_backends.argtypes = []
lib.transcribe_backend_available.restype = _c.c_bool
lib.transcribe_backend_available.argtypes = [_c.c_int]
lib.transcribe_backend_init_params_init.restype = None
lib.transcribe_backend_init_params_init.argtypes = [_c.POINTER(transcribe_backend_init_params)]
lib.transcribe_batch_detected_language.restype = _c.c_char_p
lib.transcribe_batch_detected_language.argtypes = [_c.c_void_p, _c.c_int]
lib.transcribe_batch_full_text.restype = _c.c_char_p
Expand Down Expand Up @@ -315,6 +332,8 @@ def configure(lib):
lib.transcribe_init_backends.argtypes = [_c.c_char_p]
lib.transcribe_init_backends_default.restype = _c.c_int
lib.transcribe_init_backends_default.argtypes = []
lib.transcribe_init_backends_ex.restype = _c.c_int
lib.transcribe_init_backends_ex.argtypes = [_c.POINTER(transcribe_backend_init_params)]
lib.transcribe_log_set.restype = None
lib.transcribe_log_set.argtypes = [_c.CFUNCTYPE(None, _c.c_int, _c.c_char_p, _c.c_void_p), _c.c_void_p]
lib.transcribe_model_accepts_ext_kind.restype = _c.c_bool
Expand Down
43 changes: 41 additions & 2 deletions bindings/rust/sys/src/transcribe_sys.rs
Original file line number Diff line number Diff line change
@@ -1,14 +1,21 @@
// @generated by `cargo xtask bindgen` from include/transcribe/extensions.h
// DO NOT EDIT BY HAND. Regenerate: `cargo xtask bindgen`.
// Pinned to include/transcribe.abihash = 59b9a92b47074666
// Pinned to include/transcribe.abihash = 57b1af43650f195d

/// The public-ABI digest these bindings were generated against
/// (sha256/16 over the normalized FFI surface). The load-time version
/// gate and the CI drift check both anchor on this value.
pub const PUBLIC_HEADER_HASH: &str = "59b9a92b47074666";
pub const PUBLIC_HEADER_HASH: &str = "57b1af43650f195d";

/* automatically generated by rust-bindgen 0.72.1 */

pub const TRANSCRIBE_BACKEND_MASK_CPU: u32 = 1;
pub const TRANSCRIBE_BACKEND_MASK_METAL: u32 = 2;
pub const TRANSCRIBE_BACKEND_MASK_VULKAN: u32 = 4;
pub const TRANSCRIBE_BACKEND_MASK_CUDA: u32 = 8;
pub const TRANSCRIBE_BACKEND_MASK_ROCM: u32 = 16;
pub const TRANSCRIBE_BACKEND_MASK_OTHER: u32 = 2147483648;
pub const TRANSCRIBE_BACKEND_MASK_ALL: u32 = 4294967295;
pub const TRANSCRIBE_EXT_KIND_MOONSHINE_STREAMING_STREAM: u32 = 1414746957;
pub const TRANSCRIBE_EXT_KIND_PARAKEET_STREAM: u32 = 1414744912;
pub const TRANSCRIBE_EXT_KIND_PARAKEET_BUFFERED_STREAM: u32 = 1396853584;
Expand Down Expand Up @@ -66,6 +73,7 @@ impl transcribe_abi_struct {
pub const TRANSCRIBE_ABI_EXT: transcribe_abi_struct = transcribe_abi_struct(12);
pub const TRANSCRIBE_ABI_DEVICE_INFO: transcribe_abi_struct = transcribe_abi_struct(13);
pub const TRANSCRIBE_ABI_SPEAKER_SEGMENT: transcribe_abi_struct = transcribe_abi_struct(14);
pub const TRANSCRIBE_ABI_BACKEND_INIT_PARAMS: transcribe_abi_struct = transcribe_abi_struct(15);
}
#[repr(transparent)]
#[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)]
Expand Down Expand Up @@ -217,6 +225,37 @@ unsafe extern "C" {
}
#[repr(C)]
#[derive(Debug, Copy, Clone)]
pub struct transcribe_backend_init_params {
pub struct_size: u64,
pub artifact_dir: *const ::std::os::raw::c_char,
pub allowed_backends: u32,
}
#[allow(clippy::unnecessary_operation, clippy::identity_op)]
const _: () = {
["Size of transcribe_backend_init_params"]
[::std::mem::size_of::<transcribe_backend_init_params>() - 24usize];
["Alignment of transcribe_backend_init_params"]
[::std::mem::align_of::<transcribe_backend_init_params>() - 8usize];
["Offset of field: transcribe_backend_init_params::struct_size"]
[::std::mem::offset_of!(transcribe_backend_init_params, struct_size) - 0usize];
["Offset of field: transcribe_backend_init_params::artifact_dir"]
[::std::mem::offset_of!(transcribe_backend_init_params, artifact_dir) - 8usize];
["Offset of field: transcribe_backend_init_params::allowed_backends"]
[::std::mem::offset_of!(transcribe_backend_init_params, allowed_backends) - 16usize];
};
unsafe extern "C" {
pub fn transcribe_backend_init_params_init(p: *mut transcribe_backend_init_params);
}
unsafe extern "C" {
pub fn transcribe_init_backends_ex(
params: *const transcribe_backend_init_params,
) -> transcribe_status;
}
unsafe extern "C" {
pub fn transcribe_allowed_backends() -> u32;
}
#[repr(C)]
#[derive(Debug, Copy, Clone)]
pub struct transcribe_device {
_unused: [u8; 0],
}
Expand Down
107 changes: 107 additions & 0 deletions bindings/rust/transcribe-cpp/src/backend.rs
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,113 @@ pub fn init_backends_default() -> Result<()> {
check(status, "init_backends_default")
}

/// Backend kinds allowed to register in this process. A backend outside the
/// mask never runs any code, so a broken GPU driver can be kept out of a
/// worker entirely. CPU is always allowed; `TRANSCRIBE_BACKENDS` can only
/// narrow the mask.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct BackendMask(u32);

impl BackendMask {
/// CPU plus host-memory accelerators (BLAS). Always implied.
pub const CPU: BackendMask = BackendMask(sys::TRANSCRIBE_BACKEND_MASK_CPU);
/// Apple Metal.
pub const METAL: BackendMask = BackendMask(sys::TRANSCRIBE_BACKEND_MASK_METAL);
/// Vulkan.
pub const VULKAN: BackendMask = BackendMask(sys::TRANSCRIBE_BACKEND_MASK_VULKAN);
/// NVIDIA CUDA.
pub const CUDA: BackendMask = BackendMask(sys::TRANSCRIBE_BACKEND_MASK_CUDA);
/// AMD ROCm / HIP.
pub const ROCM: BackendMask = BackendMask(sys::TRANSCRIBE_BACKEND_MASK_ROCM);
/// Every backend without a dedicated bit (SYCL, OpenCL, RPC, ...).
pub const OTHER: BackendMask = BackendMask(sys::TRANSCRIBE_BACKEND_MASK_OTHER);
/// Everything. The default.
pub const ALL: BackendMask = BackendMask(sys::TRANSCRIBE_BACKEND_MASK_ALL);

/// The smallest mask that can serve a model-load `backend` request
/// ([`Backend::Auto`] needs everything).
pub const fn for_backend(backend: Backend) -> BackendMask {
match backend {
Backend::Auto => BackendMask::ALL,
Backend::Cpu | Backend::CpuAccel => BackendMask::CPU,
Backend::Metal => BackendMask::METAL.union(BackendMask::CPU),
Backend::Vulkan => BackendMask::VULKAN.union(BackendMask::CPU),
Backend::Cuda => BackendMask::CUDA.union(BackendMask::CPU),
Backend::Rocm => BackendMask::ROCM.union(BackendMask::CPU),
}
}

/// The raw `TRANSCRIBE_BACKEND_MASK_*` bits.
pub const fn bits(self) -> u32 {
self.0
}

/// A mask from raw `TRANSCRIBE_BACKEND_MASK_*` bits.
pub const fn from_bits(bits: u32) -> BackendMask {
BackendMask(bits)
}

/// Both masks' backends.
pub const fn union(self, other: BackendMask) -> BackendMask {
BackendMask(self.0 | other.0)
}

/// Whether every backend in `other` is in `self`.
pub const fn contains(self, other: BackendMask) -> bool {
self.0 & other.0 == other.0
}
}

impl Default for BackendMask {
fn default() -> Self {
BackendMask::ALL
}
}

impl std::ops::BitOr for BackendMask {
type Output = BackendMask;
fn bitor(self, rhs: BackendMask) -> BackendMask {
self.union(rhs)
}
}

impl std::ops::BitOrAssign for BackendMask {
fn bitor_assign(&mut self, rhs: BackendMask) {
*self = self.union(rhs);
}
}

/// [`init_backends`] / [`init_backends_default`] (`dir` = `None`) with an
/// allowed-backend mask. The mask is fixed at first backend registration;
/// call this first, once per process. A later call with a different
/// effective mask returns [`crate::Error::Backend`].
///
/// ```no_run
/// use transcribe_cpp::{init_backends_with, Backend, BackendMask};
/// // A CPU fallback worker: never let the GPU driver load.
/// init_backends_with(None::<&std::path::Path>, BackendMask::for_backend(Backend::Cpu))?;
/// # Ok::<(), transcribe_cpp::Error>(())
/// ```
pub fn init_backends_with(dir: Option<impl AsRef<Path>>, allowed: BackendMask) -> Result<()> {
let c_dir = match dir {
Some(dir) => Some(CString::new(crate::model::path_bytes(dir.as_ref())?)?),
None => None,
};
let mut params: sys::transcribe_backend_init_params = unsafe { std::mem::zeroed() };
unsafe { sys::transcribe_backend_init_params_init(&mut params) };
params.artifact_dir = c_dir.as_ref().map_or(std::ptr::null(), |d| d.as_ptr());
params.allowed_backends = allowed.bits();
let status = unsafe { sys::transcribe_init_backends_ex(&params) };
check(status, "init_backends_with")
}

/// The effective mask: [`init_backends_with`]'s (ALL until then) narrowed
/// by `TRANSCRIBE_BACKENDS`.
pub fn allowed_backends() -> BackendMask {
BackendMask(unsafe { sys::transcribe_allowed_backends() })
}

/// The number of compute devices currently registered.
///
/// Do not race this query with [`init_backends`] or [`init_backends_default`].
Expand Down
4 changes: 2 additions & 2 deletions bindings/rust/transcribe-cpp/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -58,8 +58,8 @@ mod types;
mod version;

pub use backend::{
backend_available, device_count, devices, init_backends, init_backends_default, Device,
DeviceType,
allowed_backends, backend_available, device_count, devices, init_backends,
init_backends_default, init_backends_with, BackendMask, Device, DeviceType,
};
pub use cancel::CancelToken;
pub use error::{Error, Result};
Expand Down
2 changes: 2 additions & 0 deletions bindings/rust/transcribe-cpp/src/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -319,6 +319,7 @@ pub enum AbiStruct {
Ext,
DeviceInfo,
SpeakerSegment,
BackendInitParams,
}

impl AbiStruct {
Expand All @@ -340,6 +341,7 @@ impl AbiStruct {
AbiStruct::Ext => A::TRANSCRIBE_ABI_EXT,
AbiStruct::DeviceInfo => A::TRANSCRIBE_ABI_DEVICE_INFO,
AbiStruct::SpeakerSegment => A::TRANSCRIBE_ABI_SPEAKER_SEGMENT,
AbiStruct::BackendInitParams => A::TRANSCRIBE_ABI_BACKEND_INIT_PARAMS,
}
}
}
Expand Down
47 changes: 47 additions & 0 deletions bindings/rust/transcribe-cpp/tests/backend_mask.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
//! Allowed-backend mask. The mask is fixed per process, so every assertion
//! that touches the native registry lives in ONE test function in this file
//! (its own test binary, hence its own process).

use transcribe_cpp::{
allowed_backends, backend_available, devices, init_backends_with, Backend, BackendMask,
DeviceType, Error,
};

#[test]
fn for_backend_is_minimal() {
assert_eq!(BackendMask::for_backend(Backend::Auto), BackendMask::ALL);
assert_eq!(BackendMask::for_backend(Backend::Cpu), BackendMask::CPU);
assert_eq!(
BackendMask::for_backend(Backend::CpuAccel),
BackendMask::CPU
);
assert_eq!(
BackendMask::for_backend(Backend::Vulkan),
BackendMask::VULKAN | BackendMask::CPU
);
assert!(BackendMask::ALL.contains(BackendMask::METAL | BackendMask::OTHER));
assert!(!BackendMask::CPU.contains(BackendMask::VULKAN));
}

#[test]
fn cpu_only_worker() {
let cpu = BackendMask::for_backend(Backend::Cpu);
init_backends_with(None::<&std::path::Path>, cpu).expect("cpu-only init");
assert_eq!(allowed_backends(), BackendMask::CPU);

let devs = devices();
assert!(!devs.is_empty());
assert!(devs
.iter()
.all(|d| matches!(d.device_type, DeviceType::Cpu | DeviceType::Accel)));
assert!(backend_available(Backend::Cpu));
assert!(!backend_available(Backend::Vulkan));
assert!(!backend_available(Backend::Metal));

// Fixed for the process: the same mask is fine, a different one is not.
init_backends_with(None::<&std::path::Path>, cpu).expect("same mask again");
assert!(matches!(
init_backends_with(None::<&std::path::Path>, BackendMask::ALL),
Err(Error::Backend { .. })
));
}
1 change: 1 addition & 0 deletions bindings/rust/transcribe-cpp/tests/no_model.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ fn abi_struct_sizes_are_live() {
AbiStruct::Segment,
AbiStruct::SpeakerSegment,
AbiStruct::SessionLimits,
AbiStruct::BackendInitParams,
] {
assert!(abi_struct_size(which) > 0, "{which:?} reported size 0");
}
Expand Down
2 changes: 1 addition & 1 deletion bindings/swift/Sources/TranscribeCpp/ABIHash.swift
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ import CTranscribe
extension Transcribe {
/// sha256/16 of the normalized public FFI surface, pinned to the value in
/// include/transcribe.abihash at the time this binding was last reviewed.
public static let pinnedHeaderHash = "59b9a92b47074666"
public static let pinnedHeaderHash = "57b1af43650f195d"

/// The public-ABI digest this binding was reviewed against (16 hex chars).
public static func headerHash() -> String { pinnedHeaderHash }
Expand Down
Loading
Loading