Skip to main content

context_search/
context_search.rs

1//! Reuse an encoder's context for corpus loading and resident query tensors.
2use hrx::{
3    Access,
4    execution::{GpuAccess, RuntimeOptions},
5    inference::ModelContext,
6    residency::ResidencyManager,
7    tensor::{DType, Layout, TensorDesc},
8};
9use hrxdb::Corpus;
10
11fn main() -> hrxdb::Result<()> {
12    let manager = ResidencyManager::new(64 * 1024 * 1024)?;
13    let context = ModelContext::new(RuntimeOptions {
14        memory_budget: Some(manager.budget()),
15        ..Default::default()
16    })?;
17    let rows = [[0, 0x3c, 0, 0, 0, 0], [0, 0, 0, 0x3c, 0, 0]];
18    let corpus = Corpus::load_resident_fp16_in(&context, "example-v1", 3, rows)?;
19    let search = corpus.prepare_search(&context, 1, 2, 1)?;
20
21    // An encoder loaded in `context` can supply its output tensor here directly.
22    let query = context.upload(
23        TensorDesc::new(DType::F32, vec![1, 3])?.with_layout(Layout::Rows)?,
24        &[1f32, 0., 0.]
25            .into_iter()
26            .flat_map(f32::to_le_bytes)
27            .collect::<Vec<_>>(),
28    )?;
29    let results = search.submit(&query)?.read()?;
30    assert_eq!(results[0][0].id, 0);
31    println!("{results:?}");
32
33    // A native kernel can read corpus buffers and write a context tensor in
34    // one dispatch. Here gather's kernel produces a query for resident search.
35    let desc = TensorDesc::new(DType::F32, vec![1, 3])?.with_layout(Layout::Rows)?;
36    let output = context.allocate(desc.clone())?;
37    let binding = output.binding().unwrap();
38    let retained = (*corpus).clone();
39    let mut stream = retained.stream()?;
40    let mut graph = context.runtime().graph();
41    // SAFETY: the callback retains its stream and immutable corpus, writes the
42    // entire declared output, and drains the stream before returning. No scoped
43    // view escapes. Other corpus users only read the shared native storage.
44    unsafe {
45        graph.gpu_scoped(
46            &[GpuAccess {
47                view: binding.clone(),
48                access: Access::Write,
49            }],
50            move |views| {
51                let result = retained.gather_into(&mut stream, &[1], views[0]);
52                stream.synchronize()?;
53                result?;
54                Ok(())
55            },
56        )?;
57    }
58    let producer = graph.prepare()?.submit()?;
59    let gathered = context.tensor(desc, binding, producer)?;
60    assert_eq!(search.submit(&gathered)?.read()?[0][0].id, 1);
61    Ok(())
62}