use crate::alphabet::ALPHABET_SIZE;
use crate::error::FmIndexError;
use crate::fm_index::FmIndex;
use crate::gpu::GpuContext;
use wgpu;
const LOCATE_SEARCH_SHADER: &str = include_str!("../../shaders/locate_search.wgsl");
const LOCATE_RESOLVE_SHADER: &str = include_str!("../../shaders/locate_resolve.wgsl");
#[repr(C)]
#[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
struct SearchParams {
num_queries: u32,
text_len: u32,
num_blocks: u32,
_pad: u32,
c: [u32; 16],
}
#[repr(C)]
#[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
struct ResolveParams {
num_queries: u32,
text_len: u32,
num_blocks: u32,
num_seqs: u32,
sample_rate: u32,
total_matches: u32,
c: [u32; 16],
_pad2: [u32; 2],
}
pub async fn locate_batch_gpu(
ctx: &GpuContext,
index: &FmIndex,
queries: &[&[u8]],
) -> Result<Vec<Vec<(u32, u32)>>, FmIndexError> {
if queries.is_empty() {
return Ok(vec![]);
}
let num_queries = queries.len() as u32;
let text_len = index.text_len;
let block_size: u32 = 64;
let num_blocks = text_len.div_ceil(block_size);
let num_seqs = index.num_sequences;
let sample_rate = index.sa_samples.sample_rate;
let alpha = ALPHABET_SIZE as u32;
let c_arr: [u32; 16] = index.c_array.data;
let mut checkpoints_flat: Vec<u32> = Vec::with_capacity((num_blocks * alpha) as usize);
for block in index.occ.flat_block_checkpoints() {
checkpoints_flat.extend_from_slice(&block);
}
let mut bitvectors_flat: Vec<u32> = Vec::with_capacity((num_blocks * alpha * 2) as usize);
for block in &index.occ.bitvectors_full16() {
for &bv64 in block.iter() {
bitvectors_flat.push(bv64 as u32);
bitvectors_flat.push((bv64 >> 32) as u32);
}
}
let mut queries_flat: Vec<u32> = Vec::new();
let mut query_offsets: Vec<u32> = Vec::with_capacity(queries.len() + 1);
query_offsets.push(0);
for &q in queries {
for &b in q {
queries_flat.push(b as u32);
}
query_offsets.push(queries_flat.len() as u32);
}
if queries_flat.is_empty() {
queries_flat.push(0); }
let bwt_u32: Vec<u32> = index.occ.reconstruct_bwt_u32();
let sa_samples_data: Vec<u32> = index.sa_samples.to_flat_vec(index.text_len as usize);
let seq_bounds_data: Vec<u32> = index.seq_boundaries.clone();
let bwt_buf = ctx.create_buffer_init("locate_bwt", &bwt_u32);
let chk_buf = ctx.create_buffer_init("locate_chk", &checkpoints_flat);
let bv_buf = ctx.create_buffer_init("locate_bv", &bitvectors_flat);
let sa_buf = ctx.create_buffer_init("locate_sa", &sa_samples_data);
let seqb_buf = ctx.create_buffer_init("locate_seqb", &seq_bounds_data);
let qflat_buf = ctx.create_buffer_init("locate_qflat", &queries_flat);
let qoff_buf = ctx.create_buffer_init("locate_qoff", &query_offsets);
let intervals_buf = ctx.create_buffer_empty("locate_intervals", num_queries * 16 * 2);
let search_params = SearchParams {
num_queries,
text_len,
num_blocks,
_pad: 0,
c: c_arr,
};
let search_params_buf = ctx.create_uniform_buffer("locate_search_params", &search_params);
let search_pipeline =
ctx.create_compute_pipeline("locate_search", LOCATE_SEARCH_SHADER, "locate_search");
let search_bg = ctx.create_bind_group(
&search_pipeline,
0,
&[
wgpu::BindGroupEntry {
binding: 0,
resource: qflat_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: qoff_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: chk_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: bv_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: intervals_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 5,
resource: search_params_buf.as_entire_binding(),
},
],
);
let wg_size: u32 = 64;
ctx.dispatch(
&search_pipeline,
&search_bg,
(num_queries.div_ceil(wg_size), 1, 1),
);
let intervals = ctx
.download_buffer(&intervals_buf, num_queries * 16 * 2)
.await;
let match_counts: Vec<u32> = (0..num_queries as usize)
.map(|q| {
(0..16usize)
.map(|i| {
let lo = intervals[q * 32 + i * 2];
let hi = intervals[q * 32 + i * 2 + 1];
hi.saturating_sub(lo)
})
.sum()
})
.collect();
let total_matches: u32 = match_counts.iter().sum();
if total_matches == 0 {
return Ok(vec![vec![]; queries.len()]);
}
let mut match_offsets: Vec<u32> = Vec::with_capacity(queries.len() + 1);
match_offsets.push(0);
for &c in &match_counts {
let prev = *match_offsets.last().unwrap();
match_offsets.push(prev + c);
}
let results_buf = ctx.create_buffer_empty("locate_results", total_matches * 2);
let match_off_buf = ctx.create_buffer_init("locate_match_offsets", &match_offsets);
let resolve_params = ResolveParams {
num_queries,
text_len,
num_blocks,
num_seqs,
sample_rate,
total_matches,
c: c_arr,
_pad2: [0, 0],
};
let resolve_params_buf = ctx.create_uniform_buffer("locate_resolve_params", &resolve_params);
let resolve_pipeline =
ctx.create_compute_pipeline("locate_resolve", LOCATE_RESOLVE_SHADER, "locate_resolve");
let resolve_bg = ctx.create_bind_group(
&resolve_pipeline,
0,
&[
wgpu::BindGroupEntry {
binding: 0,
resource: bwt_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: chk_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: bv_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: sa_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: seqb_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 5,
resource: intervals_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 6,
resource: match_off_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 7,
resource: results_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 8,
resource: resolve_params_buf.as_entire_binding(),
},
],
);
ctx.dispatch(
&resolve_pipeline,
&resolve_bg,
(total_matches.div_ceil(wg_size), 1, 1),
);
let results_flat = ctx.download_buffer(&results_buf, total_matches * 2).await;
let mut output: Vec<Vec<(u32, u32)>> = vec![vec![]; queries.len()];
for (q, &count) in match_counts.iter().enumerate() {
let off = match_offsets[q] as usize;
let mut hits = Vec::with_capacity(count as usize);
for k in 0..count as usize {
let seq_id = results_flat[(off + k) * 2];
let pos = results_flat[(off + k) * 2 + 1];
hits.push((seq_id, pos));
}
output[q] = hits;
}
Ok(output)
}