#[cfg(feature = "gpu")]
use crate::gpu::GpuAccelerator;
#[cfg(feature = "gpu")]
use roaring::RoaringBitmap;
#[cfg(feature = "gpu")]
use std::collections::HashSet;
#[cfg(feature = "gpu")]
use std::sync::Arc;
#[cfg(feature = "gpu")]
pub struct GpuTrigramAccelerator {
accelerator: Arc<GpuAccelerator>,
}
#[cfg(feature = "gpu")]
impl GpuTrigramAccelerator {
pub fn new() -> Result<Self, String> {
let accelerator = GpuAccelerator::global().ok_or("GPU not available")?;
Ok(Self { accelerator })
}
#[must_use]
pub fn is_available() -> bool {
GpuAccelerator::is_available()
}
#[must_use]
pub fn batch_search(
&self,
patterns: &[&str],
inverted_index: &std::collections::HashMap<[u8; 3], RoaringBitmap>,
) -> Vec<RoaringBitmap> {
patterns
.iter()
.map(|pattern| Self::search_single(pattern, inverted_index))
.collect()
}
fn search_single(
pattern: &str,
inverted_index: &std::collections::HashMap<[u8; 3], RoaringBitmap>,
) -> RoaringBitmap {
let trigrams = Self::extract_trigrams_cpu(pattern);
if trigrams.is_empty() {
return RoaringBitmap::new();
}
let mut result: Option<RoaringBitmap> = None;
for trigram in &trigrams {
if let Some(bitmap) = inverted_index.get(trigram) {
result = Some(match result {
Some(r) => r & bitmap,
None => bitmap.clone(),
});
} else {
return RoaringBitmap::new();
}
}
result.unwrap_or_default()
}
#[must_use]
pub fn batch_extract_trigrams(&self, documents: &[&str]) -> Vec<HashSet<[u8; 3]>> {
documents
.iter()
.map(|doc| Self::extract_trigrams_cpu(doc))
.collect()
}
fn extract_trigrams_cpu(text: &str) -> HashSet<[u8; 3]> {
let bytes = text.as_bytes();
if bytes.len() < 3 {
return HashSet::new();
}
let mut trigrams = HashSet::with_capacity(bytes.len().saturating_sub(2));
for window in bytes.windows(3) {
trigrams.insert([window[0], window[1], window[2]]);
}
trigrams
}
#[must_use]
pub fn accelerator(&self) -> &GpuAccelerator {
&self.accelerator
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[allow(dead_code)] pub(crate) enum TrigramComputeBackend {
#[default]
CpuSimd,
#[cfg(feature = "gpu")]
Gpu,
}
#[allow(dead_code)] impl TrigramComputeBackend {
#[must_use]
pub fn auto_select(doc_count: usize, pattern_count: usize) -> Self {
#[cfg(not(feature = "gpu"))]
let _ = (doc_count, pattern_count);
#[cfg(feature = "gpu")]
{
if doc_count > 500_000 || (doc_count > 100_000 && pattern_count > 10) {
if crate::gpu::ComputeBackend::gpu_available() {
return Self::Gpu;
}
}
}
Self::CpuSimd
}
#[must_use]
pub const fn name(self) -> &'static str {
match self {
Self::CpuSimd => "CPU SIMD",
#[cfg(feature = "gpu")]
Self::Gpu => "GPU (wgpu)",
}
}
}
#[cfg(test)]
#[path = "gpu_tests.rs"]
mod tests;
#[cfg(all(test, feature = "gpu"))]
#[path = "gpu_feature_tests.rs"]
mod gpu_tests;