CI / test (push) Failing after 2s
The GPU path worked but was effectively hidden. cudarc's build script shells out
to `nvcc`, which ships in /usr/local/cuda/bin — a directory the reference host
had installed but never exported to the login shell, so `--features
embeddings-cuda` failed with a bare "`nvcc --version` failed" panic from a
dependency's build script, and the runtime fallback then reported only
"Embedder: CPU (...)" before spending hours on work a GPU does in minutes.
Two changes, both about making the failure legible rather than changing what the
code does:
- The CPU fallback now says why it fell back and what that costs, with the
concrete fix. A run that silently takes two orders of magnitude longer reads
as a hang, not as a configuration choice.
- BENCHMARKS.md states the build-time nvcc requirement, where the toolkit
actually installs, and that a shell file read non-interactively is the place
to export it — `~/.zshenv` rather than `~/.zshrc`, because build scripts do
not run in an interactive shell.
Host-side, the reference machine's CUDA exports lived in ~/.bashrc below its
non-interactive guard while the login shell is zsh, so they never applied to
anything. Moved to ~/.zshenv with duplicate-prepend guards; `nvcc --version`
and `cargo build --features embeddings-cuda` now both work over a plain
non-interactive ssh with no manual export.
171 lines
7.0 KiB
Rust
171 lines
7.0 KiB
Rust
//! Optional MiniLM sentence embedder for the LongMemEval bench.
|
|
//!
|
|
//! Compiled only under the `embeddings` feature, so the default build of a
|
|
//! project that prides itself on having no heavyweight dependencies stays
|
|
//! exactly as it was. Without it the bench runs BM25-only, as it always has.
|
|
//!
|
|
//! Loads `sentence-transformers/all-MiniLM-L6-v2` — the same checkpoint
|
|
//! omni-cortex uses — and produces 384-d mean-pooled, L2-normalised sentence
|
|
//! embeddings, which is the published recipe for this model (mean over token
|
|
//! states weighted by the attention mask, *not* the `[CLS]` pooler output).
|
|
|
|
use std::collections::HashMap;
|
|
use std::path::Path;
|
|
|
|
use candle_core::{DType, Device, Tensor};
|
|
use candle_nn::VarBuilder;
|
|
use candle_transformers::models::bert::{BertModel, Config, HiddenAct};
|
|
use tokenizers::Tokenizer;
|
|
|
|
/// Sequences encoded per forward pass. Larger batches amortise the transformer
|
|
/// call; 64 keeps peak memory modest while still saturating a CPU.
|
|
const BATCH: usize = 64;
|
|
|
|
/// A loaded MiniLM encoder.
|
|
pub struct Embedder {
|
|
model: BertModel,
|
|
tokenizer: Tokenizer,
|
|
device: Device,
|
|
}
|
|
|
|
impl Embedder {
|
|
/// Load from a directory holding `model.safetensors` and `tokenizer.json`.
|
|
///
|
|
/// `config.json` is read when present; otherwise the published MiniLM-L6-v2
|
|
/// architecture constants are used, which are pinned rather than guessed.
|
|
pub fn load(dir: &Path) -> Result<Self, Box<dyn std::error::Error>> {
|
|
// CUDA when the feature is on and a device is actually present; the CPU
|
|
// path is correct but roughly two orders of magnitude slower, which is
|
|
// the difference between minutes and most of a day on the full haystack.
|
|
let device = match Device::new_cuda(0) {
|
|
Ok(d) => {
|
|
eprintln!("Embedder: CUDA device 0");
|
|
d
|
|
}
|
|
Err(e) => {
|
|
// Loud, because the CPU path is correct but ~100x slower: the
|
|
// full longmemeval_s haystack is minutes on a GPU and most of a
|
|
// day on 8 cores. Silently falling back looks like a hang.
|
|
eprintln!("Embedder: CPU — CUDA unavailable ({e})");
|
|
eprintln!(
|
|
" WARNING: CPU embedding is roughly two orders of magnitude slower.\n Expect minutes for longmemeval_oracle and many hours for the full\n longmemeval_s haystack. For the GPU path, rebuild with\n `--features embeddings-cuda` and make sure `nvcc` is on PATH\n (it ships in /usr/local/cuda/bin, which is often not exported)."
|
|
);
|
|
Device::Cpu
|
|
}
|
|
};
|
|
let weights = dir.join("model.safetensors");
|
|
let tok_path = dir.join("tokenizer.json");
|
|
|
|
let config: Config = match std::fs::read_to_string(dir.join("config.json")) {
|
|
Ok(raw) => serde_json::from_str(&raw)?,
|
|
Err(_) => Config {
|
|
vocab_size: 30_522,
|
|
hidden_size: 384,
|
|
num_hidden_layers: 6,
|
|
num_attention_heads: 12,
|
|
intermediate_size: 1_536,
|
|
hidden_act: HiddenAct::Gelu,
|
|
hidden_dropout_prob: 0.0,
|
|
max_position_embeddings: 512,
|
|
type_vocab_size: 2,
|
|
initializer_range: 0.02,
|
|
layer_norm_eps: 1e-12,
|
|
pad_token_id: 0,
|
|
position_embedding_type: Default::default(),
|
|
use_cache: false,
|
|
classifier_dropout: None,
|
|
model_type: None,
|
|
},
|
|
};
|
|
|
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights], DType::F32, &device)? };
|
|
let model = BertModel::load(vb, &config)?;
|
|
let tokenizer = Tokenizer::from_file(&tok_path).map_err(|e| e.to_string())?;
|
|
|
|
Ok(Self {
|
|
model,
|
|
tokenizer,
|
|
device,
|
|
})
|
|
}
|
|
|
|
/// Encode `texts` into 384-d unit vectors, in order.
|
|
fn encode_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, Box<dyn std::error::Error>> {
|
|
let mut tk = self.tokenizer.clone();
|
|
let tk = tk
|
|
.with_padding(Some(tokenizers::PaddingParams::default()))
|
|
.with_truncation(Some(tokenizers::TruncationParams {
|
|
max_length: 512,
|
|
..Default::default()
|
|
}))
|
|
.map_err(|e| e.to_string())?;
|
|
let encodings = tk
|
|
.encode_batch(texts.to_vec(), true)
|
|
.map_err(|e| e.to_string())?;
|
|
|
|
let ids: Vec<u32> = encodings
|
|
.iter()
|
|
.flat_map(|e| e.get_ids().to_vec())
|
|
.collect();
|
|
let mask: Vec<u32> = encodings
|
|
.iter()
|
|
.flat_map(|e| e.get_attention_mask().to_vec())
|
|
.collect();
|
|
let (b, l) = (encodings.len(), encodings[0].get_ids().len());
|
|
|
|
let ids = Tensor::from_vec(ids, (b, l), &self.device)?;
|
|
let mask = Tensor::from_vec(mask, (b, l), &self.device)?;
|
|
let type_ids = ids.zeros_like()?;
|
|
|
|
let hidden = self.model.forward(&ids, &type_ids, Some(&mask))?;
|
|
|
|
// Mean-pool over real tokens only: sum(hidden * mask) / sum(mask).
|
|
let mask_f = mask.to_dtype(DType::F32)?.unsqueeze(2)?;
|
|
let summed = hidden.broadcast_mul(&mask_f)?.sum(1)?;
|
|
let counts = mask_f.sum(1)?.clamp(1e-9, f32::INFINITY)?;
|
|
let pooled = summed.broadcast_div(&counts)?;
|
|
|
|
// L2-normalise so cosine similarity is a plain dot product.
|
|
let norm = pooled
|
|
.sqr()?
|
|
.sum_keepdim(1)?
|
|
.sqrt()?
|
|
.clamp(1e-12, f32::INFINITY)?;
|
|
let normed = pooled.broadcast_div(&norm)?;
|
|
|
|
Ok(normed.to_vec2::<f32>()?)
|
|
}
|
|
|
|
/// Encode every distinct string in `texts` once, returning a lookup map.
|
|
///
|
|
/// LongMemEval's haystack sessions are drawn from a shared pool, so the same
|
|
/// turn text recurs across many questions. Deduplicating before encoding is
|
|
/// the difference between encoding the corpus once and encoding it per
|
|
/// question.
|
|
pub fn encode_unique(
|
|
&self,
|
|
texts: impl IntoIterator<Item = String>,
|
|
) -> Result<HashMap<String, Vec<f32>>, Box<dyn std::error::Error>> {
|
|
let mut unique: Vec<String> = texts.into_iter().collect();
|
|
unique.sort_unstable();
|
|
unique.dedup();
|
|
|
|
let total = unique.len();
|
|
eprintln!("Embedding {total} unique texts with MiniLM (batch {BATCH})...");
|
|
|
|
let mut out = HashMap::with_capacity(total);
|
|
for (n, chunk) in unique.chunks(BATCH).enumerate() {
|
|
let refs: Vec<&str> = chunk.iter().map(String::as_str).collect();
|
|
let vecs = self.encode_batch(&refs)?;
|
|
for (text, v) in chunk.iter().zip(vecs) {
|
|
out.insert(text.clone(), v);
|
|
}
|
|
if n % 50 == 0 {
|
|
eprint!("\r [{}/{}] embedded...", (n * BATCH).min(total), total);
|
|
}
|
|
}
|
|
eprintln!("\r [{total}/{total}] embedded. ");
|
|
Ok(out)
|
|
}
|
|
}
|