use super::types::{
copy_bytes_into_tape, rust_error_from_c_message, SzSequenceFromBytes, SzSequenceU32Tape, SzSequenceU64Tape,
};
use super::*;
use alloc::vec::Vec;
use core::ffi::{c_char, c_void};
use core::ptr;
pub type FingerprintsHandle = *mut c_void;
extern "C" {
fn szs_fingerprints_init(
dimensions: usize,
alphabet_size: usize,
window_widths: *const usize,
window_widths_count: usize,
seed: u64,
alloc: *const c_void, capabilities: Capability,
engine: *mut FingerprintsHandle,
error_message: *mut *const c_char,
) -> Status;
fn szs_fingerprints_sequence(
engine: FingerprintsHandle,
device: *mut c_void, texts: *const c_void, min_hashes: *mut u32,
min_hashes_stride: usize,
min_counts: *mut u32,
min_counts_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_fingerprints_u32tape(
engine: FingerprintsHandle,
device: *mut c_void, texts: *const c_void, min_hashes: *mut u32,
min_hashes_stride: usize,
min_counts: *mut u32,
min_counts_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_fingerprints_u64tape(
engine: FingerprintsHandle,
device: *mut c_void, texts: *const c_void, min_hashes: *mut u32,
min_hashes_stride: usize,
min_counts: *mut u32,
min_counts_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_fingerprints_free(engine: FingerprintsHandle);
}
pub struct FingerprintsBuilder {
alphabet_size: usize,
window_widths: Option<Vec<usize>>,
dimensions: usize,
seed: u64,
}
impl FingerprintsBuilder {
pub fn new() -> Self {
Self {
alphabet_size: 0,
window_widths: None,
dimensions: 1024, seed: 0, }
}
pub fn binary(mut self) -> Self {
self.alphabet_size = 256;
self
}
pub fn ascii(mut self) -> Self {
self.alphabet_size = 128;
self
}
pub fn dna(mut self) -> Self {
self.alphabet_size = 4;
self
}
pub fn protein(mut self) -> Self {
self.alphabet_size = 22;
self
}
pub fn alphabet_size(mut self, size: usize) -> Self {
self.alphabet_size = size;
self
}
pub fn window_widths(mut self, widths: &[usize]) -> Self {
self.window_widths = Some(widths.to_vec());
self
}
pub fn dimensions(mut self, dimensions: usize) -> Self {
self.dimensions = dimensions;
self
}
pub fn seed(mut self, seed: u64) -> Self {
self.seed = seed;
self
}
pub fn build(self, device: &DeviceScope) -> Result<Fingerprints, Error> {
let mut engine: FingerprintsHandle = ptr::null_mut();
let capabilities = device.get_capabilities().unwrap_or(0);
let (widths_ptr, widths_len) = match &self.window_widths {
Some(widths) => (widths.as_ptr(), widths.len()),
None => (ptr::null(), 0),
};
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_fingerprints_init(
self.dimensions,
self.alphabet_size,
widths_ptr,
widths_len,
self.seed,
ptr::null(), capabilities,
&mut engine,
&mut error_msg,
)
};
match status {
Status::Success => Ok(Fingerprints { handle: engine }),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
}
pub struct Fingerprints {
handle: FingerprintsHandle,
}
impl Fingerprints {
pub fn builder() -> FingerprintsBuilder {
FingerprintsBuilder::new()
}
pub fn compute<Sequences, Sequence>(
&self,
device: &DeviceScope,
strings: Sequences,
dimensions: usize,
) -> Result<(UnifiedVec<u32>, UnifiedVec<u32>), Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<[u8]>,
{
let strings_slice = strings.as_ref();
let num_strings = strings_slice.len();
let hashes_size = num_strings * dimensions;
let counts_size = num_strings * dimensions;
let mut min_hashes = UnifiedVec::with_capacity_in(hashes_size, UnifiedAlloc);
min_hashes.resize(hashes_size, 0);
let mut min_counts = UnifiedVec::with_capacity_in(counts_size, UnifiedAlloc);
min_counts.resize(counts_size, 0);
let hashes_stride = dimensions * core::mem::size_of::<u32>();
let counts_stride = dimensions * core::mem::size_of::<u32>();
if device.is_gpu() {
let total_size: usize = strings_slice.iter().map(|s| s.as_ref().len()).sum();
let force_64bit = total_size > u32::MAX as usize || strings_slice.len() > u32::MAX as usize;
let tape = copy_bytes_into_tape(strings_slice, force_64bit)?;
self.compute_into(device, tape, dimensions, &mut min_hashes[..], &mut min_counts[..])?;
Ok((min_hashes, min_counts))
} else {
let sequence = SzSequenceFromBytes::to_sz_sequence(strings_slice);
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_fingerprints_sequence(
self.handle,
device.handle,
&sequence as *const _ as *const c_void,
min_hashes.as_mut_ptr(),
hashes_stride,
min_counts.as_mut_ptr(),
counts_stride,
&mut error_msg,
)
};
match status {
Status::Success => Ok((min_hashes, min_counts)),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
}
pub fn compute_into<'a>(
&self,
device: &DeviceScope,
texts: AnyBytesTape<'a>,
dimensions: usize,
min_hashes: &mut [u32],
min_counts: &mut [u32],
) -> Result<(), Error> {
let mut error_msg: *const c_char = ptr::null();
let count = match &texts {
AnyBytesTape::Tape64(t) => SzSequenceU64Tape::from(t).count,
AnyBytesTape::View64(v) => SzSequenceU64Tape::from(v).count,
AnyBytesTape::Tape32(t) => SzSequenceU32Tape::from(t).count,
AnyBytesTape::View32(v) => SzSequenceU32Tape::from(v).count,
};
let need = count * dimensions;
if min_hashes.len() < need || min_counts.len() < need {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let hashes_stride = dimensions * core::mem::size_of::<u32>();
let counts_stride = dimensions * core::mem::size_of::<u32>();
let status = match &texts {
AnyBytesTape::Tape64(t) => {
let v = SzSequenceU64Tape::from(t);
unsafe {
szs_fingerprints_u64tape(
self.handle,
device.handle,
&v as *const _ as *const c_void,
min_hashes.as_mut_ptr(),
hashes_stride,
min_counts.as_mut_ptr(),
counts_stride,
&mut error_msg,
)
}
}
AnyBytesTape::View64(vv) => {
let v = SzSequenceU64Tape::from(vv);
unsafe {
szs_fingerprints_u64tape(
self.handle,
device.handle,
&v as *const _ as *const c_void,
min_hashes.as_mut_ptr(),
hashes_stride,
min_counts.as_mut_ptr(),
counts_stride,
&mut error_msg,
)
}
}
AnyBytesTape::Tape32(t) => {
let v = SzSequenceU32Tape::from(t);
unsafe {
szs_fingerprints_u32tape(
self.handle,
device.handle,
&v as *const _ as *const c_void,
min_hashes.as_mut_ptr(),
hashes_stride,
min_counts.as_mut_ptr(),
counts_stride,
&mut error_msg,
)
}
}
AnyBytesTape::View32(vv) => {
let v = SzSequenceU32Tape::from(vv);
unsafe {
szs_fingerprints_u32tape(
self.handle,
device.handle,
&v as *const _ as *const c_void,
min_hashes.as_mut_ptr(),
hashes_stride,
min_counts.as_mut_ptr(),
counts_stride,
&mut error_msg,
)
}
}
};
match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
}
impl Drop for Fingerprints {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe { szs_fingerprints_free(self.handle) };
}
}
}
unsafe impl Send for Fingerprints {}
unsafe impl Sync for Fingerprints {}
#[cfg(test)]
mod tests {
use super::*;
use crate::stringzillas::fixtures::device_or_skip;
const TEST_FINGERPRINT_DIMS_SMALL: usize = 64;
const TEST_FINGERPRINT_DIMS_LARGE: usize = 128;
const TEST_LARGE_BATCH_SIZE: usize = 1000;
#[test]
fn fingerprint_builder_configurations() {
let Some(device) = device_or_skip("fingerprint_builder_configurations") else {
return;
};
let default_engine = Fingerprints::builder().build(&device);
assert!(default_engine.is_ok(), "Default fingerprint engine should initialize");
let binary_engine = Fingerprints::builder().binary().dimensions(256).build(&device);
assert!(binary_engine.is_ok(), "Binary fingerprint engine should initialize");
let ascii_engine = Fingerprints::builder().ascii().dimensions(256).build(&device);
assert!(ascii_engine.is_ok(), "ASCII fingerprint engine should initialize");
let dna_engine = Fingerprints::builder()
.dna()
.window_widths(&[3, 5, 7])
.dimensions(192) .build(&device);
assert!(dna_engine.is_ok(), "DNA fingerprint engine should initialize");
let protein_engine = Fingerprints::builder()
.protein()
.window_widths(&[5, 7])
.dimensions(128) .build(&device);
assert!(protein_engine.is_ok(), "Protein fingerprint engine should initialize");
let custom_engine = Fingerprints::builder()
.alphabet_size(16) .window_widths(&[4, 6, 8])
.dimensions(192) .build(&device);
assert!(custom_engine.is_ok(), "Custom fingerprint engine should initialize");
}
#[test]
fn fingerprint_computation() {
let Some(device) = device_or_skip("fingerprint_computation") else {
return;
};
let engine = Fingerprints::builder()
.binary()
.dimensions(TEST_FINGERPRINT_DIMS_SMALL) .build(&device)
.expect("binary fingerprint engine should build on CPU");
let test_strings = vec!["hello", "world", "test"];
let (hashes, counts) = engine
.compute(&device, &test_strings, TEST_FINGERPRINT_DIMS_SMALL)
.expect("fingerprint computation should succeed on CPU");
assert_eq!(hashes.len(), 3 * TEST_FINGERPRINT_DIMS_SMALL);
assert_eq!(counts.len(), 3 * TEST_FINGERPRINT_DIMS_SMALL);
}
#[test]
fn thread_safety() {
use std::sync::Arc;
use std::thread;
const THREAD_COUNT: usize = 4;
let Some(device) = device_or_skip("thread_safety") else {
return;
};
let device = Arc::new(device);
let engine = Fingerprints::builder()
.dimensions(TEST_FINGERPRINT_DIMS_SMALL)
.build(&device)
.expect("fingerprint engine should build on CPU");
let engine = Arc::new(engine);
let handles: Vec<_> = (0..THREAD_COUNT)
.map(|i| {
let device = Arc::clone(&device);
let engine = Arc::clone(&engine);
thread::spawn(move || {
let test_data = vec![format!("thread_{}_data", i)];
engine.compute(&device, &test_data, TEST_FINGERPRINT_DIMS_SMALL)
})
})
.collect();
let mut success_count = 0;
for handle in handles {
match handle.join().expect("worker thread should not panic") {
Ok(_) => success_count += 1,
Err(e) => println!("Thread computation failed: {:?}", e),
}
}
assert_eq!(success_count, THREAD_COUNT, "not all threads succeeded");
}
#[test]
fn large_batch_processing() {
let Some(device) = device_or_skip("large_batch_processing") else {
return;
};
let engine = Fingerprints::builder()
.dimensions(TEST_FINGERPRINT_DIMS_SMALL)
.build(&device)
.expect("fingerprint engine should build on CPU");
let large_batch: Vec<String> = (0..TEST_LARGE_BATCH_SIZE)
.map(|i| format!("test_string_{}", i))
.collect();
let large_batch_refs: Vec<&str> = large_batch.iter().map(|s| s.as_str()).collect();
let (hashes, counts) = engine
.compute(&device, &large_batch_refs, TEST_FINGERPRINT_DIMS_SMALL)
.expect("fingerprint computation should succeed on CPU");
assert_eq!(hashes.len(), TEST_LARGE_BATCH_SIZE * TEST_FINGERPRINT_DIMS_SMALL);
assert_eq!(counts.len(), TEST_LARGE_BATCH_SIZE * TEST_FINGERPRINT_DIMS_SMALL);
}
#[test]
fn similarity_estimation() {
let Some(device) = device_or_skip("similarity_estimation") else {
return;
};
let engine = Fingerprints::builder()
.dimensions(TEST_FINGERPRINT_DIMS_LARGE)
.build(&device)
.expect("fingerprint engine should build on CPU");
let test_strings = vec![
"the quick brown fox",
"the quick brown fox", "the quick brown dog", "completely different", ];
let (hashes, _counts) = engine
.compute(&device, &test_strings, TEST_FINGERPRINT_DIMS_LARGE)
.expect("fingerprint computation should succeed on CPU");
{
let dimensions = TEST_FINGERPRINT_DIMS_LARGE;
let mut matches_identical = 0;
for i in 0..dimensions {
if hashes[i] == hashes[1 * dimensions + i] {
matches_identical += 1;
}
}
let similarity_identical = matches_identical as f64 / dimensions as f64;
let mut matches_similar = 0;
for i in 0..dimensions {
if hashes[i] == hashes[2 * dimensions + i] {
matches_similar += 1;
}
}
let similarity_similar = matches_similar as f64 / dimensions as f64;
let mut matches_different = 0;
for i in 0..dimensions {
if hashes[i] == hashes[3 * dimensions + i] {
matches_different += 1;
}
}
let similarity_different = matches_different as f64 / dimensions as f64;
println!("Similarity identical: {:.3}", similarity_identical);
println!("Similarity similar: {:.3}", similarity_similar);
println!("Similarity different: {:.3}", similarity_different);
assert!(similarity_identical >= similarity_similar);
assert!(similarity_similar >= similarity_different);
}
}
}