use super::types::{
copy_bytes_into_tape, copy_chars_into_tape, rust_error_from_c_message, should_use_64bit_for_bytes,
should_use_64bit_for_strings, SzSequenceFromBytes, SzSequenceFromChars, SzSequenceU32Tape, SzSequenceU64Tape,
};
use super::*;
use core::ffi::{c_char, c_void};
use core::ptr;
pub type LevenshteinDistancesHandle = *mut c_void;
pub type LevenshteinDistancesUtf8Handle = *mut c_void;
pub type NeedlemanWunschScoresHandle = *mut c_void;
pub type SmithWatermanScoresHandle = *mut c_void;
extern "C" {
fn szs_levenshtein_distances_init(
match_cost: i8,
mismatch_cost: i8,
open_cost: i8,
extend_cost: i8,
alloc: *const c_void,
capabilities: Capability,
engine: *mut LevenshteinDistancesHandle,
error_message: *mut *const c_char,
) -> Status;
fn szs_levenshtein_distances(
engine: LevenshteinDistancesHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut usize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_levenshtein_distances_u32tape(
engine: LevenshteinDistancesHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut usize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_levenshtein_distances_u64tape(
engine: LevenshteinDistancesHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut usize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_levenshtein_distances_free(engine: LevenshteinDistancesHandle);
fn szs_levenshtein_distances_utf8_init(
match_cost: i8,
mismatch_cost: i8,
open_cost: i8,
extend_cost: i8,
alloc: *const c_void,
capabilities: Capability,
engine: *mut LevenshteinDistancesUtf8Handle,
error_message: *mut *const c_char,
) -> Status;
fn szs_levenshtein_distances_utf8(
engine: LevenshteinDistancesUtf8Handle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut usize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_levenshtein_distances_utf8_u32tape(
engine: LevenshteinDistancesUtf8Handle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut usize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_levenshtein_distances_utf8_u64tape(
engine: LevenshteinDistancesUtf8Handle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut usize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_levenshtein_distances_utf8_free(engine: LevenshteinDistancesUtf8Handle);
fn szs_needleman_wunsch_scores_init(
byte_to_class: *const u8, class_substitution_costs: *const i8, open_cost: i8,
extend_cost: i8,
alloc: *const c_void,
capabilities: Capability,
engine: *mut NeedlemanWunschScoresHandle,
error_message: *mut *const c_char,
) -> Status;
fn szs_needleman_wunsch_scores(
engine: NeedlemanWunschScoresHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut isize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_needleman_wunsch_scores_u32tape(
engine: NeedlemanWunschScoresHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut isize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_needleman_wunsch_scores_u64tape(
engine: NeedlemanWunschScoresHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut isize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_needleman_wunsch_scores_free(engine: NeedlemanWunschScoresHandle);
fn szs_smith_waterman_scores_init(
byte_to_class: *const u8, class_substitution_costs: *const i8, open_cost: i8,
extend_cost: i8,
alloc: *const c_void,
capabilities: Capability,
engine: *mut SmithWatermanScoresHandle,
error_message: *mut *const c_char,
) -> Status;
fn szs_smith_waterman_scores(
engine: SmithWatermanScoresHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut isize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_smith_waterman_scores_u32tape(
engine: SmithWatermanScoresHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut isize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_smith_waterman_scores_u64tape(
engine: SmithWatermanScoresHandle,
device: *mut c_void,
queries: *const c_void, candidates: *const c_void, results: *mut isize,
results_row_stride: usize,
error_message: *mut *const c_char,
) -> Status;
fn szs_smith_waterman_scores_free(engine: SmithWatermanScoresHandle);
}
pub struct LevenshteinDistances {
handle: LevenshteinDistancesHandle,
}
impl LevenshteinDistances {
pub fn new(
device: &DeviceScope,
match_cost: i8,
mismatch_cost: i8,
open_cost: i8,
extend_cost: i8,
) -> Result<Self, Error> {
let mut handle = ptr::null_mut();
let capabilities = device.get_capabilities().unwrap_or(0);
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_levenshtein_distances_init(
match_cost,
mismatch_cost,
open_cost,
extend_cost,
ptr::null(),
capabilities,
&mut handle,
&mut error_msg,
)
};
match status {
Status::Success => Ok(Self { handle }),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
pub fn compute<Sequences, Sequence>(
&self,
device: &DeviceScope,
queries: Sequences,
candidates: Sequences,
) -> Result<UnifiedMat<usize>, Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<[u8]>,
{
let queries_slice = queries.as_ref();
let candidates_slice = candidates.as_ref();
let mut matrix = UnifiedMat::<usize>::try_allocate(queries_slice.len(), candidates_slice.len())?;
self.compute_pair(device, queries_slice, Some(candidates_slice), &mut matrix)?;
Ok(matrix)
}
pub fn compute_symmetric<Sequences, Sequence>(
&self,
device: &DeviceScope,
sequences: Sequences,
) -> Result<UnifiedMat<usize>, Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<[u8]>,
{
let sequences_slice = sequences.as_ref();
let mut matrix = UnifiedMat::<usize>::try_allocate(sequences_slice.len(), sequences_slice.len())?;
self.compute_pair(device, sequences_slice, None, &mut matrix)?;
Ok(matrix)
}
fn compute_pair<Sequence>(
&self,
device: &DeviceScope,
queries: &[Sequence],
candidates: Option<&[Sequence]>,
matrix: &mut UnifiedMat<usize>,
) -> Result<(), Error>
where
Sequence: AsRef<[u8]>,
{
if device.is_gpu() {
let force_64bit = match candidates {
Some(candidates_slice) => should_use_64bit_for_bytes(queries, candidates_slice),
None => should_use_64bit_for_bytes(queries, queries),
};
let queries_tape = copy_bytes_into_tape(queries, force_64bit)?;
let candidates_tape = match candidates {
Some(candidates_slice) => Some(copy_bytes_into_tape(candidates_slice, force_64bit)?),
None => None,
};
return self.compute_into(device, queries_tape, candidates_tape, matrix);
}
let queries_sequence = SzSequenceFromBytes::to_sz_sequence(queries);
let candidates_sequence = candidates.map(SzSequenceFromBytes::to_sz_sequence);
let candidates_ptr = match &candidates_sequence {
Some(sequence) => sequence as *const _ as *const c_void,
None => ptr::null(),
};
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_levenshtein_distances(
self.handle,
device.handle,
&queries_sequence as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
pub fn compute_into<'a>(
&self,
device: &DeviceScope,
queries: AnyBytesTape<'a>,
candidates: Option<AnyBytesTape<'a>>,
matrix: &mut UnifiedMat<usize>,
) -> Result<(), Error> {
let mut error_msg: *const c_char = ptr::null();
let queries64 = match &queries {
AnyBytesTape::Tape64(tape) => Some(SzSequenceU64Tape::from(tape)),
AnyBytesTape::View64(view) => Some(SzSequenceU64Tape::from(view)),
_ => None,
};
if let Some(queries_view) = queries64 {
let candidates_view = match &candidates {
Some(AnyBytesTape::Tape64(tape)) => Some(SzSequenceU64Tape::from(tape)),
Some(AnyBytesTape::View64(view)) => Some(SzSequenceU64Tape::from(view)),
Some(_) => return Err(Error::from(SzStatus::UnexpectedDimensions)),
None => None,
};
let candidates_count = candidates_view.map(|view| view.count).unwrap_or(queries_view.count);
if matrix.queries_count != queries_view.count || matrix.candidates_count != candidates_count {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let candidates_ptr = match &candidates_view {
Some(view) => view as *const _ as *const c_void,
None => ptr::null(),
};
let status = unsafe {
szs_levenshtein_distances_u64tape(
self.handle,
device.handle,
&queries_view as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
return match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
};
}
let queries32 = match &queries {
AnyBytesTape::Tape32(tape) => Some(SzSequenceU32Tape::from(tape)),
AnyBytesTape::View32(view) => Some(SzSequenceU32Tape::from(view)),
_ => None,
};
if let Some(queries_view) = queries32 {
let candidates_view = match &candidates {
Some(AnyBytesTape::Tape32(tape)) => Some(SzSequenceU32Tape::from(tape)),
Some(AnyBytesTape::View32(view)) => Some(SzSequenceU32Tape::from(view)),
Some(_) => return Err(Error::from(SzStatus::UnexpectedDimensions)),
None => None,
};
let candidates_count = candidates_view.map(|view| view.count).unwrap_or(queries_view.count);
if matrix.queries_count != queries_view.count || matrix.candidates_count != candidates_count {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let candidates_ptr = match &candidates_view {
Some(view) => view as *const _ as *const c_void,
None => ptr::null(),
};
let status = unsafe {
szs_levenshtein_distances_u32tape(
self.handle,
device.handle,
&queries_view as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
return match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
};
}
Err(Error::from(SzStatus::UnexpectedDimensions))
}
}
impl Drop for LevenshteinDistances {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe { szs_levenshtein_distances_free(self.handle) };
}
}
}
unsafe impl Send for LevenshteinDistances {}
unsafe impl Sync for LevenshteinDistances {}
pub struct LevenshteinDistancesUtf8 {
handle: LevenshteinDistancesUtf8Handle,
}
impl LevenshteinDistancesUtf8 {
pub fn new(
device: &DeviceScope,
match_cost: i8,
mismatch_cost: i8,
open_cost: i8,
extend_cost: i8,
) -> Result<Self, Error> {
let mut handle = ptr::null_mut();
let capabilities = device.get_capabilities().unwrap_or(0);
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_levenshtein_distances_utf8_init(
match_cost,
mismatch_cost,
open_cost,
extend_cost,
ptr::null(),
capabilities,
&mut handle,
&mut error_msg,
)
};
match status {
Status::Success => Ok(Self { handle }),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
pub fn compute<Sequences, Sequence>(
&self,
device: &DeviceScope,
queries: Sequences,
candidates: Sequences,
) -> Result<UnifiedMat<usize>, Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<str>,
{
let queries_slice = queries.as_ref();
let candidates_slice = candidates.as_ref();
let mut matrix = UnifiedMat::<usize>::try_allocate(queries_slice.len(), candidates_slice.len())?;
self.compute_pair(device, queries_slice, Some(candidates_slice), &mut matrix)?;
Ok(matrix)
}
pub fn compute_symmetric<Sequences, Sequence>(
&self,
device: &DeviceScope,
sequences: Sequences,
) -> Result<UnifiedMat<usize>, Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<str>,
{
let sequences_slice = sequences.as_ref();
let mut matrix = UnifiedMat::<usize>::try_allocate(sequences_slice.len(), sequences_slice.len())?;
self.compute_pair(device, sequences_slice, None, &mut matrix)?;
Ok(matrix)
}
fn compute_pair<Sequence>(
&self,
device: &DeviceScope,
queries: &[Sequence],
candidates: Option<&[Sequence]>,
matrix: &mut UnifiedMat<usize>,
) -> Result<(), Error>
where
Sequence: AsRef<str>,
{
if device.is_gpu() {
let force_64bit = match candidates {
Some(candidates_slice) => should_use_64bit_for_strings(queries, candidates_slice),
None => should_use_64bit_for_strings(queries, queries),
};
let queries_tape = copy_chars_into_tape(queries, force_64bit)?;
let candidates_tape = match candidates {
Some(candidates_slice) => Some(copy_chars_into_tape(candidates_slice, force_64bit)?),
None => None,
};
return self.compute_into(device, queries_tape, candidates_tape, matrix);
}
let queries_sequence = SzSequenceFromChars::to_sz_sequence(queries);
let candidates_sequence = candidates.map(SzSequenceFromChars::to_sz_sequence);
let candidates_ptr = match &candidates_sequence {
Some(sequence) => sequence as *const _ as *const c_void,
None => ptr::null(),
};
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_levenshtein_distances_utf8(
self.handle,
device.handle,
&queries_sequence as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
pub fn compute_into<'a>(
&self,
device: &DeviceScope,
queries: AnyCharsTape<'a>,
candidates: Option<AnyCharsTape<'a>>,
matrix: &mut UnifiedMat<usize>,
) -> Result<(), Error> {
let mut error_msg: *const c_char = ptr::null();
let queries64 = match &queries {
AnyCharsTape::Tape64(tape) => Some(SzSequenceU64Tape::from(tape)),
AnyCharsTape::View64(view) => Some(SzSequenceU64Tape::from(view)),
_ => None,
};
if let Some(queries_view) = queries64 {
let candidates_view = match &candidates {
Some(AnyCharsTape::Tape64(tape)) => Some(SzSequenceU64Tape::from(tape)),
Some(AnyCharsTape::View64(view)) => Some(SzSequenceU64Tape::from(view)),
Some(_) => return Err(Error::from(SzStatus::UnexpectedDimensions)),
None => None,
};
let candidates_count = candidates_view.map(|view| view.count).unwrap_or(queries_view.count);
if matrix.queries_count != queries_view.count || matrix.candidates_count != candidates_count {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let candidates_ptr = match &candidates_view {
Some(view) => view as *const _ as *const c_void,
None => ptr::null(),
};
let status = unsafe {
szs_levenshtein_distances_utf8_u64tape(
self.handle,
device.handle,
&queries_view as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
return match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
};
}
let queries32 = match &queries {
AnyCharsTape::Tape32(tape) => Some(SzSequenceU32Tape::from(tape)),
AnyCharsTape::View32(view) => Some(SzSequenceU32Tape::from(view)),
_ => None,
};
if let Some(queries_view) = queries32 {
let candidates_view = match &candidates {
Some(AnyCharsTape::Tape32(tape)) => Some(SzSequenceU32Tape::from(tape)),
Some(AnyCharsTape::View32(view)) => Some(SzSequenceU32Tape::from(view)),
Some(_) => return Err(Error::from(SzStatus::UnexpectedDimensions)),
None => None,
};
let candidates_count = candidates_view.map(|view| view.count).unwrap_or(queries_view.count);
if matrix.queries_count != queries_view.count || matrix.candidates_count != candidates_count {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let candidates_ptr = match &candidates_view {
Some(view) => view as *const _ as *const c_void,
None => ptr::null(),
};
let status = unsafe {
szs_levenshtein_distances_utf8_u32tape(
self.handle,
device.handle,
&queries_view as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
return match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
};
}
Err(Error::from(SzStatus::UnexpectedDimensions))
}
}
impl Drop for LevenshteinDistancesUtf8 {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe { szs_levenshtein_distances_utf8_free(self.handle) };
}
}
}
unsafe impl Send for LevenshteinDistancesUtf8 {}
unsafe impl Sync for LevenshteinDistancesUtf8 {}
pub struct NeedlemanWunschScores {
handle: NeedlemanWunschScoresHandle,
}
impl NeedlemanWunschScores {
pub fn new(
device: &DeviceScope,
byte_to_class: &[u8; 256],
class_substitution_costs: &[[i8; 32]; 32],
open_cost: i8,
extend_cost: i8,
) -> Result<Self, Error> {
let mut handle = ptr::null_mut();
let capabilities = device.get_capabilities().unwrap_or(0);
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_needleman_wunsch_scores_init(
byte_to_class.as_ptr() as *const u8,
class_substitution_costs.as_ptr() as *const i8,
open_cost,
extend_cost,
ptr::null(),
capabilities,
&mut handle,
&mut error_msg,
)
};
match status {
Status::Success => Ok(Self { handle }),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
pub fn compute<Sequences, Sequence>(
&self,
device: &DeviceScope,
queries: Sequences,
candidates: Sequences,
) -> Result<UnifiedMat<isize>, Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<[u8]>,
{
let queries_slice = queries.as_ref();
let candidates_slice = candidates.as_ref();
let mut matrix = UnifiedMat::<isize>::try_allocate(queries_slice.len(), candidates_slice.len())?;
self.compute_pair(device, queries_slice, Some(candidates_slice), &mut matrix)?;
Ok(matrix)
}
pub fn compute_symmetric<Sequences, Sequence>(
&self,
device: &DeviceScope,
sequences: Sequences,
) -> Result<UnifiedMat<isize>, Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<[u8]>,
{
let sequences_slice = sequences.as_ref();
let mut matrix = UnifiedMat::<isize>::try_allocate(sequences_slice.len(), sequences_slice.len())?;
self.compute_pair(device, sequences_slice, None, &mut matrix)?;
Ok(matrix)
}
fn compute_pair<Sequence>(
&self,
device: &DeviceScope,
queries: &[Sequence],
candidates: Option<&[Sequence]>,
matrix: &mut UnifiedMat<isize>,
) -> Result<(), Error>
where
Sequence: AsRef<[u8]>,
{
if device.is_gpu() {
let force_64bit = match candidates {
Some(candidates_slice) => should_use_64bit_for_bytes(queries, candidates_slice),
None => should_use_64bit_for_bytes(queries, queries),
};
let queries_tape = copy_bytes_into_tape(queries, force_64bit)?;
let candidates_tape = match candidates {
Some(candidates_slice) => Some(copy_bytes_into_tape(candidates_slice, force_64bit)?),
None => None,
};
return self.compute_into(device, queries_tape, candidates_tape, matrix);
}
let queries_sequence = SzSequenceFromBytes::to_sz_sequence(queries);
let candidates_sequence = candidates.map(SzSequenceFromBytes::to_sz_sequence);
let candidates_ptr = match &candidates_sequence {
Some(sequence) => sequence as *const _ as *const c_void,
None => ptr::null(),
};
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_needleman_wunsch_scores(
self.handle,
device.handle,
&queries_sequence as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
pub fn compute_into<'a>(
&self,
device: &DeviceScope,
queries: AnyBytesTape<'a>,
candidates: Option<AnyBytesTape<'a>>,
matrix: &mut UnifiedMat<isize>,
) -> Result<(), Error> {
let mut error_msg: *const c_char = ptr::null();
let queries64 = match &queries {
AnyBytesTape::Tape64(tape) => Some(SzSequenceU64Tape::from(tape)),
AnyBytesTape::View64(view) => Some(SzSequenceU64Tape::from(view)),
_ => None,
};
if let Some(queries_view) = queries64 {
let candidates_view = match &candidates {
Some(AnyBytesTape::Tape64(tape)) => Some(SzSequenceU64Tape::from(tape)),
Some(AnyBytesTape::View64(view)) => Some(SzSequenceU64Tape::from(view)),
Some(_) => return Err(Error::from(SzStatus::UnexpectedDimensions)),
None => None,
};
let candidates_count = candidates_view.map(|view| view.count).unwrap_or(queries_view.count);
if matrix.queries_count != queries_view.count || matrix.candidates_count != candidates_count {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let candidates_ptr = match &candidates_view {
Some(view) => view as *const _ as *const c_void,
None => ptr::null(),
};
let status = unsafe {
szs_needleman_wunsch_scores_u64tape(
self.handle,
device.handle,
&queries_view as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
return match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
};
}
let queries32 = match &queries {
AnyBytesTape::Tape32(tape) => Some(SzSequenceU32Tape::from(tape)),
AnyBytesTape::View32(view) => Some(SzSequenceU32Tape::from(view)),
_ => None,
};
if let Some(queries_view) = queries32 {
let candidates_view = match &candidates {
Some(AnyBytesTape::Tape32(tape)) => Some(SzSequenceU32Tape::from(tape)),
Some(AnyBytesTape::View32(view)) => Some(SzSequenceU32Tape::from(view)),
Some(_) => return Err(Error::from(SzStatus::UnexpectedDimensions)),
None => None,
};
let candidates_count = candidates_view.map(|view| view.count).unwrap_or(queries_view.count);
if matrix.queries_count != queries_view.count || matrix.candidates_count != candidates_count {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let candidates_ptr = match &candidates_view {
Some(view) => view as *const _ as *const c_void,
None => ptr::null(),
};
let status = unsafe {
szs_needleman_wunsch_scores_u32tape(
self.handle,
device.handle,
&queries_view as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
return match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
};
}
Err(Error::from(SzStatus::UnexpectedDimensions))
}
}
impl Drop for NeedlemanWunschScores {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe { szs_needleman_wunsch_scores_free(self.handle) };
}
}
}
unsafe impl Send for NeedlemanWunschScores {}
unsafe impl Sync for NeedlemanWunschScores {}
pub struct SmithWatermanScores {
handle: SmithWatermanScoresHandle,
}
impl SmithWatermanScores {
pub fn new(
device: &DeviceScope,
byte_to_class: &[u8; 256],
class_substitution_costs: &[[i8; 32]; 32],
open_cost: i8,
extend_cost: i8,
) -> Result<Self, Error> {
let mut handle = ptr::null_mut();
let capabilities = device.get_capabilities().unwrap_or(0);
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_smith_waterman_scores_init(
byte_to_class.as_ptr() as *const u8,
class_substitution_costs.as_ptr() as *const i8,
open_cost,
extend_cost,
ptr::null(),
capabilities,
&mut handle,
&mut error_msg,
)
};
match status {
Status::Success => Ok(Self { handle }),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
pub fn compute<Sequences, Sequence>(
&self,
device: &DeviceScope,
queries: Sequences,
candidates: Sequences,
) -> Result<UnifiedMat<isize>, Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<[u8]>,
{
let queries_slice = queries.as_ref();
let candidates_slice = candidates.as_ref();
let mut matrix = UnifiedMat::<isize>::try_allocate(queries_slice.len(), candidates_slice.len())?;
self.compute_pair(device, queries_slice, Some(candidates_slice), &mut matrix)?;
Ok(matrix)
}
pub fn compute_symmetric<Sequences, Sequence>(
&self,
device: &DeviceScope,
sequences: Sequences,
) -> Result<UnifiedMat<isize>, Error>
where
Sequences: AsRef<[Sequence]>,
Sequence: AsRef<[u8]>,
{
let sequences_slice = sequences.as_ref();
let mut matrix = UnifiedMat::<isize>::try_allocate(sequences_slice.len(), sequences_slice.len())?;
self.compute_pair(device, sequences_slice, None, &mut matrix)?;
Ok(matrix)
}
fn compute_pair<Sequence>(
&self,
device: &DeviceScope,
queries: &[Sequence],
candidates: Option<&[Sequence]>,
matrix: &mut UnifiedMat<isize>,
) -> Result<(), Error>
where
Sequence: AsRef<[u8]>,
{
if device.is_gpu() {
let force_64bit = match candidates {
Some(candidates_slice) => should_use_64bit_for_bytes(queries, candidates_slice),
None => should_use_64bit_for_bytes(queries, queries),
};
let queries_tape = copy_bytes_into_tape(queries, force_64bit)?;
let candidates_tape = match candidates {
Some(candidates_slice) => Some(copy_bytes_into_tape(candidates_slice, force_64bit)?),
None => None,
};
return self.compute_into(device, queries_tape, candidates_tape, matrix);
}
let queries_sequence = SzSequenceFromBytes::to_sz_sequence(queries);
let candidates_sequence = candidates.map(SzSequenceFromBytes::to_sz_sequence);
let candidates_ptr = match &candidates_sequence {
Some(sequence) => sequence as *const _ as *const c_void,
None => ptr::null(),
};
let mut error_msg: *const c_char = ptr::null();
let status = unsafe {
szs_smith_waterman_scores(
self.handle,
device.handle,
&queries_sequence as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
}
}
pub fn compute_into<'a>(
&self,
device: &DeviceScope,
queries: AnyBytesTape<'a>,
candidates: Option<AnyBytesTape<'a>>,
matrix: &mut UnifiedMat<isize>,
) -> Result<(), Error> {
let mut error_msg: *const c_char = ptr::null();
let queries64 = match &queries {
AnyBytesTape::Tape64(tape) => Some(SzSequenceU64Tape::from(tape)),
AnyBytesTape::View64(view) => Some(SzSequenceU64Tape::from(view)),
_ => None,
};
if let Some(queries_view) = queries64 {
let candidates_view = match &candidates {
Some(AnyBytesTape::Tape64(tape)) => Some(SzSequenceU64Tape::from(tape)),
Some(AnyBytesTape::View64(view)) => Some(SzSequenceU64Tape::from(view)),
Some(_) => return Err(Error::from(SzStatus::UnexpectedDimensions)),
None => None,
};
let candidates_count = candidates_view.map(|view| view.count).unwrap_or(queries_view.count);
if matrix.queries_count != queries_view.count || matrix.candidates_count != candidates_count {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let candidates_ptr = match &candidates_view {
Some(view) => view as *const _ as *const c_void,
None => ptr::null(),
};
let status = unsafe {
szs_smith_waterman_scores_u64tape(
self.handle,
device.handle,
&queries_view as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
return match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
};
}
let queries32 = match &queries {
AnyBytesTape::Tape32(tape) => Some(SzSequenceU32Tape::from(tape)),
AnyBytesTape::View32(view) => Some(SzSequenceU32Tape::from(view)),
_ => None,
};
if let Some(queries_view) = queries32 {
let candidates_view = match &candidates {
Some(AnyBytesTape::Tape32(tape)) => Some(SzSequenceU32Tape::from(tape)),
Some(AnyBytesTape::View32(view)) => Some(SzSequenceU32Tape::from(view)),
Some(_) => return Err(Error::from(SzStatus::UnexpectedDimensions)),
None => None,
};
let candidates_count = candidates_view.map(|view| view.count).unwrap_or(queries_view.count);
if matrix.queries_count != queries_view.count || matrix.candidates_count != candidates_count {
return Err(Error::from(SzStatus::UnexpectedDimensions));
}
let candidates_ptr = match &candidates_view {
Some(view) => view as *const _ as *const c_void,
None => ptr::null(),
};
let status = unsafe {
szs_smith_waterman_scores_u32tape(
self.handle,
device.handle,
&queries_view as *const _ as *const c_void,
candidates_ptr,
matrix.data.as_mut_ptr(),
matrix.row_stride,
&mut error_msg,
)
};
return match status {
Status::Success => Ok(()),
err => Err(rust_error_from_c_message(err, error_msg)),
};
}
Err(Error::from(SzStatus::UnexpectedDimensions))
}
}
impl Drop for SmithWatermanScores {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe { szs_smith_waterman_scores_free(self.handle) };
}
}
}
unsafe impl Send for SmithWatermanScores {}
unsafe impl Sync for SmithWatermanScores {}
pub fn error_costs_classes_diagonal(match_score: i8, mismatch_score: i8) -> ([u8; 256], [[i8; 32]; 32]) {
let mut byte_to_class = [0u8; 256];
for i in 0..256 {
byte_to_class[i] = (i % 32) as u8;
}
let mut class_costs = [[0i8; 32]; 32];
for i in 0..32 {
for j in 0..32 {
class_costs[i][j] = if i == j { match_score } else { mismatch_score };
}
}
(byte_to_class, class_costs)
}
pub fn error_costs_classes_unary() -> ([u8; 256], [[i8; 32]; 32]) {
error_costs_classes_diagonal(0, -1)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::stringzillas::fixtures::device_or_skip;
use stringtape::BytesTape;
#[test]
fn levenshtein_distance_engine() {
let Some(device) = device_or_skip("levenshtein_distance_engine") else {
return;
};
let engine = LevenshteinDistances::new(
&device, 0, 1, 1, 1, )
.expect("Levenshtein engine should build on CPU");
let queries = vec!["kitten", "saturday"];
let candidates = vec!["sitting", "sunday", "kitten"];
let matrix = engine
.compute(&device, &queries, &candidates)
.expect("Levenshtein computation should succeed on CPU");
assert_eq!(matrix.dimensions(), (2, 3));
assert_eq!(matrix[(0, 0)], 3);
assert_eq!(matrix[(0, 2)], 0);
assert_eq!(matrix[(1, 1)], 3);
assert_eq!(matrix.row(0), &[3usize, 6, 0][..]);
}
#[test]
fn levenshtein_distance_symmetric() {
let Some(device) = device_or_skip("levenshtein_distance_symmetric") else {
return;
};
let engine = LevenshteinDistances::new(&device, 0, 1, 1, 1).expect("Levenshtein engine should build on CPU");
let sequences = vec!["cat", "bat", "cart"];
let matrix = engine
.compute_symmetric(&device, &sequences)
.expect("symmetric Levenshtein computation should succeed on CPU");
assert_eq!(matrix.dimensions(), (3, 3));
for diagonal_index in 0..3 {
assert_eq!(matrix[(diagonal_index, diagonal_index)], 0);
}
for first_index in 0..3 {
for second_index in 0..3 {
assert_eq!(matrix[(first_index, second_index)], matrix[(second_index, first_index)]);
}
}
assert_eq!(matrix[(0, 1)], 1);
let single = vec!["cat"];
let single_matrix = engine
.compute(&device, &single, &single)
.expect("Levenshtein computation should succeed on CPU");
assert_eq!(single_matrix[(0, 0)], matrix[(0, 0)]);
}
#[test]
fn levenshtein_utf8_engine() {
let Some(device) = device_or_skip("levenshtein_utf8_engine") else {
return;
};
let engine =
LevenshteinDistancesUtf8::new(&device, 0, 1, 1, 1).expect("UTF-8 Levenshtein engine should build on CPU");
let queries = vec!["café", "naïve"];
let candidates = vec!["cafe", "naive"];
let matrix = engine
.compute(&device, &queries, &candidates)
.expect("UTF-8 Levenshtein computation should succeed on CPU");
assert_eq!(matrix.dimensions(), (2, 2));
assert_eq!(matrix[(0, 0)], 1); assert_eq!(matrix[(1, 1)], 1); }
#[test]
fn needleman_wunsch_engine() {
let Some(device) = device_or_skip("needleman_wunsch_engine") else {
return;
};
let (byte_to_class, class_costs) = error_costs_classes_diagonal(2, -1);
let engine = NeedlemanWunschScores::new(&device, &byte_to_class, &class_costs, -2, -1)
.expect("Needleman-Wunsch engine should build on CPU");
let queries = vec!["ACGT", "ACGT"];
let candidates = vec!["ACGT", "TTTT"];
let matrix = engine
.compute(&device, &queries, &candidates)
.expect("Needleman-Wunsch computation should succeed on CPU");
assert_eq!(matrix.dimensions(), (2, 2));
assert!(matrix[(0, 0)] > matrix[(0, 1)]);
assert!(matrix[(0, 0)] > 0, "Identical sequences should score positively");
let sequences = vec!["ACGT", "AGGT", "TTTT"];
let symmetric = engine
.compute_symmetric(&device, &sequences)
.expect("Needleman-Wunsch symmetric computation should succeed on CPU");
assert_eq!(symmetric.dimensions(), (3, 3));
for first_index in 0..3 {
for second_index in 0..3 {
assert_eq!(
symmetric[(first_index, second_index)],
symmetric[(second_index, first_index)]
);
}
}
for diagonal_index in 0..3 {
let diagonal_score = symmetric[(diagonal_index, diagonal_index)];
for candidate_index in 0..3 {
assert!(diagonal_score >= symmetric[(diagonal_index, candidate_index)]);
}
}
}
#[test]
fn smith_waterman_engine() {
let Some(device) = device_or_skip("smith_waterman_engine") else {
return;
};
let (byte_to_class, class_costs) = error_costs_classes_diagonal(3, -1);
let engine = SmithWatermanScores::new(&device, &byte_to_class, &class_costs, -2, -1)
.expect("Smith-Waterman engine should build on CPU");
let queries = vec!["ACGTACGT"];
let candidates = vec!["ACGT", "TTTT"];
let matrix = engine
.compute(&device, &queries, &candidates)
.expect("Smith-Waterman computation should succeed on CPU");
assert_eq!(matrix.dimensions(), (1, 2));
assert!(matrix[(0, 0)] > matrix[(0, 1)]);
assert!(matrix[(0, 0)] > 0, "Local alignment should be positive");
let sequences = vec!["ACGTACGT", "ACGT", "TTTT"];
let symmetric = engine
.compute_symmetric(&device, &sequences)
.expect("Smith-Waterman symmetric computation should succeed on CPU");
assert_eq!(symmetric.dimensions(), (3, 3));
for first_index in 0..3 {
for second_index in 0..3 {
assert_eq!(
symmetric[(first_index, second_index)],
symmetric[(second_index, first_index)]
);
}
}
}
#[test]
fn error_costs_for_needleman_wunsch() {
let Some(device) = device_or_skip("error_costs_for_needleman_wunsch") else {
return;
};
let (byte_to_class, class_costs) = error_costs_classes_diagonal(2, -1);
let engine = NeedlemanWunschScores::new(&device, &byte_to_class, &class_costs, -2, -1)
.expect("Needleman-Wunsch engine should build on CPU");
let queries = vec!["ABCD"];
let candidates = vec!["ABCD"];
let matrix = engine
.compute(&device, &queries, &candidates)
.expect("Needleman-Wunsch computation should succeed on CPU");
assert_eq!(matrix.dimensions(), (1, 1));
assert!(matrix[(0, 0)] > 0, "Identical sequences should have positive score");
}
#[test]
fn levenshtein_compute_into_u32_bytes() {
let Some(device) = device_or_skip("levenshtein_compute_into_u32_bytes") else {
return;
};
let engine = LevenshteinDistances::new(&device, 0, 1, 1, 1).expect("Levenshtein engine should build on CPU");
let queries = [b"kitten".as_ref(), b"saturday".as_ref()];
let candidates = [b"sitting".as_ref(), b"sunday".as_ref()];
let mut queries_tape = BytesTape::<u32, UnifiedAlloc>::new_in(UnifiedAlloc);
queries_tape.extend(queries).unwrap();
let mut candidates_tape = BytesTape::<u32, UnifiedAlloc>::new_in(UnifiedAlloc);
candidates_tape.extend(candidates).unwrap();
let mut matrix = UnifiedMat::<usize>::try_allocate(2, 2).expect("matrix allocation");
engine
.compute_into(
&device,
AnyBytesTape::Tape32(queries_tape),
Some(AnyBytesTape::Tape32(candidates_tape)),
&mut matrix,
)
.expect("Levenshtein compute_into should succeed on CPU");
assert_eq!(matrix[(0, 0)], 3);
assert_eq!(matrix[(1, 1)], 3);
}
#[test]
fn levenshtein_compute_into_u64_bytes() {
let Some(device) = device_or_skip("levenshtein_compute_into_u64_bytes") else {
return;
};
let engine = LevenshteinDistances::new(&device, 0, 1, 1, 1).expect("Levenshtein engine should build on CPU");
let queries = [b"abc".as_ref(), b"abcdef".as_ref()];
let candidates = [b"yabd".as_ref(), b"abcxef".as_ref()];
let mut queries_tape = BytesTape::<u64, UnifiedAlloc>::new_in(UnifiedAlloc);
queries_tape.extend(queries).unwrap();
let mut candidates_tape = BytesTape::<u64, UnifiedAlloc>::new_in(UnifiedAlloc);
candidates_tape.extend(candidates).unwrap();
let mut matrix = UnifiedMat::<usize>::try_allocate(2, 2).expect("matrix allocation");
engine
.compute_into(
&device,
AnyBytesTape::Tape64(queries_tape),
Some(AnyBytesTape::Tape64(candidates_tape)),
&mut matrix,
)
.expect("Levenshtein compute_into should succeed on CPU");
assert_eq!(matrix[(0, 0)], 2);
assert_eq!(matrix[(1, 1)], 1);
}
}