496 lines
18 KiB
Rust
496 lines
18 KiB
Rust
//! # Standalone ColBERT Test
|
|
//!
|
|
//! This is a standalone test to validate the ColBERT implementation
|
|
//! without depending on the full RTX ecosystem.
|
|
|
|
use std::collections::HashMap;
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub enum SimilarityMetric {
|
|
Cosine,
|
|
DotProduct,
|
|
L2,
|
|
}
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub struct TokenEmbedding {
|
|
pub token: String,
|
|
pub vector: Vec<f32>,
|
|
pub position: usize,
|
|
}
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub struct MaxSimResult {
|
|
pub score: f32,
|
|
pub query_token_contributions: HashMap<usize, f32>,
|
|
pub best_matches: HashMap<usize, (usize, f32)>,
|
|
}
|
|
|
|
impl MaxSimResult {
|
|
pub fn compute(
|
|
query_embeddings: &[TokenEmbedding],
|
|
doc_embeddings: &[TokenEmbedding],
|
|
similarity_metric: SimilarityMetric,
|
|
) -> Result<Self, String> {
|
|
if query_embeddings.is_empty() || doc_embeddings.is_empty() {
|
|
return Err("Query and document embeddings cannot be empty".to_string());
|
|
}
|
|
|
|
let mut query_token_contributions = HashMap::new();
|
|
let mut best_matches = HashMap::new();
|
|
let mut total_score = 0.0;
|
|
|
|
// For each query token, find maximum similarity with any document token
|
|
for (q_idx, query_token) in query_embeddings.iter().enumerate() {
|
|
let mut max_sim = f32::NEG_INFINITY;
|
|
let mut best_doc_idx = 0;
|
|
|
|
for (d_idx, doc_token) in doc_embeddings.iter().enumerate() {
|
|
let similarity = calculate_similarity(
|
|
&query_token.vector,
|
|
&doc_token.vector,
|
|
&similarity_metric,
|
|
);
|
|
|
|
if similarity > max_sim {
|
|
max_sim = similarity;
|
|
best_doc_idx = d_idx;
|
|
}
|
|
}
|
|
|
|
query_token_contributions.insert(q_idx, max_sim);
|
|
best_matches.insert(q_idx, (best_doc_idx, max_sim));
|
|
total_score += max_sim;
|
|
}
|
|
|
|
// Average score across query tokens
|
|
let final_score = total_score / query_embeddings.len() as f32;
|
|
|
|
Ok(MaxSimResult {
|
|
score: final_score,
|
|
query_token_contributions,
|
|
best_matches,
|
|
})
|
|
}
|
|
}
|
|
|
|
fn calculate_similarity(vec1: &[f32], vec2: &[f32], metric: &SimilarityMetric) -> f32 {
|
|
match metric {
|
|
SimilarityMetric::Cosine => {
|
|
let dot_product: f32 = vec1.iter().zip(vec2.iter()).map(|(a, b)| a * b).sum();
|
|
let norm1: f32 = vec1.iter().map(|x| x * x).sum::<f32>().sqrt();
|
|
let norm2: f32 = vec2.iter().map(|x| x * x).sum::<f32>().sqrt();
|
|
|
|
if norm1 == 0.0 || norm2 == 0.0 {
|
|
0.0
|
|
} else {
|
|
dot_product / (norm1 * norm2)
|
|
}
|
|
}
|
|
SimilarityMetric::DotProduct => {
|
|
vec1.iter().zip(vec2.iter()).map(|(a, b)| a * b).sum()
|
|
}
|
|
SimilarityMetric::L2 => {
|
|
let squared_diff: f32 = vec1.iter().zip(vec2.iter()).map(|(a, b)| (a - b).powi(2)).sum();
|
|
1.0 / (1.0 + squared_diff.sqrt()) // Convert distance to similarity
|
|
}
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_token_embedding_creation() {
|
|
let token = TokenEmbedding {
|
|
token: "machine".to_string(),
|
|
vector: vec![0.1, 0.2, 0.3, 0.4],
|
|
position: 0,
|
|
};
|
|
|
|
assert_eq!(token.token, "machine");
|
|
assert_eq!(token.vector.len(), 4);
|
|
assert_eq!(token.position, 0);
|
|
assert_eq!(token.vector[0], 0.1);
|
|
}
|
|
|
|
#[test]
|
|
fn test_cosine_similarity() {
|
|
let vec1 = vec![1.0, 0.0, 0.0];
|
|
let vec2 = vec![1.0, 0.0, 0.0]; // Identical
|
|
let vec3 = vec![0.0, 1.0, 0.0]; // Orthogonal
|
|
|
|
let sim_identical = calculate_similarity(&vec1, &vec2, &SimilarityMetric::Cosine);
|
|
let sim_orthogonal = calculate_similarity(&vec1, &vec3, &SimilarityMetric::Cosine);
|
|
|
|
assert!((sim_identical - 1.0).abs() < 1e-6, "Identical vectors should have similarity 1.0");
|
|
assert!(sim_orthogonal.abs() < 1e-6, "Orthogonal vectors should have similarity 0.0");
|
|
}
|
|
|
|
#[test]
|
|
fn test_maxsim_computation() {
|
|
let query_embeddings = vec![
|
|
TokenEmbedding {
|
|
token: "machine".to_string(),
|
|
vector: vec![1.0, 0.0, 0.0],
|
|
position: 0,
|
|
},
|
|
TokenEmbedding {
|
|
token: "learning".to_string(),
|
|
vector: vec![0.0, 1.0, 0.0],
|
|
position: 1,
|
|
},
|
|
];
|
|
|
|
let doc_embeddings = vec![
|
|
TokenEmbedding {
|
|
token: "artificial".to_string(),
|
|
vector: vec![0.9, 0.1, 0.0], // Similar to "machine"
|
|
position: 0,
|
|
},
|
|
TokenEmbedding {
|
|
token: "intelligence".to_string(),
|
|
vector: vec![0.1, 0.9, 0.0], // Similar to "learning"
|
|
position: 1,
|
|
},
|
|
TokenEmbedding {
|
|
token: "algorithms".to_string(),
|
|
vector: vec![0.5, 0.5, 0.0],
|
|
position: 2,
|
|
},
|
|
];
|
|
|
|
let maxsim = MaxSimResult::compute(
|
|
&query_embeddings,
|
|
&doc_embeddings,
|
|
SimilarityMetric::Cosine,
|
|
).expect("MaxSim computation should succeed");
|
|
|
|
// Check basic properties
|
|
assert!(maxsim.score > 0.0, "MaxSim score should be positive");
|
|
assert!(maxsim.score <= 1.0, "MaxSim score should be <= 1.0 for cosine similarity");
|
|
assert_eq!(maxsim.query_token_contributions.len(), 2, "Should have contributions for each query token");
|
|
assert_eq!(maxsim.best_matches.len(), 2, "Should have best matches for each query token");
|
|
|
|
// Check that each query token has a contribution
|
|
assert!(maxsim.query_token_contributions.contains_key(&0), "Should have contribution for query token 0");
|
|
assert!(maxsim.query_token_contributions.contains_key(&1), "Should have contribution for query token 1");
|
|
|
|
// Check that each query token has a best match
|
|
assert!(maxsim.best_matches.contains_key(&0), "Should have best match for query token 0");
|
|
assert!(maxsim.best_matches.contains_key(&1), "Should have best match for query token 1");
|
|
|
|
println!("MaxSim score: {:.4}", maxsim.score);
|
|
println!("Query token contributions: {:?}", maxsim.query_token_contributions);
|
|
println!("Best matches: {:?}", maxsim.best_matches);
|
|
}
|
|
|
|
#[test]
|
|
fn test_maxsim_different_metrics() {
|
|
let query_embeddings = vec![
|
|
TokenEmbedding {
|
|
token: "test".to_string(),
|
|
vector: vec![1.0, 2.0],
|
|
position: 0,
|
|
},
|
|
];
|
|
|
|
let doc_embeddings = vec![
|
|
TokenEmbedding {
|
|
token: "example".to_string(),
|
|
vector: vec![2.0, 4.0], // Scaled version
|
|
position: 0,
|
|
},
|
|
];
|
|
|
|
// Test different similarity metrics
|
|
let cosine_result = MaxSimResult::compute(
|
|
&query_embeddings,
|
|
&doc_embeddings,
|
|
SimilarityMetric::Cosine,
|
|
).expect("Cosine MaxSim should succeed");
|
|
|
|
let dot_result = MaxSimResult::compute(
|
|
&query_embeddings,
|
|
&doc_embeddings,
|
|
SimilarityMetric::DotProduct,
|
|
).expect("Dot product MaxSim should succeed");
|
|
|
|
let l2_result = MaxSimResult::compute(
|
|
&query_embeddings,
|
|
&doc_embeddings,
|
|
SimilarityMetric::L2,
|
|
).expect("L2 MaxSim should succeed");
|
|
|
|
// Cosine similarity should be 1.0 for scaled vectors
|
|
assert!((cosine_result.score - 1.0).abs() < 1e-6, "Cosine similarity should be 1.0 for scaled vectors");
|
|
|
|
// Dot product should be higher than cosine for scaled vectors
|
|
assert!(dot_result.score > cosine_result.score, "Dot product should be larger for scaled vectors");
|
|
|
|
// L2 should produce a valid similarity
|
|
assert!(l2_result.score > 0.0 && l2_result.score <= 1.0, "L2 similarity should be in (0, 1]");
|
|
|
|
println!("Cosine: {:.4}, Dot: {:.4}, L2: {:.4}",
|
|
cosine_result.score, dot_result.score, l2_result.score);
|
|
}
|
|
|
|
#[test]
|
|
fn test_maxsim_error_cases() {
|
|
let empty_query: Vec<TokenEmbedding> = vec![];
|
|
let empty_doc: Vec<TokenEmbedding> = vec![];
|
|
|
|
let query_embeddings = vec![
|
|
TokenEmbedding {
|
|
token: "test".to_string(),
|
|
vector: vec![1.0, 0.0],
|
|
position: 0,
|
|
},
|
|
];
|
|
|
|
// Test empty inputs
|
|
let result = MaxSimResult::compute(&empty_query, &empty_doc, SimilarityMetric::Cosine);
|
|
assert!(result.is_err(), "Should fail with empty inputs");
|
|
|
|
let result = MaxSimResult::compute(&query_embeddings, &empty_doc, SimilarityMetric::Cosine);
|
|
assert!(result.is_err(), "Should fail with empty document embeddings");
|
|
|
|
let result = MaxSimResult::compute(&empty_query, &query_embeddings, SimilarityMetric::Cosine);
|
|
assert!(result.is_err(), "Should fail with empty query embeddings");
|
|
}
|
|
|
|
#[test]
|
|
fn test_realistic_colbert_scenario() {
|
|
// Simulate a realistic ColBERT scenario
|
|
let query = "What is machine learning?";
|
|
let query_tokens = vec!["[CLS]", "what", "is", "machine", "learning", "[MASK]", "[MASK]", "[SEP]"];
|
|
|
|
let document = "Machine learning is a method of data analysis that automates analytical model building.";
|
|
let doc_tokens = vec!["[CLS]", "machine", "learning", "is", "a", "method", "of", "data", "analysis", "[SEP]"];
|
|
|
|
// Create mock embeddings (in reality these would come from BERT)
|
|
let query_embeddings: Vec<TokenEmbedding> = query_tokens.iter().enumerate().map(|(i, &token)| {
|
|
TokenEmbedding {
|
|
token: token.to_string(),
|
|
vector: create_mock_embedding(token, 128),
|
|
position: i,
|
|
}
|
|
}).collect();
|
|
|
|
let doc_embeddings: Vec<TokenEmbedding> = doc_tokens.iter().enumerate().map(|(i, &token)| {
|
|
TokenEmbedding {
|
|
token: token.to_string(),
|
|
vector: create_mock_embedding(token, 128),
|
|
position: i,
|
|
}
|
|
}).collect();
|
|
|
|
let maxsim = MaxSimResult::compute(
|
|
&query_embeddings,
|
|
&doc_embeddings,
|
|
SimilarityMetric::Cosine,
|
|
).expect("MaxSim computation should succeed");
|
|
|
|
// Verify expected properties
|
|
assert!(maxsim.score > 0.0, "Should have positive similarity for relevant document");
|
|
assert_eq!(maxsim.query_token_contributions.len(), query_embeddings.len());
|
|
assert_eq!(maxsim.best_matches.len(), query_embeddings.len());
|
|
|
|
// The tokens "machine" and "learning" should have high similarity
|
|
// Find the indices of these tokens in the query
|
|
let machine_idx = query_tokens.iter().position(|&t| t == "machine").unwrap();
|
|
let learning_idx = query_tokens.iter().position(|&t| t == "learning").unwrap();
|
|
|
|
let machine_contribution = maxsim.query_token_contributions[&machine_idx];
|
|
let learning_contribution = maxsim.query_token_contributions[&learning_idx];
|
|
|
|
println!("Machine token contribution: {:.4}", machine_contribution);
|
|
println!("Learning token contribution: {:.4}", learning_contribution);
|
|
println!("Overall MaxSim score: {:.4}", maxsim.score);
|
|
|
|
// These should be relatively high since the document contains the same tokens
|
|
assert!(machine_contribution > 0.5, "Machine token should have high similarity");
|
|
assert!(learning_contribution > 0.5, "Learning token should have high similarity");
|
|
}
|
|
|
|
// Helper function to create mock embeddings
|
|
fn create_mock_embedding(token: &str, dim: usize) -> Vec<f32> {
|
|
use std::collections::hash_map::DefaultHasher;
|
|
use std::hash::{Hash, Hasher};
|
|
|
|
let mut hasher = DefaultHasher::new();
|
|
token.hash(&mut hasher);
|
|
let hash = hasher.finish();
|
|
|
|
let mut embedding = Vec::with_capacity(dim);
|
|
for i in 0..dim {
|
|
let mut h = DefaultHasher::new();
|
|
(hash.wrapping_add(i as u64)).hash(&mut h);
|
|
let val = (h.finish() as f32) / (u64::MAX as f32);
|
|
embedding.push((val - 0.5) * 2.0); // Normalize to [-1, 1]
|
|
}
|
|
|
|
// L2 normalize for cosine similarity
|
|
let norm: f32 = embedding.iter().map(|x| x * x).sum::<f32>().sqrt();
|
|
if norm > 0.0 {
|
|
for val in &mut embedding {
|
|
*val /= norm;
|
|
}
|
|
}
|
|
|
|
embedding
|
|
}
|
|
}
|
|
|
|
fn main() {
|
|
println!("ColBERT Standalone Test");
|
|
println!("====================");
|
|
|
|
// Run a simple demonstration
|
|
let query_embeddings = vec![
|
|
TokenEmbedding {
|
|
token: "machine".to_string(),
|
|
vector: vec![1.0, 0.0, 0.0],
|
|
position: 0,
|
|
},
|
|
TokenEmbedding {
|
|
token: "learning".to_string(),
|
|
vector: vec![0.0, 1.0, 0.0],
|
|
position: 1,
|
|
},
|
|
];
|
|
|
|
let doc_embeddings = vec![
|
|
TokenEmbedding {
|
|
token: "artificial".to_string(),
|
|
vector: vec![0.8, 0.2, 0.0],
|
|
position: 0,
|
|
},
|
|
TokenEmbedding {
|
|
token: "intelligence".to_string(),
|
|
vector: vec![0.2, 0.8, 0.0],
|
|
position: 1,
|
|
},
|
|
];
|
|
|
|
match MaxSimResult::compute(&query_embeddings, &doc_embeddings, SimilarityMetric::Cosine) {
|
|
Ok(maxsim) => {
|
|
println!("MaxSim Score: {:.4}", maxsim.score);
|
|
println!("Query Token Contributions:");
|
|
for (idx, contribution) in &maxsim.query_token_contributions {
|
|
let token = &query_embeddings[*idx].token;
|
|
println!(" {}: {:.4}", token, contribution);
|
|
}
|
|
println!("Best Matches:");
|
|
for (q_idx, (d_idx, score)) in &maxsim.best_matches {
|
|
let q_token = &query_embeddings[*q_idx].token;
|
|
let d_token = &doc_embeddings[*d_idx].token;
|
|
println!(" {} -> {}: {:.4}", q_token, d_token, score);
|
|
}
|
|
}
|
|
Err(e) => {
|
|
println!("Error: {}", e);
|
|
}
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod integration_tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_colbert_pipeline_simulation() {
|
|
// Simulate a complete ColBERT pipeline
|
|
let queries = vec![
|
|
"What is artificial intelligence?",
|
|
"How do neural networks work?",
|
|
"Machine learning algorithms"
|
|
];
|
|
|
|
let documents = vec![
|
|
"Artificial intelligence (AI) is intelligence demonstrated by machines.",
|
|
"Neural networks are computing systems inspired by biological neural networks.",
|
|
"Machine learning is a method of data analysis that automates model building."
|
|
];
|
|
|
|
for (q_idx, query) in queries.iter().enumerate() {
|
|
println!("Query {}: {}", q_idx + 1, query);
|
|
let mut doc_scores = Vec::new();
|
|
|
|
for (d_idx, document) in documents.iter().enumerate() {
|
|
// Mock tokenization and embedding
|
|
let query_tokens = tokenize_simple(query);
|
|
let doc_tokens = tokenize_simple(document);
|
|
|
|
let query_embeddings: Vec<TokenEmbedding> = query_tokens.iter().enumerate().map(|(i, token)| {
|
|
TokenEmbedding {
|
|
token: token.clone(),
|
|
vector: create_mock_embedding_simple(&token, 64),
|
|
position: i,
|
|
}
|
|
}).collect();
|
|
|
|
let doc_embeddings: Vec<TokenEmbedding> = doc_tokens.iter().enumerate().map(|(i, token)| {
|
|
TokenEmbedding {
|
|
token: token.clone(),
|
|
vector: create_mock_embedding_simple(&token, 64),
|
|
position: i,
|
|
}
|
|
}).collect();
|
|
|
|
if let Ok(maxsim) = MaxSimResult::compute(&query_embeddings, &doc_embeddings, SimilarityMetric::Cosine) {
|
|
doc_scores.push((d_idx, maxsim.score, document));
|
|
}
|
|
}
|
|
|
|
// Sort by score (descending)
|
|
doc_scores.sort_by(|a, b| b.1.total_cmp(&a.1));
|
|
|
|
println!(" Top results:");
|
|
for (rank, (doc_idx, score, doc)) in doc_scores.iter().take(2).enumerate() {
|
|
println!(" {}. Doc {}: {:.4} - {}", rank + 1, doc_idx + 1, score, doc);
|
|
}
|
|
println!();
|
|
}
|
|
}
|
|
|
|
fn tokenize_simple(text: &str) -> Vec<String> {
|
|
let mut tokens = vec!["[CLS]".to_string()];
|
|
for word in text.to_lowercase().split_whitespace() {
|
|
let cleaned = word.chars().filter(|c| c.is_alphanumeric()).collect::<String>();
|
|
if !cleaned.is_empty() {
|
|
tokens.push(cleaned);
|
|
}
|
|
}
|
|
tokens.push("[SEP]".to_string());
|
|
tokens
|
|
}
|
|
|
|
fn create_mock_embedding_simple(token: &str, dim: usize) -> Vec<f32> {
|
|
use std::collections::hash_map::DefaultHasher;
|
|
use std::hash::{Hash, Hasher};
|
|
|
|
let mut hasher = DefaultHasher::new();
|
|
token.hash(&mut hasher);
|
|
let hash = hasher.finish();
|
|
|
|
let mut embedding = Vec::with_capacity(dim);
|
|
for i in 0..dim {
|
|
let mut h = DefaultHasher::new();
|
|
(hash + i as u64).hash(&mut h);
|
|
let val = (h.finish() as f32) / (u64::MAX as f32);
|
|
embedding.push((val - 0.5) * 2.0);
|
|
}
|
|
|
|
// L2 normalize
|
|
let norm: f32 = embedding.iter().map(|x| x * x).sum::<f32>().sqrt();
|
|
if norm > 0.0 {
|
|
for val in &mut embedding {
|
|
*val /= norm;
|
|
}
|
|
}
|
|
|
|
embedding
|
|
}
|
|
} |