From 96b0a635988f3dc923a1b45ad44e5993d1e5785a Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Sun, 12 Apr 2026 21:29:38 -0500 Subject: [PATCH 01/34] working! --- .claude/settings.local.json | 16 + Cargo.lock | 20 + Cargo.toml | 2 + fusor-ml/core/src/device.rs | 5 +- fusor-ml/core/src/matmul/mod.rs | 6 +- fusor-ml/core/src/quantized/matmul/mod.rs | 2 +- .../core/src/quantized/matmul/sgemv/mod.rs | 30 +- fusor-ml/fusor/src/composite/activations.rs | 20 +- fusor-ml/fusor/src/layers/linear.rs | 38 +- fusor-ml/fusor/src/lib.rs | 36 +- fusor-ml/fusor/src/quantized.rs | 83 + models/rbert/src/language_model.rs | 28 +- models/rbert/src/lib.rs | 252 ++- models/rbert/src/raw/attention.rs | 24 + models/rbert/src/raw/encoder.rs | 41 + models/rbert/src/raw/layer.rs | 27 + models/rbert/src/raw/mod.rs | 49 + models/rbert/src/raw/self_attention.rs | 77 +- models/rgliner/Cargo.toml | 33 + .../convert_to_gguf.cpython-311.pyc | Bin 0 -> 27125 bytes models/rgliner/convert_to_gguf.py | 475 ++++++ models/rgliner/examples/basic.rs | 46 + models/rgliner/src/config.rs | 143 ++ models/rgliner/src/decoding.rs | 255 +++ models/rgliner/src/error.rs | 43 + models/rgliner/src/lib.rs | 1467 +++++++++++++++++ models/rgliner/src/raw/label_encoder.rs | 308 ++++ models/rgliner/src/raw/mod.rs | 13 + .../rgliner/src/raw/modern_bert/attention.rs | 99 ++ models/rgliner/src/raw/modern_bert/config.rs | 101 ++ .../src/raw/modern_bert/feed_forward.rs | 45 + models/rgliner/src/raw/modern_bert/layer.rs | 77 + models/rgliner/src/raw/modern_bert/mod.rs | 16 + models/rgliner/src/raw/modern_bert/model.rs | 127 ++ models/rgliner/src/raw/scorer.rs | 62 + models/rgliner/src/raw/span_layer.rs | 221 +++ models/rgliner/src/raw/text_encoder.rs | 54 + models/rgliner/src/source.rs | 310 ++++ models/rgliner/src/tokenization.rs | 278 ++++ models/rgliner/tests/example_regression.rs | 86 + 40 files changed, 4957 insertions(+), 58 deletions(-) create mode 100644 .claude/settings.local.json create mode 100644 models/rgliner/Cargo.toml create mode 100644 models/rgliner/__pycache__/convert_to_gguf.cpython-311.pyc create mode 100644 models/rgliner/convert_to_gguf.py create mode 100644 models/rgliner/examples/basic.rs create mode 100644 models/rgliner/src/config.rs create mode 100644 models/rgliner/src/decoding.rs create mode 100644 models/rgliner/src/error.rs create mode 100644 models/rgliner/src/lib.rs create mode 100644 models/rgliner/src/raw/label_encoder.rs create mode 100644 models/rgliner/src/raw/mod.rs create mode 100644 models/rgliner/src/raw/modern_bert/attention.rs create mode 100644 models/rgliner/src/raw/modern_bert/config.rs create mode 100644 models/rgliner/src/raw/modern_bert/feed_forward.rs create mode 100644 models/rgliner/src/raw/modern_bert/layer.rs create mode 100644 models/rgliner/src/raw/modern_bert/mod.rs create mode 100644 models/rgliner/src/raw/modern_bert/model.rs create mode 100644 models/rgliner/src/raw/scorer.rs create mode 100644 models/rgliner/src/raw/span_layer.rs create mode 100644 models/rgliner/src/raw/text_encoder.rs create mode 100644 models/rgliner/src/source.rs create mode 100644 models/rgliner/src/tokenization.rs create mode 100644 models/rgliner/tests/example_regression.rs diff --git a/.claude/settings.local.json b/.claude/settings.local.json new file mode 100644 index 000000000..076b30c7c --- /dev/null +++ b/.claude/settings.local.json @@ -0,0 +1,16 @@ +{ + "permissions": { + "allow": [ + "Bash(GLINER_MODEL=/tmp/gliner-edge.gguf cargo run:*)", + "Bash(GLINER_MODEL=/private/tmp/gliner-edge.gguf cargo run:*)", + "Bash(RUST_BACKTRACE=1 GLINER_MODEL=/private/tmp/gliner-edge.gguf cargo run:*)", + "Bash(RUST_BACKTRACE=full GLINER_MODEL=./models/rgliner/weights/gliner-edge.gguf cargo run:*)", + "Bash(GLINER_MODEL=./models/rgliner/weights/gliner-edge.gguf cargo run:*)", + "Bash(git -C /Users/evanalmloff/Desktop/Github/ner diff HEAD -- fusor-ml/core/src/matmul/mod.rs)", + "Bash(GLINER_MODEL=/Users/evanalmloff/Desktop/Github/ner/models/rgliner/weights/gliner-edge.gguf cargo run:*)", + "Bash(RUST_BACKTRACE=1 GLINER_MODEL=/Users/evanalmloff/Desktop/Github/ner/models/rgliner/weights/gliner-edge.gguf cargo run:*)", + "Bash(git -C /Users/evanalmloff/Desktop/Github/ner checkout -- fusor-ml/core/src/matmul/mod.rs)", + "WebFetch(domain:api.github.com)" + ] + } +} diff --git a/Cargo.lock b/Cargo.lock index 906a2da05..dff86fcb6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7670,6 +7670,26 @@ version = "0.8.53" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "47b34b781b31e5d73e9fbc8689c70551fd1ade9a19e3e28cfec8580a79290cc4" +[[package]] +name = "rgliner" +version = "0.4.0" +dependencies = [ + "anyhow", + "fusor", + "fusor-core", + "fusor-gguf", + "kalosm-common", + "kalosm-language-model", + "kalosm-model-types", + "rbert", + "serde", + "serde_json", + "thiserror 2.0.18", + "tokenizers", + "tokio", + "tracing", +] + [[package]] name = "ring" version = "0.17.14" diff --git a/Cargo.toml b/Cargo.toml index e5c8aa78e..208e09fb4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,6 +13,7 @@ bench = false [workspace] members = [ "models/rbert", + "models/rgliner", "models/kalosm-llama", "models/rwhisper", "models/rwuerstchen", @@ -55,6 +56,7 @@ kalosm-vision = { path = "./interfaces/kalosm-vision", version = "0.4.0" } kalosm-learning = { path = "./interfaces/kalosm-learning", version = "0.4.0" } kalosm-learning-macro = { path = "./interfaces/kalosm-learning-macro", version = "0.4.0" } rbert = { path = "./models/rbert", version = "0.4.0" } +rgliner = { path = "./models/rgliner", version = "0.4.0" } kalosm-llama = { path = "./models/kalosm-llama", version = "0.4.0" } rwhisper = { path = "./models/rwhisper", version = "0.4.0" } rwuerstchen = { path = "./models/rwuerstchen", version = "0.4.0" } diff --git a/fusor-ml/core/src/device.rs b/fusor-ml/core/src/device.rs index aaa89601c..aac43b8b5 100644 --- a/fusor-ml/core/src/device.rs +++ b/fusor-ml/core/src/device.rs @@ -86,6 +86,9 @@ pub struct Device { impl Device { pub async fn new() -> Result { + let disable_shader_f16 = std::env::var_os("FUSOR_DISABLE_SHADER_F16") + .map(|value| value != "0") + .unwrap_or(false); let dx_compiler = wgpu::Dx12Compiler::from_env().unwrap_or(wgpu::Dx12Compiler::StaticDxc); let backends = wgpu::Backends::from_env().unwrap_or(wgpu::Backends::all()); let instance = wgpu::Instance::new(wgpu::InstanceDescriptor { @@ -106,7 +109,7 @@ impl Device { if adapter.features().contains(wgpu::Features::SUBGROUP) { required_features |= wgpu::Features::SUBGROUP; } - if adapter.features().contains(wgpu::Features::SHADER_F16) { + if !disable_shader_f16 && adapter.features().contains(wgpu::Features::SHADER_F16) { required_features |= wgpu::Features::SHADER_F16; } let (device, queue) = adapter diff --git a/fusor-ml/core/src/matmul/mod.rs b/fusor-ml/core/src/matmul/mod.rs index 638d619f2..cf58fc806 100644 --- a/fusor-ml/core/src/matmul/mod.rs +++ b/fusor-ml/core/src/matmul/mod.rs @@ -76,6 +76,10 @@ impl MatMulOperation { let last_dim = first_shape.len() - 1; let second_to_last_dim = first_shape.len() - 2; let mut out_shape = first_shape.to_vec(); + // Handle broadcasting for batch dimensions + for i in 0..second_to_last_dim { + out_shape[i] = std::cmp::max(first_shape[i], second_shape[i]); + } out_shape[second_to_last_dim] = first_shape[second_to_last_dim]; out_shape[last_dim] = second_shape[last_dim]; assert_eq!(first_shape[last_dim], second_shape[second_to_last_dim]); @@ -85,7 +89,7 @@ impl MatMulOperation { .rev() .skip(2) .zip(second_shape.iter().rev().skip(2)) - .all(|(a, b)| a == b) + .all(|(a, b)| *a == *b || *a == 1 || *b == 1) ); Self { diff --git a/fusor-ml/core/src/quantized/matmul/mod.rs b/fusor-ml/core/src/quantized/matmul/mod.rs index ab2a2f627..de3e91e0e 100644 --- a/fusor-ml/core/src/quantized/matmul/mod.rs +++ b/fusor-ml/core/src/quantized/matmul/mod.rs @@ -808,7 +808,7 @@ impl Operation for QMatMulOperation { .product(); if self.sgemv() { - sgemv::dispatch_size(&self.matrix, n, m, batch_size) + sgemv::dispatch_size(self, n, m, batch_size, &self.matrix.device) } else { sgemm::dispatch_size(self, workgroup_shape, &self.matrix, n, m, batch_size) } diff --git a/fusor-ml/core/src/quantized/matmul/sgemv/mod.rs b/fusor-ml/core/src/quantized/matmul/sgemv/mod.rs index abc2b4866..ddc84c7fd 100644 --- a/fusor-ml/core/src/quantized/matmul/sgemv/mod.rs +++ b/fusor-ml/core/src/quantized/matmul/sgemv/mod.rs @@ -77,6 +77,20 @@ fn can_use_specialized_sgemv(device: &Device) -> bool { device.max_subgroup_size() >= 2 * device.min_subgroup_size() } +fn batch_size(op: &QMatMulOperation) -> u32 { + op.in_shape + .iter() + .rev() + .skip(2) + .map(|x| *x as u32) + .product::() + .max(1) +} + +fn use_specialized_sgemv(op: &QMatMulOperation, device: &Device) -> bool { + batch_size(op) == 1 && can_use_specialized_sgemv(device) +} + #[allow(clippy::too_many_arguments)] pub(crate) fn sgemv( op: &QMatMulOperation, @@ -92,7 +106,7 @@ pub(crate) fn sgemv( ) { let device = graph.device(); // Check if we can use specialized SGEMV (requires 2 subgroups per workgroup) - let use_specialized = can_use_specialized_sgemv(&device); + let use_specialized = use_specialized_sgemv(op, &device); match op.matrix.datatype { GgmlType::Q6K if use_specialized => q6k_sgemv( op, @@ -165,9 +179,9 @@ pub(crate) fn sgemv( } /// Calculate the number of N-dimension workgroups based on matrix type -pub(crate) fn n_workgroups(matrix: &QMatrix, n: u32) -> u32 { +pub(crate) fn n_workgroups(op: &QMatMulOperation, matrix: &QMatrix, n: u32, device: &Device) -> u32 { // Only use specialized dispatch sizes if we can use specialized SGEMV - if can_use_specialized_sgemv(&matrix.device) { + if use_specialized_sgemv(op, device) { if matrix.datatype == GgmlType::Q6K { n.div_ceil(Q6K_SGEMV_CHUNK_SIZE * 2) } else if matrix.datatype == GgmlType::Q4K { @@ -186,10 +200,16 @@ pub(crate) fn n_workgroups(matrix: &QMatrix, n: u32) -> u32 { } } -pub(crate) fn dispatch_size(matrix: &QMatrix, n: u32, m: u32, batch_size: u32) -> [u32; 3] { +pub(crate) fn dispatch_size( + op: &QMatMulOperation, + n: u32, + m: u32, + batch_size: u32, + device: &Device, +) -> [u32; 3] { // Calculate total workgroups: n_workgroups * m * batch // Use distribute_workgroups to spread across all 3 dimensions if needed - let n_wg = n_workgroups(matrix, n); + let n_wg = n_workgroups(op, &op.matrix, n, device); let total_workgroups = n_wg * m * batch_size; distribute_workgroups(total_workgroups) } diff --git a/fusor-ml/fusor/src/composite/activations.rs b/fusor-ml/fusor/src/composite/activations.rs index d5b9c902b..6cd16578e 100644 --- a/fusor-ml/fusor/src/composite/activations.rs +++ b/fusor-ml/fusor/src/composite/activations.rs @@ -50,15 +50,19 @@ where D: fusor_cpu::Scalar, { let coeff = D::from_f32((2.0 / std::f32::consts::PI).sqrt()); + // Clamp the tanh approximation input range before the cubic term. + // Large transformer activations can otherwise drive the backend tanh into + // unstable territory on GPU, even though GELU is already saturated there. + let clamped = self.clamp(D::from_f32(-5.5), D::from_f32(5.5)); // x^2 - let x_squared = self * self; + let x_squared = &clamped * &clamped; // 0.044715 * x^2 + 1.0 let inner_factor = x_squared * D::from_f32(0.044715) + D::from_f32(1.0); // x * (1 + 0.044715 * x^2) - let inner = self * &inner_factor; + let inner = &clamped * &inner_factor; // sqrt(2/pi) * (x * (1 + 0.044715 * x^2)) let tanh_input = inner * coeff; @@ -146,9 +150,13 @@ mod tests { async fn test_gelu_cpu_vs_gpu() { use crate::Device; - // Create random-ish data similar to FFN activations + // Use a wider activation range so backend tanh/gelu instability shows up. let data: Vec = (0..1 * 100 * 1536) - .map(|i| (i as f32 * 0.001).sin() * 5.0) + .map(|i| { + let base = (i as f32 * 0.001).sin() * 40.0; + let offset = ((i % 29) as f32 - 14.0) * 2.0; + base + offset + }) .collect(); // CPU version @@ -170,8 +178,8 @@ mod tests { let mut sum_diff = 0.0f32; let mut count = 0; for i in 0..cpu_slice.shape()[0] { - for j in 0..cpu_slice.shape()[1].min(50) { - for k in 0..cpu_slice.shape()[2].min(100) { + for j in 0..cpu_slice.shape()[1] { + for k in 0..cpu_slice.shape()[2] { let cpu_val: f32 = cpu_slice[[i, j, k]].into(); let gpu_val: f32 = gpu_slice[[i, j, k]].into(); let diff = (cpu_val - gpu_val).abs(); diff --git a/fusor-ml/fusor/src/layers/linear.rs b/fusor-ml/fusor/src/layers/linear.rs index 6bec403b5..e4f8ee518 100644 --- a/fusor-ml/fusor/src/layers/linear.rs +++ b/fusor-ml/fusor/src/layers/linear.rs @@ -73,12 +73,27 @@ impl Linear { where B: fusor_cpu::TensorBacking<3, Elem = f32>, { - let output = input.q_mat_mul(&self.weight); + let [batch_size, seq_len, _in_features] = input.shape(); + let out_features = self.weight.shape()[0]; + let flattened_input: Tensor<2, f32> = input + .reshape([batch_size * seq_len, self.weight.shape()[1]]) + .to_concrete(); + let output = flattened_input.q_mat_mul(&self.weight); if let Some(bias) = &self.bias { - output.add_(bias) + let bias_broadcast: Tensor<2, f32> = bias + .unsqueeze(0) + .to_concrete() + .broadcast_as(output.shape()) + .to_concrete(); + output + .add_(&bias_broadcast) + .reshape([batch_size, seq_len, out_features]) + .to_concrete() } else { output + .reshape([batch_size, seq_len, out_features]) + .to_concrete() } } } @@ -96,21 +111,34 @@ where where B: fusor_cpu::TensorBacking<3, Elem = T>, { + let [batch_size, seq_len, _in_features] = input.shape(); + let out_features = self.weight.shape()[0]; // Cast input to f32 let input_f32 = input.cast::(); + let flattened_input: Tensor<2, f32> = input_f32 + .reshape([batch_size * seq_len, self.weight.shape()[1]]) + .to_concrete(); // Do quantized matmul in f32 - let output_f32 = input_f32.q_mat_mul(&self.weight); + let output_f32 = flattened_input.q_mat_mul(&self.weight); // Add bias if present (in f32) let output_f32 = if let Some(bias) = &self.bias { let bias_f32: Tensor<1, f32> = bias.cast(); - output_f32.add_(&bias_f32) + let bias_broadcast: Tensor<2, f32> = bias_f32 + .unsqueeze(0) + .to_concrete() + .broadcast_as(output_f32.shape()) + .to_concrete(); + output_f32.add_(&bias_broadcast) } else { output_f32 }; // Cast back to T - output_f32.cast() + output_f32 + .reshape([batch_size, seq_len, out_features]) + .to_concrete() + .cast() } } diff --git a/fusor-ml/fusor/src/lib.rs b/fusor-ml/fusor/src/lib.rs index 5e326fdad..e58ae2cc8 100644 --- a/fusor-ml/fusor/src/lib.rs +++ b/fusor-ml/fusor/src/lib.rs @@ -1093,7 +1093,9 @@ where let rhs_tensor = fusor_cpu::Tensor::new(rhs_concrete); let rhs_transposed = rhs_tensor.transpose(0, 1); - // Reshape to R dimensions: [1, 1, ..., K, N] + // Reshape to R dimensions: [1, 1, ..., K, N], then broadcast + // across the lhs batch dimensions so CPU batched matmul sees + // matching batch shapes. let weight_shape: [usize; R] = std::array::from_fn(|i| { if i < R - 2 { 1 // Broadcast batch dimensions @@ -1103,10 +1105,22 @@ where n // N dimension } }); - let rhs_broadcast = rhs_transposed.reshape(weight_shape); // Do regular matmul let lhs_eval = lhs.to_concrete(); + let lhs_shape = lhs_eval.shape(); + let broadcast_shape: [usize; R] = std::array::from_fn(|i| { + if i < R - 2 { + lhs_shape[i] + } else if i == R - 2 { + k + } else { + n + } + }); + let rhs_broadcast = rhs_transposed + .reshape(weight_shape) + .broadcast_as(broadcast_shape); let result = lhs_eval.matmul(rhs_broadcast); Tensor::Cpu(result) } @@ -1133,7 +1147,9 @@ where let rhs_tensor = fusor_cpu::Tensor::new(rhs_concrete); let rhs_transposed = rhs_tensor.transpose(0, 1); - // Reshape to R dimensions: [1, 1, ..., K, N] + // Reshape to R dimensions: [1, 1, ..., K, N], then broadcast + // across the lhs batch dimensions so CPU batched matmul sees + // matching batch shapes. let weight_shape: [usize; R] = std::array::from_fn(|i| { if i < R - 2 { 1 // Broadcast batch dimensions @@ -1143,10 +1159,22 @@ where n // N dimension } }); - let rhs_broadcast = rhs_transposed.reshape(weight_shape); // Do regular matmul let lhs_eval = lhs.to_concrete(); + let lhs_shape = lhs_eval.shape(); + let broadcast_shape: [usize; R] = std::array::from_fn(|i| { + if i < R - 2 { + lhs_shape[i] + } else if i == R - 2 { + k + } else { + n + } + }); + let rhs_broadcast = rhs_transposed + .reshape(weight_shape) + .broadcast_as(broadcast_shape); let result = lhs_eval.matmul(rhs_broadcast); Tensor::Cpu(result) } diff --git a/fusor-ml/fusor/src/quantized.rs b/fusor-ml/fusor/src/quantized.rs index bfc8396e9..2814f5655 100644 --- a/fusor-ml/fusor/src/quantized.rs +++ b/fusor-ml/fusor/src/quantized.rs @@ -446,4 +446,87 @@ mod tests { result_1 ); } + + #[test] + fn test_cpu_f32_qmatmul_broadcasts_batch_dims() { + let shape = [2, 4]; // [N, K] = [out_features, in_features] + let weight_data: Vec = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]; + let weight_bytes: Vec = weight_data.iter().flat_map(|f| f.to_le_bytes()).collect(); + + let qmatrix: QMatrix = + QMatrix::from_raw_bytes(&Device::Cpu, shape, &weight_bytes, GgmlType::F32).unwrap(); + + let input_data: Vec = vec![ + 1.0, 1.0, 1.0, 1.0, 0.0, 1.0, 0.0, 1.0, 2.0, 0.0, 0.0, 0.0, 1.0, 0.0, 1.0, 0.0, 0.0, + 0.0, 1.0, 1.0, 3.0, 1.0, 0.0, 0.0, + ]; + let input: Tensor<3, f32> = Tensor::from_slice(&Device::Cpu, [2, 3, 4], &input_data); + + let output = input.q_mat_mul(&qmatrix).unwrap_cpu(); + assert_eq!(output.shape(), [2, 3, 2]); + + let expected = [ + [[10.0, 26.0], [6.0, 14.0], [2.0, 10.0]], + [[4.0, 12.0], [7.0, 15.0], [5.0, 21.0]], + ]; + + for batch in 0..2 { + for row in 0..3 { + for col in 0..2 { + let actual = output.get([batch, row, col]); + let expected = expected[batch][row][col]; + assert!( + (actual - expected).abs() < 0.1, + "output[{batch}, {row}, {col}] = {actual}, expected {expected}" + ); + } + } + } + } + + #[test] + fn test_cpu_f16_qmatmul_broadcasts_batch_dims() { + let shape = [2, 4]; // [N, K] = [out_features, in_features] + let weight_data: Vec = vec![ + f16::from_f32(1.0), + f16::from_f32(2.0), + f16::from_f32(3.0), + f16::from_f32(4.0), + f16::from_f32(5.0), + f16::from_f32(6.0), + f16::from_f32(7.0), + f16::from_f32(8.0), + ]; + let weight_bytes: Vec = weight_data.iter().flat_map(|f| f.to_le_bytes()).collect(); + + let qmatrix: QMatrix = + QMatrix::from_raw_bytes(&Device::Cpu, shape, &weight_bytes, GgmlType::F16).unwrap(); + + let input_data: Vec = vec![ + 1.0, 1.0, 1.0, 1.0, 0.0, 1.0, 0.0, 1.0, 2.0, 0.0, 0.0, 0.0, 1.0, 0.0, 1.0, 0.0, 0.0, + 0.0, 1.0, 1.0, 3.0, 1.0, 0.0, 0.0, + ]; + let input: Tensor<3, f32> = Tensor::from_slice(&Device::Cpu, [2, 3, 4], &input_data); + + let output = input.q_mat_mul(&qmatrix).unwrap_cpu(); + assert_eq!(output.shape(), [2, 3, 2]); + + let expected = [ + [[10.0, 26.0], [6.0, 14.0], [2.0, 10.0]], + [[4.0, 12.0], [7.0, 15.0], [5.0, 21.0]], + ]; + + for batch in 0..2 { + for row in 0..3 { + for col in 0..2 { + let actual = output.get([batch, row, col]); + let expected = expected[batch][row][col]; + assert!( + (actual - expected).abs() < 0.1, + "output[{batch}, {row}, {col}] = {actual}, expected {expected}" + ); + } + } + } + } } diff --git a/models/rbert/src/language_model.rs b/models/rbert/src/language_model.rs index 7201d0051..9b04b7bea 100644 --- a/models/rbert/src/language_model.rs +++ b/models/rbert/src/language_model.rs @@ -44,6 +44,24 @@ impl Bert { )) } + /// Embed a batch of sentences with a specific pooling strategy and explicit + /// control over whether the pooled vectors are L2-normalized. + pub async fn embed_batch_with_pooling_and_normalization( + &self, + inputs: Vec<&str>, + pooling: Pooling, + normalize: bool, + ) -> Result, BertError> { + let tensors = self.embed_batch_raw_with_options(inputs, pooling, normalize)?; + + let mut embeddings = Vec::with_capacity(tensors.len()); + for tensor in tensors { + embeddings.push(self.tensor_to_embedding(tensor).await?); + } + + Ok(embeddings) + } + /// Embed a sentence with a specific pooling strategy. pub async fn embed_with_pooling( &self, @@ -60,14 +78,8 @@ impl Bert { inputs: Vec<&str>, pooling: Pooling, ) -> Result, BertError> { - let tensors = self.embed_batch_raw(inputs, pooling)?; - - let mut embeddings = Vec::with_capacity(tensors.len()); - for tensor in tensors { - embeddings.push(self.tensor_to_embedding(tensor).await?); - } - - Ok(embeddings) + self.embed_batch_with_pooling_and_normalization(inputs, pooling, true) + .await } } diff --git a/models/rbert/src/lib.rs b/models/rbert/src/lib.rs index 509fea494..b68e65630 100644 --- a/models/rbert/src/lib.rs +++ b/models/rbert/src/lib.rs @@ -375,6 +375,15 @@ impl Bert { &self, sentences: Vec<&str>, pooling: Pooling, + ) -> Result>, BertError> { + self.embed_batch_raw_with_options(sentences, pooling, true) + } + + pub(crate) fn embed_batch_raw_with_options( + &self, + sentences: Vec<&str>, + pooling: Pooling, + normalize: bool, ) -> Result>, BertError> { let embedding_dim = self.model.embedding_dim(); // Approximates the quadratic attention memory cost (seq_len^2). @@ -424,7 +433,7 @@ impl Bert { } for (indices, encodings) in chunks { - let embeddings = self.embed_batch_raw_inner(encodings, pooling)?; + let embeddings = self.embed_batch_raw_inner(encodings, pooling, normalize)?; for (i, embedding) in indices.iter().zip(embeddings) { combined[*i] = Some(embedding); } @@ -436,6 +445,7 @@ impl Bert { &self, mut tokens: Vec, pooling: Pooling, + normalize: bool, ) -> Result>, BertError> { if tokens.is_empty() { return Ok(Vec::new()); @@ -485,13 +495,19 @@ impl Bert { // Cast mask u32→f32, unsqueeze to [batch, seq_len, 1] for broadcasting let mask_f32: Tensor<2, f32> = attention_mask.cast(); let mask_3d: Tensor<3, f32, _> = mask_f32.unsqueeze(2).to_concrete(); + // Broadcast mask to match embedding shape [batch, seq, hidden] + let mask_3d: Tensor<3, f32> = mask_3d.broadcast_as(shape).to_concrete(); // Zero out padding positions, sum along seq dim → [batch, hidden_dim] let masked_embeddings = (embeddings * mask_3d).to_concrete(); let summed = masked_embeddings.sum::<2>(1); // Divide by valid token count [batch, 1] — broadcasts with [batch, hidden_dim] let valid_count = mask_f32.sum_keepdim::<1>(1); let embeddings = summed.div_(&valid_count); - let embeddings = normalize_l2(&embeddings); + let embeddings = if normalize { + normalize_l2(&embeddings) + } else { + embeddings + }; Ok(embeddings .chunk(n_sentences, 0) .into_iter() @@ -510,7 +526,11 @@ impl Bert { Pooling::Last => { // With left padding, the last token is always at the final position let indexed_embeddings = embeddings.to_concrete().i((.., n_tokens - 1, ..)); - let normalized = normalize_l2(&indexed_embeddings); + let normalized = if normalize { + normalize_l2(&indexed_embeddings) + } else { + indexed_embeddings + }; Ok(normalized .chunk(n_sentences, 0) .into_iter() @@ -519,6 +539,232 @@ impl Bert { } } } + + #[doc(hidden)] + pub fn debug_batch_forward( + &self, + sentences: Vec, + ) -> Result<(Tensor<3, f32>, Tensor<2, u32>), BertError> { + let mut tokens = { + let tokenizer_read = self.tokenizer.read().unwrap(); + tokenizer_read.encode_batch(sentences, true) + } + .map_err(BertError::TokenizerError)?; + + let pp = PaddingParams { + strategy: tokenizers::PaddingStrategy::BatchLongest, + direction: PaddingDirection::Right, + ..Default::default() + }; + tokenizers::pad_encodings(&mut tokens, &pp).map_err(BertError::TokenizerError)?; + + let device = self.model.device(); + let max_seq_len = self.model.max_seq_len(); + let token_ids = tokens.iter().map(|tokens| { + let tokens = tokens.get_ids().to_vec(); + Tensor::new( + device, + &tokens.as_slice()[..max_seq_len.min(tokens.as_slice().len())], + ) + }); + let token_ids = Tensor::stack(token_ids, 0); + + let attention_masks = tokens.iter().map(|tokens| { + let attention_mask = tokens.get_attention_mask(); + Tensor::new( + device, + &attention_mask[..max_seq_len.min(attention_mask.len())], + ) + }); + let attention_mask = Tensor::stack(attention_masks, 0); + + let embeddings = self.model.forward(&token_ids, Some(&attention_mask)); + Ok((embeddings, attention_mask)) + } + + #[doc(hidden)] + pub fn debug_batch_mean_pool( + &self, + sentences: Vec, + normalize: bool, + ) -> Result, BertError> { + let (embeddings, attention_mask) = self.debug_batch_forward(sentences)?; + let shape = embeddings.shape(); + let mask_f32: Tensor<2, f32> = attention_mask.cast(); + let mask_3d: Tensor<3, f32, _> = mask_f32.unsqueeze(2).to_concrete(); + let mask_3d: Tensor<3, f32> = mask_3d.broadcast_as(shape).to_concrete(); + let masked_embeddings = (embeddings * mask_3d).to_concrete(); + let summed = masked_embeddings.sum::<2>(1); + let valid_count = mask_f32.sum_keepdim::<1>(1); + let pooled = summed.div_(&valid_count); + Ok(if normalize { + normalize_l2(&pooled) + } else { + pooled + }) + } + + #[doc(hidden)] + pub fn debug_batch_hidden_states( + &self, + sentences: Vec, + ) -> Result<(Vec>, Tensor<2, u32>), BertError> { + let mut tokens = { + let tokenizer_read = self.tokenizer.read().unwrap(); + tokenizer_read.encode_batch(sentences, true) + } + .map_err(BertError::TokenizerError)?; + + let pp = PaddingParams { + strategy: tokenizers::PaddingStrategy::BatchLongest, + direction: PaddingDirection::Right, + ..Default::default() + }; + tokenizers::pad_encodings(&mut tokens, &pp).map_err(BertError::TokenizerError)?; + + let device = self.model.device(); + let max_seq_len = self.model.max_seq_len(); + let token_ids = tokens.iter().map(|tokens| { + let tokens = tokens.get_ids().to_vec(); + Tensor::new( + device, + &tokens.as_slice()[..max_seq_len.min(tokens.as_slice().len())], + ) + }); + let token_ids = Tensor::stack(token_ids, 0); + + let attention_masks = tokens.iter().map(|tokens| { + let attention_mask = tokens.get_attention_mask(); + Tensor::new( + device, + &attention_mask[..max_seq_len.min(attention_mask.len())], + ) + }); + let attention_mask = Tensor::stack(attention_masks, 0); + + let states = match &*self.model { + EmbeddingModel::Bert(model) => { + let token_type_ids = token_ids.zeros_like(); + model.debug_hidden_states(&token_ids, &token_type_ids, Some(&attention_mask)) + } + EmbeddingModel::Qwen(_) => { + return Err(BertError::Fusor(fusor::Error::msg( + "debug_batch_hidden_states is only implemented for BERT models", + ))); + } + }; + + Ok((states, attention_mask)) + } + + #[doc(hidden)] + pub fn debug_batch_first_layer( + &self, + sentences: Vec, + ) -> Result<(Tensor<3, f32>, Tensor<3, f32>, Tensor<3, f32>), BertError> { + let mut tokens = { + let tokenizer_read = self.tokenizer.read().unwrap(); + tokenizer_read.encode_batch(sentences, true) + } + .map_err(BertError::TokenizerError)?; + + let pp = PaddingParams { + strategy: tokenizers::PaddingStrategy::BatchLongest, + direction: PaddingDirection::Right, + ..Default::default() + }; + tokenizers::pad_encodings(&mut tokens, &pp).map_err(BertError::TokenizerError)?; + + let device = self.model.device(); + let max_seq_len = self.model.max_seq_len(); + let token_ids = tokens.iter().map(|tokens| { + let tokens = tokens.get_ids().to_vec(); + Tensor::new( + device, + &tokens.as_slice()[..max_seq_len.min(tokens.as_slice().len())], + ) + }); + let token_ids = Tensor::stack(token_ids, 0); + + let attention_masks = tokens.iter().map(|tokens| { + let attention_mask = tokens.get_attention_mask(); + Tensor::new( + device, + &attention_mask[..max_seq_len.min(attention_mask.len())], + ) + }); + let attention_mask = Tensor::stack(attention_masks, 0); + + match &*self.model { + EmbeddingModel::Bert(model) => { + let token_type_ids = token_ids.zeros_like(); + model.debug_first_layer(&token_ids, &token_type_ids, Some(&attention_mask)) + .ok_or_else(|| BertError::Fusor(fusor::Error::msg("BERT encoder has no layers"))) + } + EmbeddingModel::Qwen(_) => Err(BertError::Fusor(fusor::Error::msg( + "debug_batch_first_layer is only implemented for BERT models", + ))), + } + } + + #[doc(hidden)] + pub fn debug_batch_first_layer_attention( + &self, + sentences: Vec, + ) -> Result< + ( + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<3, f32>, + Tensor<3, f32>, + ), + BertError, + > { + let mut tokens = { + let tokenizer_read = self.tokenizer.read().unwrap(); + tokenizer_read.encode_batch(sentences, true) + } + .map_err(BertError::TokenizerError)?; + + let pp = PaddingParams { + strategy: tokenizers::PaddingStrategy::BatchLongest, + direction: PaddingDirection::Right, + ..Default::default() + }; + tokenizers::pad_encodings(&mut tokens, &pp).map_err(BertError::TokenizerError)?; + + let device = self.model.device(); + let max_seq_len = self.model.max_seq_len(); + let token_ids = tokens.iter().map(|tokens| { + let tokens = tokens.get_ids().to_vec(); + Tensor::new( + device, + &tokens.as_slice()[..max_seq_len.min(tokens.as_slice().len())], + ) + }); + let token_ids = Tensor::stack(token_ids, 0); + + let attention_masks = tokens.iter().map(|tokens| { + let attention_mask = tokens.get_attention_mask(); + Tensor::new( + device, + &attention_mask[..max_seq_len.min(attention_mask.len())], + ) + }); + let attention_mask = Tensor::stack(attention_masks, 0); + + match &*self.model { + EmbeddingModel::Bert(model) => { + let token_type_ids = token_ids.zeros_like(); + model.debug_first_layer_attention(&token_ids, &token_type_ids, Some(&attention_mask)) + .ok_or_else(|| BertError::Fusor(fusor::Error::msg("BERT encoder has no layers"))) + } + EmbeddingModel::Qwen(_) => Err(BertError::Fusor(fusor::Error::msg( + "debug_batch_first_layer_attention is only implemented for BERT models", + ))), + } + } } fn normalize_l2(v: &Tensor<2, f32>) -> Tensor<2, f32> { diff --git a/models/rbert/src/raw/attention.rs b/models/rbert/src/raw/attention.rs index 99f4a5626..2aeb7ec31 100644 --- a/models/rbert/src/raw/attention.rs +++ b/models/rbert/src/raw/attention.rs @@ -34,4 +34,28 @@ impl BertAttention { let self_outputs = self.self_attention.forward(hidden_states, attention_mask); self.self_output.forward(&self_outputs, hidden_states) } + + pub(crate) fn debug_forward( + &self, + hidden_states: &Tensor<3, f32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> ( + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<3, f32>, + Tensor<3, f32>, + ) { + let _enter = self.span.enter(); + let (query_layer, key_layer, value_layer, self_outputs) = + self.self_attention.debug_forward(hidden_states, attention_mask); + let attention_output = self.self_output.forward(&self_outputs, hidden_states); + ( + query_layer, + key_layer, + value_layer, + self_outputs, + attention_output, + ) + } } diff --git a/models/rbert/src/raw/encoder.rs b/models/rbert/src/raw/encoder.rs index 20aff7974..64165efe2 100644 --- a/models/rbert/src/raw/encoder.rs +++ b/models/rbert/src/raw/encoder.rs @@ -35,4 +35,45 @@ impl BertEncoder { } hidden_states } + + pub(crate) fn debug_hidden_states( + &self, + hidden_states: &Tensor<3, f32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Vec> { + let _enter = self.span.enter(); + let mut hidden_states = hidden_states.clone(); + let mut states = Vec::with_capacity(self.layers.len()); + for layer in self.layers.iter() { + hidden_states = layer.forward(&hidden_states, attention_mask); + states.push(hidden_states.clone()); + } + states + } + + pub(crate) fn debug_first_layer( + &self, + hidden_states: &Tensor<3, f32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Option<(Tensor<3, f32>, Tensor<3, f32>, Tensor<3, f32>)> { + self.layers + .first() + .map(|layer| layer.debug_forward(hidden_states, attention_mask)) + } + + pub(crate) fn debug_first_layer_attention( + &self, + hidden_states: &Tensor<3, f32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Option<( + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<3, f32>, + Tensor<3, f32>, + )> { + self.layers + .first() + .map(|layer| layer.debug_attention_forward(hidden_states, attention_mask)) + } } diff --git a/models/rbert/src/raw/layer.rs b/models/rbert/src/raw/layer.rs index 207ba2f5c..1f48adbba 100644 --- a/models/rbert/src/raw/layer.rs +++ b/models/rbert/src/raw/layer.rs @@ -38,4 +38,31 @@ impl BertLayer { let intermediate_output = self.intermediate.forward(&attention_output); self.output.forward(&intermediate_output, &attention_output) } + + pub(crate) fn debug_forward( + &self, + hidden_states: &Tensor<3, f32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> (Tensor<3, f32>, Tensor<3, f32>, Tensor<3, f32>) { + let _enter = self.span.enter(); + let attention_output = self.attention.forward(hidden_states, attention_mask); + let intermediate_output = self.intermediate.forward(&attention_output); + let layer_output = self.output.forward(&intermediate_output, &attention_output); + (attention_output, intermediate_output, layer_output) + } + + pub(crate) fn debug_attention_forward( + &self, + hidden_states: &Tensor<3, f32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> ( + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<3, f32>, + Tensor<3, f32>, + ) { + let _enter = self.span.enter(); + self.attention.debug_forward(hidden_states, attention_mask) + } } diff --git a/models/rbert/src/raw/mod.rs b/models/rbert/src/raw/mod.rs index 9b8113f5b..adf89b896 100644 --- a/models/rbert/src/raw/mod.rs +++ b/models/rbert/src/raw/mod.rs @@ -145,6 +145,55 @@ impl BertModel { self.encoder.forward(&embedding_output, attention_mask) } + #[doc(hidden)] + pub fn debug_hidden_states( + &self, + input_ids: &Tensor<2, u32>, + token_type_ids: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Vec> { + let _enter = self.span.enter(); + let embedding_output = self.embeddings.forward(input_ids, token_type_ids); + let mut states = vec![embedding_output.clone()]; + states.extend( + self.encoder + .debug_hidden_states(&embedding_output, attention_mask), + ); + states + } + + #[doc(hidden)] + pub fn debug_first_layer( + &self, + input_ids: &Tensor<2, u32>, + token_type_ids: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Option<(Tensor<3, f32>, Tensor<3, f32>, Tensor<3, f32>)> { + let _enter = self.span.enter(); + let embedding_output = self.embeddings.forward(input_ids, token_type_ids); + self.encoder + .debug_first_layer(&embedding_output, attention_mask) + } + + #[doc(hidden)] + pub fn debug_first_layer_attention( + &self, + input_ids: &Tensor<2, u32>, + token_type_ids: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Option<( + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<3, f32>, + Tensor<3, f32>, + )> { + let _enter = self.span.enter(); + let embedding_output = self.embeddings.forward(input_ids, token_type_ids); + self.encoder + .debug_first_layer_attention(&embedding_output, attention_mask) + } + pub(crate) fn max_seq_len(&self) -> usize { self.embeddings.max_seq_len() } diff --git a/models/rbert/src/raw/self_attention.rs b/models/rbert/src/raw/self_attention.rs index 2ad23ebb6..b914da45a 100644 --- a/models/rbert/src/raw/self_attention.rs +++ b/models/rbert/src/raw/self_attention.rs @@ -59,36 +59,65 @@ impl BertSelfAttention { let key_layer = self.transpose_for_scores(&key_layer); let value_layer = self.transpose_for_scores(&value_layer); - let attention_scores = query_layer.mat_mul(&key_layer.t()); - let mut attention_scores = - attention_scores.div_scalar((self.attention_head_size as f32).sqrt()); + let scale = 1.0 / (self.attention_head_size as f32).sqrt(); + const MASK_NEG_VALUE: f32 = -10000.0; + let mask: Option> = attention_mask.map(|m| { + let mask_f32: Tensor<2, f32> = m.cast(); + let zeros = mask_f32.zeros_like(); + let ones = (zeros + 1.0f32).to_concrete(); + ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete() + }); - // If there is an attention mask, filter the attention scores by that mask - if let Some(attention_mask) = attention_mask { - // The attention mask is a tensor of shape (bsize, seq_len) - // the attention scores are a tensor of shape (bsize, _, seq_len, seq_len) - // We expand the attention mask to (bsize, 1, 1, seq_len) - let mask = attention_mask - .unsqueeze::<3>(1) - .unsqueeze::<4>(2) - .to_concrete(); - let shape = attention_scores.shape(); - let mask: Tensor<4, f32> = mask.broadcast_as::<4>(shape).to_concrete().cast(); - // We use a value slightly larger that the true f32 min value to avoid NaN - const FALSE_MIN: f32 = -3.4028235e34f32; - let device = attention_scores.device(); - let on_false = Tensor::splat(&device, FALSE_MIN, shape); - attention_scores = mask.where_cond(&attention_scores, &on_false); - } - - let attention_probs = { + let context_layer = { let _enter_sm = self.span_softmax.enter(); - attention_scores.softmax_last_dim::<3>() + query_layer.flash_attention( + &key_layer, + &value_layer, + scale, + mask.as_ref().map(|m| (m, fusor::MaskKind::BatchKeyMask)), + ) }; - let context_layer = attention_probs.mat_mul(&value_layer); let context_layer = context_layer.transpose(1, 2).to_concrete(); context_layer.flatten_last_n::<1, _>() } + + pub(crate) fn debug_forward( + &self, + hidden_states: &Tensor<3, f32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> (Tensor<4, f32>, Tensor<4, f32>, Tensor<4, f32>, Tensor<3, f32>) { + let _enter = self.span.enter(); + let query_layer = self.query.forward(hidden_states); + let key_layer = self.key.forward(hidden_states); + let value_layer = self.value.forward(hidden_states); + + let query_layer = self.transpose_for_scores(&query_layer); + let key_layer = self.transpose_for_scores(&key_layer); + let value_layer = self.transpose_for_scores(&value_layer); + + let scale = 1.0 / (self.attention_head_size as f32).sqrt(); + const MASK_NEG_VALUE: f32 = -10000.0; + let mask: Option> = attention_mask.map(|m| { + let mask_f32: Tensor<2, f32> = m.cast(); + let zeros = mask_f32.zeros_like(); + let ones = (zeros + 1.0f32).to_concrete(); + ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete() + }); + + let context_layer = { + let _enter_sm = self.span_softmax.enter(); + query_layer.flash_attention( + &key_layer, + &value_layer, + scale, + mask.as_ref().map(|m| (m, fusor::MaskKind::BatchKeyMask)), + ) + }; + let context_layer = context_layer.transpose(1, 2).to_concrete(); + let context_layer = context_layer.flatten_last_n::<1, _>(); + + (query_layer, key_layer, value_layer, context_layer) + } } // attention_probs before matmul: Tensor[dims 3, 12, 13, 13; f32] diff --git a/models/rgliner/Cargo.toml b/models/rgliner/Cargo.toml new file mode 100644 index 000000000..aaaad55e6 --- /dev/null +++ b/models/rgliner/Cargo.toml @@ -0,0 +1,33 @@ +[package] +name = "rgliner" +version = "0.4.0" +edition = "2021" +description = "GLiNER bi-encoder Named Entity Recognition for Rust" +license = "MIT/Apache-2.0" +repository = "https://github.com/floneum/floneum" +authors = ["Evan Almloff "] +keywords = ["ai", "ner", "nlp", "gliner", "transformers"] + +[dependencies] +fusor = { workspace = true } +fusor-core = { workspace = true } +fusor-gguf = { workspace = true } +tokenizers = { workspace = true, features = ["fancy-regex"] } +thiserror.workspace = true +rbert = { workspace = true } + +tracing = "0.1.37" +serde_json = "1.0.106" +serde = { version = "1", features = ["derive"] } + +kalosm-common = { workspace = true } +kalosm-model-types.workspace = true +kalosm-language-model.workspace = true + +[dev-dependencies] +anyhow.workspace = true +tokio = { version = "1", features = ["full"] } + +[features] +default = ["tokio"] +tokio = ["kalosm-common/tokio"] diff --git a/models/rgliner/__pycache__/convert_to_gguf.cpython-311.pyc b/models/rgliner/__pycache__/convert_to_gguf.cpython-311.pyc new file mode 100644 index 0000000000000000000000000000000000000000..9af6e8f7d95b1773518b6ecd0101655f766c5e0e GIT binary patch literal 27125 zcmeHwS!^81m0(s?)?G!`Nfr-LMN$XP;!ROJ#HK`vIw(plQQa+eHwq0`c}m3$c)@#00qiz6f6dlCFwr_)Bj75?^rkN$3iApQxxWG@^7 zFJCO{2;vbz5fcQZqx3P|gbqLT6MFnKOc?N!oFMVjIAMgRA!eF2PndNm4jHq|S|_Zt zwh7y;eZoHLm~dc!eatyggyR%X6vNXPbIrOZ+_Rnu4}_a0yp(yuN9iW~lx3oXvQCu3 zdw{Y{lu`DHa>_ALK{+QXsiKJkRPjU=<(fE1xhH~@XQCS7)IgkCh;sj1Z&@=i49i4lVGy-QI3kM#uc3H-^|L?cy0mC#L8Df|U!2W_L9=;E8r+3+cSS=LF= zHs#xti7I>7I1!@Cam=PElB#%@oM^EUR)VUk%-y6X*(ek(gBPqz&~X-s(?OPr#1n`aKn3y@ zn~P916-?0c5hlXYv3vEl&k(Q2>U9!%Il|7=lagU5eorEYqm!&ejz$x#WE@|ZkJ0sd zsd#28JhN~kOwHYmBTAAx5s%C#X6D$OR}^voGXs$P3@VHG3{~uBBmzbRObD0}upnSX z0QI)d>+nvGBqHX@m!MiX;P0xE(`(6K2Tom{2gf5K zsZSC&vtdM1hZ3N4zy#fv>nW9{&TwQ08(G%K)K53XTd<>zVRc=1D?0>Br)cTqv$xC(T|!l(s3+rc{~R2sMpN)CyeOTzoQsB}ZhOA4 zh?AM?CZdclO3$L#+dZT2LM|H`F2M5Unu@}%-Iv<(5 zB^iK5rs^c!6w?I7k<4-gN>h2xFeVU|Q|AJT`cs9R0~(1-0H6vz-bZy0>Q;{lo?6jU zyKGKdtlu!dXa1(?p?!7uacye)drcoT39c5=)goA0MN2DhY0coQm<3Bfv;_F(XI*(eql-rPr*0h)0YFeCUr@Z_gwSb2I%(Xk$_Q^5x?@_LjY)-T-f;qj<)OBoNCtBLpkR(#h zJZ|2FzEWzAz$jVRJJ82CbFt`)oVlhnaaZ)r-|xE&!7Mc`KeT0FYx9_pYt?rZdMF+# zevh_sC{I<*33UqiJF}nB$h7}ah$8cAFm~q=>HU54_)j$onRcs_mvB=*OuXHI=L!4a zsu8~GfS+I2YHR6!V`2MM@pb2MW&eErQ;jSC{CqY4!F=6$T>I+2+w0!Pwf)UK{!qRK zpuO%ql4@F}j<{cj+aj94XD-UZ)@R*lUZrfA>X!T+;zP|C_l}NLQnIQky*zfkojaVJ z`DRy5L3!167Z#DhC<_u6Soefru6vHt_n&0Yx}8D26W(q!A=xp-1+T>u3-j}H3=7lb z2*X6~1?d<)3sMOjzczHkv3kE`izcG+1RDV%hD6?&n~O1MA}MI5!B z*QFiKZ;ZS*^35|3FQmxFalzLm`nm*1x9I5R9o?$XwJ!e1u~Z}#S&s<59?{n$IF5^s z<9zmh?krt3ewYxPwW70ji#UX+7yDLjzjtr>UfSh;WPV`&(DZ>lHT-uiLZDX+^a`#% z(bdPh`c$E7hlN0^7-$t-ZKA7<&)zST6n@pvD+Kz)K%e017hV0ltA8&f8unKGgR_q(1$Tq!ZrCDR z&ewD*%DY$N-_@?sKWKQ;AOyO^K$qa^7G2%Et2WTMWF( zTp-g!(k^_d>8g374*P@C%*nmfbYGGtaWfY$rI4p&^ABq&aOO2FFxPH-3u<^#k7qFG z!(MQ0w-$uij$$2>j{>-bvz>N|TqLK7MWd3F!^Dd2?f>}}w82HwqIuDhw6I7xoVhZq zwrUnhKQq_OT%9D%tYYPPGbf#RVpMZy(H#E{(5=Z*Ac%d@&8-Nwp*sh3krXy_W$03W zXnvFClw$3Wr4BMPWqJKse|5`PG$+lJiG%v5%pCOKlw#(cjiI4bn)NTBM5K0V3nn#> z_ubC*s=j5unVZwA^Po;e9i}5#H2KP0$sK9nm85l_Eu8kpv@`Umro_B8co#vg?fJz zjcRygGcOipMwucQWp-E>LcdeUukxv=M@wyffJ+L_ZNsk}QiJl@wkzNLw?K1lLqBzS zbVS@&#osPIp=-|f==JQI~VSLSfxeM2cvIMW*!TIFX9t&D&v>pmh zfl@Oh>F?fPdVm^U&*9bG4|I}|LMxch&$pDrocu>w^+G z2Z}1JrPQw7uu?UAacJb+DanBsVy~UPa`oJ$F}#w(dM>QplZ>EuolW2+7;_%dlbjErW~JCXoBCL5hzm|IBTwWegoYdi>7qPKNt{J{Z>$1ItXXW*T@6Q%D; zrUh6-LM&4xfyyesL#Qz^6MG${EcT{wY0lSmt=xUgOn*9V6eteG|C zU8Kn706p=K@EwLVq|?1}=wbcpxZtc2ov?~LW^GJGUZ|FlH{}_;;Aj^e?YyI1wTyZk z6MP+_uS0NjijGd+(WwrdTs!onNjP#$+xOv0h!uwlx#R+ZrTX}!$r`yv> zR~-EA@oyi0eC4C-LPZOt1cVP%r>=Z!Vy#=KY}>3jzEN@fr!_yR7b=Fu3P^U#M!5aC z0#kykQFJx(u14+RhJnXDysJwU$~W{2fdMfvAh=G5t`ofL#0yvDrmJ?tRht^;8wdD9 zCj{4^=o;i*gK4NY>Z4mAAmCj!FDef3)g9|Mgz6rl;`q}mLiyl|8DUJrk`}@RR}EBO zMdf#oeEZ0^>OQDjzK{--FP}?S9bM~tS|&7}6sk@=yC)pDygZg^h$|zjAP31W^I>7t zd2BuTZ2aeM2_09&jw@RP>1@OR6sGT6S3t%A%vl{yRj00J3bjgp*YRz~w~9U}5?n_` z*HPYe6f&CfeN?f=2-R&uMLXoUya(kJ3r1GWsiQ)mKGzhXHEvl^D$_dXU$~QD@Hy&k zzeXVMx1H-ZzxmB}-d{30!dk=^HKz`JR1Xih1xrY@g!t^8Koi=3HP&*;WB8?qyi{WT zWtk4(7#O^O0SK~n2!|zGI1Kh83o+z&gu}NNBC$+}ISd9rli@Hkga)c00w8%J7-UGm zyoSK*03>sq!s9X1kGvWYK*BUUlH(!h2LxV@%3~6m{V@m}hrh&s1F)3)rO6tetl1(B z#(~wVTLe6`z@|6WrVP;ejkRm1@IwnM#97@n9r8tEgt2Qifgfu%PeAMe@6WvXFonxP z8p~23=z@n9IO>h(b!(Gb1bS`*U(@v&TegToAeS1I73_PE;aWIKN!H28 zXm){#s|<=@0wyohr(imxUZzuET(sRso<+<0%#{Qqh-omS`@Zgp;bZXsBp0@*XHn&4 zuDroTO5vLgl?Br3DaoRNRV)ya21>K^0zm=EDpZR%b4q`hyu$n;Z~C%mq)h0iroYRS zN0Gxl+_@!wzSK}|lRB?VgHjm{(?8R;HQ>$DP?Xz5St;8SyQUt@@l%w8 za>l?^bkPD6y8V@|sGxMY2_TxbZ=XLsk%SO+$HZ;keFmKy#&fzZh9vhy?Jn%RP18z z1KLygODsaGfax7+VMFmvhs~51n;2e9b6iH zZnnOA_pkr*oxfa(3FcbS3`4iY*!`p1`T%Ql!xSiv@`MR`9jEqi$PQSw1Db&hX2DU_ZPOHVC3b|4H>z?JK(Nx|PF z`kR(*FF+-8hqoW#$piS)`?FI2P>*co`DcSVfRZCyVdxt>spz`|8VAs5{l&k5 zg)lNK#DZ7wD7=tKxn>+dLfS$^@SzM^^MhUJ!pOX5x@Jlvz;Q-=1;(RLP!kOR2z~am%RkJ>ynTm3ziLKvnG-_aGJAGj26i zvuE5|>d>BXVI#zzagR_(_l#Rd)$bX%VIS?Rk!spA-DWDZXWSO5bEQm+)^!c1z2YrYtn&rPJ{sYnnF_4z$3 z@WH@4A=L!zF-1Suuedv0EXo{iTtnqOg-el~5^; zF7Lom1=vF{dyI!(zkTaYh_m6BFm_ur0dpbFDSOMYL!2cv8=DWk9z_9DQ}HnD`JuxL z^SJuA391XZsfW;3meB15nz@H~kLodeJH+|47_k2%7Y@y*fC4s? z$zeb<2q-9IDxHF@6tG&UEDTYU$|L7N*~%a{AkC#=6BQMOW)jNKvEcTB+DuVMoq{L( zuAu7;VMc^iguOCT=E1%Rx_E}JQ{~B7nRq-DOR&Ik5;?gL+8L%888ok%{a~x zcSQT0rQx)_SoVCCzj{-!-xBS&0N0emJ^Ac{V81NdFDp<-`LSuiJ|o&^6sW7tR8p{a ziuTSN=Bhmv7VKT3y-SH_OI_09*;1o|y-l>Y<)b=MF~Qy;+B+0c`M}__fM7o_+Rv+e z`}l#&g8fy|{%S5hpz=en3-&ie`x^?+A^udMXBU4AawL~gQS#8bL$D8s_5oE|t?Lb0 z-yG{g>o&pOC))cIDQ$<>z7FXOi}qm!UN_$ZWxpueFRGH+j%3|hlVCq4+K(xsH}RoU zS(=KxH?MtFun&s%K_Hr~cc_im1bdTcZ&IanmOuZdV1G-rzoj6%t>>9v6c;e`_wd_3 zn>l7Gvi`zeymSG}J|)8x=-KShxnMIb7J5h)kT1jJhf{B=L*fm*_I=^jN=??`cu*6H{HNuWli z28O8JthER)^cnHEwxiEz}4TJ?D%$y5TnCU~9hXt-=FOc!+3uTb+=%U1; z75+eFzsGX>aJlV^W>8Ms?+}a$y>f0vxN!IVN#Gm4yeD6cw2_1i8qeLK=D zkO%CUX6cSJ15}xsiuw)mxl;+tqna&!V6frbfg6QVgrXgB3#ACfJK`2f5nMas7D^G^ zJK`2f5j;EM9;S|{dA8`?5%(xnr@{5@*vjjZ#cJGL+DyYfa2t~v$<{7(o01xd*)F)v zN%x-f8%k<=$1ZeR*qreZXk2{i9=_)kTdB4^RvWjd5*d`X)+M=Lr?WLYx@m+8XUzJBTq?# zlQVbBV`$S_EQ$Y@T0&{~c!*$r28Kl--HG`YOaGa3oHK;nCC34_IUjnnH0w^5Zf|Hw z9ZLpMNJm?6X}FdD@F+0VIO_JuHsuj-=FIb+Xe3=JMPJC{CjT|zbs zQ?+7!(hhC)c+!t&DkXUe+s#y9l&Z|9mv<%o$r3PMgx-R+ZLf%Cc%JIrn zvz;;x@5^`mb`BMMJ1dgqJI;(ws#^+-COhH&UNa-=6teLuNTXuU-}1`jPPhlOv!KNT zJK`QJ&>mFnh+9n!tMj{fa7Wx)>a+$oxFhai>Wl`rdPiQ5BoC-@cd3`NNsVS?7u<76 zjmBjc-1EtUd(Oj!Wc8kLN2!Z@#vMy)G&!37s{UG9)qMoUo=e;GKiU>kU-$MMvkzqQ2TAM0CNzRz#%yD&!W|6BWN zvh0uit<`-0_VTS&%}B1O`y$L6yFp8KmHV%!)W^GF{Yl}eZph;06hnhB9@J-dvuD16 z6;DDw>a98W7Yf!P^{GB6XVomK%!31E;eIv zZgW<Fv_>z=gt%hcMP#;=wN&cm^uZvf?+yM|7i1nl7XyJ9 zGDE)skSkYijm}8Lun7}vo37ej-^BT*X*lOA5(`DZR0>SGC&7*o4B2IyD$bD&MT*DY z>RLYg9b8;x)>oNU7@4{QBN!y};9N3+&}`%@;k!`^43VrkCxga=90pbeAVy2pywfrp zVQ$gPYpq;qhEQVRhFlV)`ao9gt^yz9RFr|!=ZZ3%&KQJA_L(RJijr^wM1|2Mu#=Iz z@rBuNPN5W@p(9j+L8kAL8--->vGygwAd!{iR;qJkO0R%qU9#N)W3d~u$FrYpj$C1J26G6_a=&f=$hL&%-3KiZPj3RR=&Jb#!;tauy zCpjZNc9HXL(=8=G?y5OuY_vUGKQEw&t#%^k^U&9quWc zsAm1Jl20<>9X1T}Jw)#Nh@2GpN$)3#0{70Zs0V90E32LAv;>${+ylL(1 zdQ>QTMJ#(oAWN3cub9AQ)o9~wl>&J{B$0t-I#9D2sM`qC34sPN0Ji-)V=o3@cPx|3 zQ`lg$XQh)b>VbQ;`-6VjwRYr5lk5^K$3@F=9^atcvba{dHp~Iu9C+@lUbZe<(^mJ& zwGC@IZ!K5&ytZMj;H?$UeKiQ@)1m-h_R9JU0Jw#sA+czPw+sQf0|$AsEL~Q?lO->L z4V%H%jbN(~Y!`!YSORJ6#z3%CE<0AH()Px+ByaD7JKYXPN>ry{yGjf1XnE1n1^bl1 zywuUeJDPw4MS)ej;5;Nc4=tU6G>rC5vUGziThEJWbkU@71C;%a4kF zShSuH>iWdGzNaU}xM0$yah{e4zWOP9S?lvX>{(eycp$hP5+4I{(A->wod= zq|h`XHjN16S&=-;lV=sN^GN4~rb}YeC4sywl9ze%vI41Xz5Czx{_9@;#AW`AAOwcR zz%WnxbI7UD^}~X@M|Af9uD^EEe{{otl&>3l)+zYUivF{Zg*@Hvrxw9IB)W%|&Z{`o zwSG?^heUFSCx;Zp?Ri=+ki#Mg2@Wf~o7Qd?Rudm+TkjLdL6IEf$w38nC*Qp-!&}y? zwyVXKfAFt!Lfc8P?W91S63J6Mc}hVIm1~>YYg~hh=n=^tp6to=u~$E=<*QnS@;0%& zjkmoj-vW6>B(Lz~l^4PG&0x<)utx~?ib3dd#m4T>5y-UW=e}An3&4kdEL*os7VE*Z zx9ritgMs%49}aGM>o&Y~g115RHf(xB8{UxMZ56$(o8ImXZ@1v>5xqUj!^^|Jd0t$x zMVOF*c^0G{f#p&3OS?)QIUhLRFMe3O={mgOIxKI>+H^H;xEke?Pc~gW8?GL~)hoJs zmxrE%9XlKy;(ug+U|&6za{q1F_shOl@lnNQb;m|^hfv)mR%ebhd0rBDcyr5W@D+Vd z0QfZm|Ib7em%KN<5_vc+6jg~uRa<7l=|-8{GDF_LZ2e0Y-i9@>0z=FI+^OsGbKUpj za@i$VPKcHhyyZmNQnYC)+pv_a)~CkTjtiD9(bC0Rz*4>I1!?>1(RW6_Humn=(%5rk zxZj+z1Ato~Ln0aC$q@EF^UZd!haY?}nCcU}Euyz&{gB}8e(Dpv1HA1_?iR=qksRU4 z5wIjEJG2>S+z2$TjqqLPg}?;&p&>Ce^sGh)&PsVk;G@AAt-1O9Lc%aW7f7&j1PKcfp(7&^&{9)9;z5SxM zA8^Z#Y?d`_l;NIC3T3CovePiG=1Ab4oE;&2)tkP<8@|K*k-=wP!FO8ponE^5B43Q- zPocAASd%LTwD-}ATNsKMV?ghsAK(YMc>Pz4?XJ=$a5lj zjwjD4(rj6;$Z<5&_g+VRuR|#B6w5n#+w1ZzkZ*|O8$9{Oi(uDgux}&SCjTPFW&W<3Cz-04D)AccIN~NWuBV_hDkC-bR=|F@VxN z(VPvj2-{A!7u2F{BrL3*vKIwgMV^>HhOH`}z@L19SbL|~7RsuLZQDm|*b<_O?Sw6k zMc>y^#ZO!xL!Lg#p=b8lH+^hDG2PjViMI@Lx<-OXI; zNu6TzpDurf5cnelCwx8eX;Vcf!Poo8xWlK^GO_Ldl=ZWaNpu3j?S&`v5?DGS4+ zq)}8LjJM>e7k~9O$r53v;r}EgXsKw3nO;DLEL=vfjC`y-gknFu30ksY%rJN+@VZ4Z zVxfm*l2_RIRVD08G+|sg6A)pmEI&!Kd}JiZ)sNA-rxd^v@!aW$i9O%;2JgHsIIoM& z>q}>H^S89cz0&#I=}No(Y3Jycl`sZACyYketOsBV0|$|vxTh_imEkP%p_d~2Rw#t* z@~2%#czmawEn9Y+gb5`90Onx;iWEusaU5UT347bO$xRXA@NC*EH|&*yy-KuKEe(H4 zmP76zmc!|14=*6sfwZUOVQ9+=CD4{7S1Q?z5H_ph1$UT1%5pHMMXhSAzJp1?li5Jv z9R!vTXafL;Jl;#7vozS5So8*iMnYKE-&YbK~9#x6oH77-a;8%TaWm`9MMneeVb+j|D zpfu3QH_S-{P9tyvfdK#oM_P7w=4CX>6Hi|447Y0D+f?sen((f$I=n{{-V<(WgPgfNOGRm~k0)Rstry-EPM@DA+=U3>3%Dq* zuLKZ)5YTYKp%DNuwVi=?sc6fT)al@yf$ibv*#Fnc(7!ZP>R=--fE1j00FQO~@yz&) z`b7tS`I>n7E&jq=;+ePk@Fag`QhbHNXj*V$NT+k5vk}m9E%TJ$RB%jB_{-D&vMuXt zdL10eiD0_Oy=6d-7jEB{5&Oe>%Clw0P7C4kDzW7_8OV^eIv96yU{pttj%9AcP{td| Z(xhc+_-kk0J+p*=TLywG2XEZe{y#v+kaz$9 literal 0 HcmV?d00001 diff --git a/models/rgliner/convert_to_gguf.py b/models/rgliner/convert_to_gguf.py new file mode 100644 index 000000000..ba0609fe0 --- /dev/null +++ b/models/rgliner/convert_to_gguf.py @@ -0,0 +1,475 @@ +#!/usr/bin/env python3 +""" +Convert GLiNER PyTorch models to GGUF format. + +Usage: + python convert_to_gguf.py --model knowledgator/gliner-bi-edge-v2.0 --output gliner-bi-edge-v2.0.gguf + +This script converts both: +1. The main text encoder (ModernBERT/Ettin) + span layer weights +2. The label encoder projection weights (sentence transformer is loaded separately) +""" + +import argparse +import json +import os +import struct +import sys +from pathlib import Path +from typing import Any, Dict, List, Tuple + +import numpy as np +import torch +from huggingface_hub import hf_hub_download, snapshot_download + + +# GGUF constants +GGUF_MAGIC = 0x46554747 # "GGUF" in little-endian +GGUF_VERSION = 3 + +# GGUF data types +GGUF_TYPE_UINT8 = 0 +GGUF_TYPE_INT8 = 1 +GGUF_TYPE_UINT16 = 2 +GGUF_TYPE_INT16 = 3 +GGUF_TYPE_UINT32 = 4 +GGUF_TYPE_INT32 = 5 +GGUF_TYPE_FLOAT32 = 6 +GGUF_TYPE_BOOL = 7 +GGUF_TYPE_STRING = 8 +GGUF_TYPE_ARRAY = 9 +GGUF_TYPE_UINT64 = 10 +GGUF_TYPE_INT64 = 11 +GGUF_TYPE_FLOAT64 = 12 + +# GGML tensor types +GGML_TYPE_F32 = 0 +GGML_TYPE_F16 = 1 +GGML_TYPE_Q4_0 = 2 +GGML_TYPE_Q4_1 = 3 +GGML_TYPE_Q5_0 = 6 +GGML_TYPE_Q5_1 = 7 +GGML_TYPE_Q8_0 = 8 +GGML_TYPE_Q8_1 = 9 +GGML_TYPE_BF16 = 30 + + +class GGUFWriter: + """Simple GGUF file writer.""" + + def __init__(self, path: str): + self.path = path + self.metadata: Dict[str, Any] = {} + self.tensors: List[Tuple[str, np.ndarray, int]] = [] # (name, data, ggml_type) + + def add_metadata(self, key: str, value: Any): + """Add metadata key-value pair.""" + self.metadata[key] = value + + def add_tensor(self, name: str, data: np.ndarray, ggml_type: int = GGML_TYPE_F32): + """Add a tensor.""" + self.tensors.append((name, data, ggml_type)) + + def _write_string(self, f, s: str): + """Write a GGUF string (length-prefixed UTF-8).""" + encoded = s.encode('utf-8') + f.write(struct.pack('> 16) & 0xFFFF).astype(np.uint16) + + # Write tensor info + # GGUF stores dimensions in reverse order (column-major) + # Reader reverses them back, so we write reversed to get original order + self._write_string(f, name) + f.write(struct.pack(' Tuple[Dict[str, torch.Tensor], Dict]: + """Load PyTorch model from HuggingFace.""" + print(f"Downloading model: {model_id}") + + # Download the model files + model_dir = snapshot_download( + model_id, + cache_dir=cache_dir, + allow_patterns=["*.bin", "*.json", "*.safetensors"] + ) + + # Load config + config_path = os.path.join(model_dir, "gliner_config.json") + with open(config_path, 'r') as f: + config = json.load(f) + + # Load weights + weights_path = os.path.join(model_dir, "pytorch_model.bin") + if os.path.exists(weights_path): + print(f"Loading weights from: {weights_path}") + state_dict = torch.load(weights_path, map_location='cpu', weights_only=True) + else: + # Try safetensors + from safetensors.torch import load_file + weights_path = os.path.join(model_dir, "model.safetensors") + print(f"Loading weights from: {weights_path}") + state_dict = load_file(weights_path) + + return state_dict, config + + +def map_weight_name(pytorch_name: str) -> str: + """Map PyTorch weight names to GGUF conventions.""" + name = pytorch_name + + # ===== Text Encoder (ModernBERT/Ettin) ===== + # token_rep_layer.bert_layer.model.embeddings.tok_embeddings.weight -> text.token_embd.weight + name = name.replace("token_rep_layer.bert_layer.model.embeddings.tok_embeddings.weight", "text.token_embd.weight") + name = name.replace("token_rep_layer.bert_layer.model.embeddings.norm.weight", "text.embd_norm.weight") + + # token_rep_layer.bert_layer.model.layers.X -> text.blk.X + name = name.replace("token_rep_layer.bert_layer.model.layers.", "text.blk.") + name = name.replace("token_rep_layer.bert_layer.model.final_norm.weight", "text.output_norm.weight") + + # ModernBERT attention (fused Wqkv) + name = name.replace(".attn.Wqkv.", ".attn_qkv.") + name = name.replace(".attn.Wo.", ".attn_output.") + + # ModernBERT FFN (GeGLU with fused Wi) + name = name.replace(".mlp.Wi.", ".ffn_gate_up.") # Fused gate+up + name = name.replace(".mlp.Wo.", ".ffn_down.") + name = name.replace(".mlp_norm.", ".ffn_norm.") + + # ===== Label Encoder (BERT/MiniLM) ===== + name = name.replace("token_rep_layer.labels_encoder.model.", "label.") + + # BERT embeddings (rbert expects token_types, token_embd_norm) + name = name.replace("label.embeddings.word_embeddings.", "label.token_embd.") + name = name.replace("label.embeddings.position_embeddings.", "label.position_embd.") + name = name.replace("label.embeddings.token_type_embeddings.", "label.token_types.") + name = name.replace("label.embeddings.LayerNorm.", "label.token_embd_norm.") + + # BERT layers + name = name.replace("label.encoder.layer.", "label.blk.") + + # BERT attention (rbert uses attn_output_norm, not attn_norm) + name = name.replace(".attention.self.query.", ".attn_q.") + name = name.replace(".attention.self.key.", ".attn_k.") + name = name.replace(".attention.self.value.", ".attn_v.") + name = name.replace(".attention.output.dense.", ".attn_output.") + name = name.replace(".attention.output.LayerNorm.", ".attn_output_norm.") + + # BERT FFN (rbert uses layer_output_norm, not ffn_norm) + name = name.replace(".intermediate.dense.", ".ffn_up.") + name = name.replace(".output.dense.", ".ffn_down.") + name = name.replace(".output.LayerNorm.", ".layer_output_norm.") + + # BERT pooler + name = name.replace("label.pooler.dense.", "label.pooler.") + + # ===== BiLSTM ===== + # Keep rnn.lstm.* as-is for now + name = name.replace("rnn.lstm.", "rnn.") + + # ===== Span Representation Layer ===== + name = name.replace("span_rep_layer.span_rep_layer.project_start.0.", "span.start_fc1.") + name = name.replace("span_rep_layer.span_rep_layer.project_start.3.", "span.start_fc2.") + name = name.replace("span_rep_layer.span_rep_layer.project_end.0.", "span.end_fc1.") + name = name.replace("span_rep_layer.span_rep_layer.project_end.3.", "span.end_fc2.") + name = name.replace("span_rep_layer.span_rep_layer.out_project.0.", "span.out_fc1.") + name = name.replace("span_rep_layer.span_rep_layer.out_project.3.", "span.out_fc2.") + + # ===== Prompt/Label Projection ===== + name = name.replace("prompt_rep_layer.0.", "label_proj.0.") + name = name.replace("prompt_rep_layer.3.", "label_proj.2.") + + return name + + +def convert_gliner_to_gguf( + model_id: str, + output_path: str, + quantize: str = "f32", + cache_dir: str = None +): + """Convert GLiNER model to GGUF format. + + Creates two GGUF files: + - {output_path}: Main model (text encoder, span layer, projection) + - {output_path_stem}-label-encoder.gguf: Label encoder (BERT/MiniLM) + """ + + # Load model + state_dict, config = load_pytorch_model(model_id, cache_dir) + + # Print model structure + print("\nModel weights:") + for name, tensor in state_dict.items(): + print(f" {name}: {tensor.shape} {tensor.dtype}") + + # Determine quantization type + if quantize == "f32": + ggml_type = GGML_TYPE_F32 + elif quantize == "f16": + ggml_type = GGML_TYPE_F16 + elif quantize == "bf16": + ggml_type = GGML_TYPE_BF16 + else: + raise ValueError(f"Unsupported quantization: {quantize}") + + # Separate label encoder weights from main model weights + label_encoder_weights = {} + main_model_weights = {} + + for pytorch_name, tensor in state_dict.items(): + if "token_rep_layer.labels_encoder" in pytorch_name: + label_encoder_weights[pytorch_name] = tensor + else: + main_model_weights[pytorch_name] = tensor + + # ============ Main Model GGUF ============ + writer = GGUFWriter(output_path) + + # Add metadata + writer.add_metadata("general.architecture", "gliner") + writer.add_metadata("general.name", model_id.split("/")[-1]) + writer.add_metadata("general.quantization_version", 2) + + # GLiNER-specific metadata + writer.add_metadata("gliner.max_width", config.get("max_width", 12)) + writer.add_metadata("gliner.span_mode", config.get("span_mode", "markerV0")) + writer.add_metadata("gliner.subtoken_pooling", config.get("subtoken_pooling", "first")) + + # Encoder config - use standard GGUF naming for model loading + encoder_config = config.get("encoder_config", {}) + hidden_size = encoder_config.get("hidden_size", 384) + num_heads = encoder_config.get("num_attention_heads", 6) + num_layers = encoder_config.get("num_hidden_layers", 10) + intermediate_size = encoder_config.get("intermediate_size", 576) + vocab_size = encoder_config.get("vocab_size", 50368) + context_length = encoder_config.get("max_position_embeddings", 8192) + rope_theta = encoder_config.get("local_rope_theta", 160000.0) + + # Standard GGUF metadata (without architecture prefix - the loader adds it) + writer.add_metadata("gliner.attention.head_count", num_heads) + writer.add_metadata("gliner.attention.head_count_kv", num_heads) # No GQA in this model + writer.add_metadata("gliner.block_count", num_layers) + writer.add_metadata("gliner.embedding_length", hidden_size) + writer.add_metadata("gliner.feed_forward_length", intermediate_size) + writer.add_metadata("gliner.context_length", context_length) + writer.add_metadata("gliner.rope.freq_base", float(rope_theta)) + writer.add_metadata("gliner.attention.layer_norm_rms_epsilon", 1e-5) + writer.add_metadata("gliner.vocab_size", vocab_size) + + # Convert main model tensors + print(f"\nConverting {len(main_model_weights)} main model tensors to GGUF...") + + for pytorch_name, tensor in main_model_weights.items(): + gguf_name = map_weight_name(pytorch_name) + + # Convert to numpy (workaround for numpy/torch incompatibility) + try: + data = tensor.detach().float().cpu().numpy() + except RuntimeError: + import array + t = tensor.detach().float().cpu().contiguous() + data = np.frombuffer( + array.array('f', t.flatten().tolist()), + dtype=np.float32 + ).reshape(t.shape) + + print(f" {pytorch_name} -> {gguf_name} {data.shape}") + writer.add_tensor(gguf_name, data, ggml_type) + + writer.write() + print(f"Main model output: {output_path}") + print(f"Size: {os.path.getsize(output_path) / 1024 / 1024:.2f} MB") + + # ============ Label Encoder GGUF ============ + # Create separate file for label encoder (without prefix, for rbert compatibility) + label_output_path = output_path.replace(".gguf", "-label-encoder.gguf") + label_writer = GGUFWriter(label_output_path) + + # Add BERT metadata + labels_config = config.get("labels_encoder_config", {}) + label_writer.add_metadata("general.architecture", "bert") + label_writer.add_metadata("general.name", model_id.split("/")[-1] + "-label-encoder") + + label_hidden = labels_config.get("hidden_size", 384) + label_heads = labels_config.get("num_attention_heads", 12) + label_layers = labels_config.get("num_hidden_layers", 6) + label_intermediate = labels_config.get("intermediate_size", 1536) + label_vocab = labels_config.get("vocab_size", 30522) + label_max_pos = labels_config.get("max_position_embeddings", 512) + + label_writer.add_metadata("bert.attention.head_count", label_heads) + label_writer.add_metadata("bert.block_count", label_layers) + label_writer.add_metadata("bert.embedding_length", label_hidden) + label_writer.add_metadata("bert.feed_forward_length", label_intermediate) + label_writer.add_metadata("bert.context_length", label_max_pos) + label_writer.add_metadata("bert.attention.layer_norm_epsilon", 1e-12) + label_writer.add_metadata("bert.vocab_size", label_vocab) + + # Convert label encoder tensors (remove prefix so rbert can load them) + print(f"\nConverting {len(label_encoder_weights)} label encoder tensors to GGUF...") + + for pytorch_name, tensor in label_encoder_weights.items(): + # Map name but remove the "label." prefix for rbert compatibility + gguf_name = map_weight_name(pytorch_name) + if gguf_name.startswith("label."): + gguf_name = gguf_name[6:] # Remove "label." prefix + + try: + data = tensor.detach().float().cpu().numpy() + except RuntimeError: + import array + t = tensor.detach().float().cpu().contiguous() + data = np.frombuffer( + array.array('f', t.flatten().tolist()), + dtype=np.float32 + ).reshape(t.shape) + + print(f" {pytorch_name} -> {gguf_name} {data.shape}") + label_writer.add_tensor(gguf_name, data, ggml_type) + + label_writer.write() + print(f"Label encoder output: {label_output_path}") + print(f"Size: {os.path.getsize(label_output_path) / 1024 / 1024:.2f} MB") + + print(f"\nConversion complete!") + + +def main(): + parser = argparse.ArgumentParser(description="Convert GLiNER PyTorch models to GGUF") + parser.add_argument( + "--model", "-m", + type=str, + required=True, + help="HuggingFace model ID (e.g., knowledgator/gliner-bi-edge-v2.0)" + ) + parser.add_argument( + "--output", "-o", + type=str, + required=True, + help="Output GGUF file path" + ) + parser.add_argument( + "--quantize", "-q", + type=str, + default="f32", + choices=["f32", "f16", "bf16"], + help="Quantization type (default: f32)" + ) + parser.add_argument( + "--cache-dir", + type=str, + default=None, + help="HuggingFace cache directory" + ) + + args = parser.parse_args() + + convert_gliner_to_gguf( + model_id=args.model, + output_path=args.output, + quantize=args.quantize, + cache_dir=args.cache_dir + ) + + +if __name__ == "__main__": + main() diff --git a/models/rgliner/examples/basic.rs b/models/rgliner/examples/basic.rs new file mode 100644 index 000000000..dd8af2f4a --- /dev/null +++ b/models/rgliner/examples/basic.rs @@ -0,0 +1,46 @@ +use rgliner::*; + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + // Use local GGUF files if GLINER_MODEL env var is set, otherwise try official HuggingFace edge variant + let source = if let Ok(model_path) = std::env::var("GLINER_MODEL") { + // Derive label encoder path from model path + let label_encoder_path = model_path.replace(".gguf", "-label-encoder.gguf"); + println!("Loading GLiNER model from: {}", model_path); + println!("Loading label encoder from: {}", label_encoder_path); + GlinerSource::local(model_path, label_encoder_path) + } else { + // Use official HuggingFace GGUF + println!("Loading GLiNER model from HuggingFace (edge variant)..."); + GlinerSource::edge() + }; + + let mut gliner = Gliner::builder().with_source(source).build().await?; + + println!("Model loaded!"); + + let labels = ["person", "organization", "location"]; + + // Test with multiple texts to see if the issue is consistent + let texts = [ + "Apple Inc. was founded by Steve Jobs in California.", + "Microsoft Corporation is headquartered in Seattle.", + "Elon Musk is the CEO of Tesla.", + "Google was founded in Mountain View.", + ]; + + for text in texts { + println!("\n--- Testing: {} ---", text); + let entities = gliner.extract(text, &labels).await?; + + println!("Found {} entities:", entities.len()); + for entity in entities { + println!( + " {}: '{}' (score: {:.2})", + entity.label, entity.text, entity.score + ); + } + } + + Ok(()) +} diff --git a/models/rgliner/src/config.rs b/models/rgliner/src/config.rs new file mode 100644 index 000000000..1f17a9ff6 --- /dev/null +++ b/models/rgliner/src/config.rs @@ -0,0 +1,143 @@ +//! GLiNER configuration parsing from gliner_config.json. + +use serde::Deserialize; + +/// GLiNER model configuration parsed from gliner_config.json. +/// +/// This differs from standard HuggingFace config.json and contains +/// GLiNER-specific parameters. +#[derive(Debug, Clone, Deserialize)] +pub struct GlinerConfig { + /// Text encoder model name (e.g., "jhu-clsp/ettin-encoder-32m") + #[serde(default)] + pub model_name: Option, + + /// Label encoder model name (e.g., "sentence-transformers/all-MiniLM-L6-v2") + /// If None, falls back to uni-encoder mode. + #[serde(default)] + pub labels_encoder: Option, + + /// Maximum span width in words (default: 12) + #[serde(default = "default_max_width")] + pub max_width: usize, + + /// Hidden dimension for span and label FFNs + #[serde(default = "default_hidden_size")] + pub hidden_size: usize, + + /// Dropout rate for FFN layers + #[serde(default = "default_dropout")] + pub dropout: f32, + + /// Subtoken pooling strategy: "first" or "mean" + #[serde(default = "default_subtoken_pooling")] + pub subtoken_pooling: String, + + /// Whether to enable cross-attention fusion (false for bi-encoder) + #[serde(default)] + pub fuse_layers: bool, + + /// Post-fusion schema (empty string for bi-encoder) + #[serde(default)] + pub post_fusion_schema: String, + + /// Span representation mode: "markerV0" for span-level + #[serde(default = "default_span_mode")] + pub span_mode: String, + + /// Index of the CLS token (typically 0, -1 means last token) + #[serde(default, deserialize_with = "deserialize_token_index")] + pub class_token_index: Option, + + /// Vocabulary size for output classes (-1 means use default) + #[serde(default, deserialize_with = "deserialize_optional_size")] + pub vocab_size: Option, +} + +fn deserialize_token_index<'de, D>(deserializer: D) -> Result, D::Error> +where + D: serde::Deserializer<'de>, +{ + let value: i64 = i64::deserialize(deserializer)?; + if value < 0 { + Ok(None) // -1 means last token or not applicable + } else { + Ok(Some(value as usize)) + } +} + +fn deserialize_optional_size<'de, D>(deserializer: D) -> Result, D::Error> +where + D: serde::Deserializer<'de>, +{ + let value: i64 = i64::deserialize(deserializer)?; + if value < 0 { + Ok(None) // -1 means use default or not applicable + } else { + Ok(Some(value as usize)) + } +} + +fn default_max_width() -> usize { + 12 +} + +fn default_hidden_size() -> usize { + 768 +} + +fn default_dropout() -> f32 { + 0.4 +} + +fn default_subtoken_pooling() -> String { + "first".to_string() +} + +fn default_span_mode() -> String { + "markerV0".to_string() +} + +fn default_vocab_size() -> usize { + 2 +} + +impl GlinerConfig { + /// Parse config from JSON bytes. + pub fn from_json(json: &[u8]) -> Result { + serde_json::from_slice(json) + } + + /// Check if this is a bi-encoder configuration. + pub fn is_bi_encoder(&self) -> bool { + self.labels_encoder.is_some() && self.post_fusion_schema.is_empty() + } + + /// Check if span mode is markerV0. + pub fn is_marker_v0(&self) -> bool { + self.span_mode == "markerV0" + } + + /// Check if subtoken pooling uses first token. + pub fn uses_first_subtoken(&self) -> bool { + self.subtoken_pooling == "first" + } +} + +impl Default for GlinerConfig { + fn default() -> Self { + Self { + model_name: None, + labels_encoder: None, + max_width: default_max_width(), + hidden_size: default_hidden_size(), + dropout: default_dropout(), + subtoken_pooling: default_subtoken_pooling(), + fuse_layers: false, + post_fusion_schema: String::new(), + span_mode: default_span_mode(), + class_token_index: Some(0), + vocab_size: Some(default_vocab_size()), + } + } +} diff --git a/models/rgliner/src/decoding.rs b/models/rgliner/src/decoding.rs new file mode 100644 index 000000000..6f1716006 --- /dev/null +++ b/models/rgliner/src/decoding.rs @@ -0,0 +1,255 @@ +//! Entity decoding with flat and nested NER support. + +/// A recognized named entity. +#[derive(Debug, Clone)] +pub struct Entity { + /// The entity text span. + pub text: String, + /// The entity label/type. + pub label: String, + /// Start character offset in the original text. + pub start_char: usize, + /// End character offset in the original text (exclusive). + pub end_char: usize, + /// Start word index. + pub start_word: usize, + /// End word index (inclusive). + pub end_word: usize, + /// Confidence score (0.0 to 1.0). + pub score: f32, +} + +/// Decoding mode for NER. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub enum DecodingMode { + /// Flat NER: no overlapping entities allowed. + /// Uses greedy non-maximum suppression. + #[default] + Flat, + /// Nested NER: overlapping entities allowed if one fully contains the other. + /// Partial overlaps are still forbidden. + Nested, +} + +/// Entity decoder using non-maximum suppression. +pub struct Decoder { + /// Confidence threshold for entity detection. + threshold: f32, + /// Decoding mode (flat or nested). + mode: DecodingMode, +} + +impl Decoder { + /// Create a new decoder with the given threshold and mode. + pub fn new(threshold: f32, mode: DecodingMode) -> Self { + Self { threshold, mode } + } + + /// Set the confidence threshold. + pub fn with_threshold(mut self, threshold: f32) -> Self { + self.threshold = threshold; + self + } + + /// Set the decoding mode. + pub fn with_mode(mut self, mode: DecodingMode) -> Self { + self.mode = mode; + self + } + + /// Decode entity predictions for a single text. + /// + /// # Arguments + /// * `scores` - Score matrix [num_spans, num_labels] after sigmoid + /// * `span_indices` - (start_word, end_word) for each span + /// * `word_offsets` - (start_char, end_char) for each word + /// * `labels` - Label strings + /// * `text` - Original input text + pub fn decode( + &self, + scores: &[f32], + num_spans: usize, + num_labels: usize, + span_indices: &[(usize, usize)], + word_offsets: &[(usize, usize)], + labels: &[&str], + text: &str, + ) -> Vec { + // Collect all predictions above threshold + let mut candidates: Vec<(usize, usize, f32)> = Vec::new(); // (span_idx, label_idx, score) + + for span_idx in 0..num_spans { + for label_idx in 0..num_labels { + let score = scores[span_idx * num_labels + label_idx]; + if score >= self.threshold { + candidates.push((span_idx, label_idx, score)); + } + } + } + + // Sort by score descending + candidates.sort_by(|a, b| b.2.partial_cmp(&a.2).unwrap_or(std::cmp::Ordering::Equal)); + + match self.mode { + DecodingMode::Flat => { + self.decode_flat(candidates, span_indices, word_offsets, labels, text) + } + DecodingMode::Nested => { + self.decode_nested(candidates, span_indices, word_offsets, labels, text) + } + } + } + + fn decode_flat( + &self, + candidates: Vec<(usize, usize, f32)>, + span_indices: &[(usize, usize)], + word_offsets: &[(usize, usize)], + labels: &[&str], + text: &str, + ) -> Vec { + let num_words = word_offsets.len(); + let mut entities = Vec::new(); + let mut used_positions: Vec = vec![false; num_words]; + + for (span_idx, label_idx, score) in candidates { + let (start_word, end_word) = span_indices[span_idx]; + + // Check if any word in span is already used + let overlaps = + (start_word..=end_word).any(|w| used_positions.get(w).copied().unwrap_or(false)); + + if !overlaps { + // Mark words as used + for w in start_word..=end_word { + if w < used_positions.len() { + used_positions[w] = true; + } + } + + if let Some(entity) = self.create_entity( + start_word, + end_word, + label_idx, + score, + word_offsets, + labels, + text, + ) { + entities.push(entity); + } + } + } + + // Sort by position + entities.sort_by_key(|e| e.start_char); + entities + } + + fn decode_nested( + &self, + candidates: Vec<(usize, usize, f32)>, + span_indices: &[(usize, usize)], + word_offsets: &[(usize, usize)], + labels: &[&str], + text: &str, + ) -> Vec { + let mut entities = Vec::new(); + // Track selected (start, end, label) triples to avoid exact duplicates + let mut selected: Vec<(usize, usize, usize)> = Vec::new(); + + for (span_idx, label_idx, score) in candidates { + let (start_word, end_word) = span_indices[span_idx]; + let key = (start_word, end_word, label_idx); + + // Check for partial overlap with already selected entities + let has_partial_overlap = selected.iter().any(|(sel_start, sel_end, _)| { + self.is_partial_overlap(start_word, end_word, *sel_start, *sel_end) + }); + + // For nested NER, allow if no partial overlap and not exact duplicate + if !has_partial_overlap && !selected.contains(&key) { + selected.push(key); + + if let Some(entity) = self.create_entity( + start_word, + end_word, + label_idx, + score, + word_offsets, + labels, + text, + ) { + entities.push(entity); + } + } + } + + // Sort by position, then by span length descending (outer spans first) + entities.sort_by(|a, b| { + a.start_char + .cmp(&b.start_char) + .then_with(|| b.end_char.cmp(&a.end_char)) + }); + entities + } + + /// Check if two spans have partial overlap (overlap but neither contains the other). + fn is_partial_overlap(&self, start1: usize, end1: usize, start2: usize, end2: usize) -> bool { + // Check if spans overlap at all + let overlaps = start1 <= end2 && start2 <= end1; + if !overlaps { + return false; + } + + // Check if one fully contains the other (not partial overlap) + let one_contains_other = + (start1 <= start2 && end1 >= end2) || (start2 <= start1 && end2 >= end1); + + overlaps && !one_contains_other + } + + fn create_entity( + &self, + start_word: usize, + end_word: usize, + label_idx: usize, + score: f32, + word_offsets: &[(usize, usize)], + labels: &[&str], + text: &str, + ) -> Option { + if start_word >= word_offsets.len() || end_word >= word_offsets.len() { + return None; + } + if label_idx >= labels.len() { + return None; + } + + let start_char = word_offsets[start_word].0; + let end_char = word_offsets[end_word].1; + + if end_char > text.len() { + return None; + } + + Some(Entity { + text: text[start_char..end_char].to_string(), + label: labels[label_idx].to_string(), + start_char, + end_char, + start_word, + end_word, + score, + }) + } +} + +impl Default for Decoder { + fn default() -> Self { + Self { + threshold: 0.5, + mode: DecodingMode::default(), + } + } +} diff --git a/models/rgliner/src/error.rs b/models/rgliner/src/error.rs new file mode 100644 index 000000000..0804efec6 --- /dev/null +++ b/models/rgliner/src/error.rs @@ -0,0 +1,43 @@ +//! Error types for rgliner. + +use kalosm_common::CacheError; + +/// An error that can occur when loading a GLiNER model. +#[derive(Debug, thiserror::Error)] +pub enum GlinerLoadingError { + /// An error that can occur when trying to download model files. + #[error("Failed to download model files: {0}")] + DownloadingError(#[from] CacheError), + /// An error that can occur when trying to load the model. + #[error("Failed to load model: {0}")] + LoadModel(#[from] fusor::Error), + /// An IO error. + #[error("IO error: {0}")] + Io(#[from] std::io::Error), + /// An error that can occur when trying to load the tokenizer. + #[error("Failed to load tokenizer: {0}")] + LoadTokenizer(tokenizers::Error), + /// An error that can occur when trying to load the config. + #[error("Failed to load config: {0}")] + LoadConfig(serde_json::Error), + /// Config file not found. + #[error("Config file not found")] + ConfigNotFound, + /// Label encoder loading error. + #[error("Failed to load label encoder: {0}")] + LabelEncoder(#[from] rbert::BertLoadingError), +} + +/// An error that can occur when running GLiNER inference. +#[derive(Debug, thiserror::Error)] +pub enum GlinerError { + /// An error that can occur when running tensor operations. + #[error("Tensor operation error: {0}")] + Fusor(#[from] fusor::Error), + /// An error that can occur when tokenizing text. + #[error("Tokenization error: {0}")] + Tokenizer(tokenizers::Error), + /// An error that can occur with the label encoder. + #[error("Label encoder error: {0}")] + LabelEncoder(#[from] rbert::BertError), +} diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs new file mode 100644 index 000000000..e4002b833 --- /dev/null +++ b/models/rgliner/src/lib.rs @@ -0,0 +1,1467 @@ +//! # rgliner +//! +//! GLiNER bi-encoder Named Entity Recognition for Rust. +//! +//! GLiNER (Generalist Lightweight Model for Named Entity Recognition) identifies +//! arbitrary entity types at inference time using natural language labels. +//! +//! ## Usage +//! +//! ```rust, no_run +//! use rgliner::*; +//! +//! #[tokio::main] +//! async fn main() -> anyhow::Result<()> { +//! let mut gliner = Gliner::new().await?; +//! +//! let labels = ["person", "organization", "location"]; +//! let text = "Apple Inc. was founded by Steve Jobs in California."; +//! +//! let entities = gliner.extract(text, &labels).await?; +//! for entity in entities { +//! println!("{}: {} ({:.2})", entity.label, entity.text, entity.score); +//! } +//! Ok(()) +//! } +//! ``` +//! +//! ## Label Caching +//! +//! For production workloads with fixed label sets, you can pre-compute label +//! embeddings for significant speedup: +//! +//! ```rust, no_run +//! use rgliner::*; +//! +//! # async fn example() -> anyhow::Result<()> { +//! let mut gliner = Gliner::new().await?; +//! +//! // Pre-compute label embeddings once +//! let labels = ["person", "organization", "location"]; +//! gliner.cache_labels(&labels).await?; +//! +//! // Fast inference with cached labels +//! for text in documents { +//! let entities = gliner.extract_with_cached_labels(&text).await?; +//! // Process entities... +//! } +//! # Ok(()) +//! # } +//! ``` + +#![warn(missing_docs)] + +mod config; +mod decoding; +mod error; +mod raw; +mod source; +mod tokenization; + +pub use config::GlinerConfig; +pub use decoding::{Decoder, DecodingMode, Entity}; +pub use error::{GlinerError, GlinerLoadingError}; +pub use raw::modern_bert::{ModernBertConfig, ModernBertModel}; +pub use source::GlinerSource; + +use fusor::{Device, Tensor, VarBuilder}; +use kalosm_common::Cache; +use kalosm_model_types::ModelLoadingProgress; +use rbert::BertSource; +use std::sync::{Arc, Mutex, OnceLock}; +use tokenizers::Tokenizer; + +use crate::raw::{CachedLabels, LabelEncoder, Scorer, SpanLayer, TextEncoder}; +use crate::tokenization::{first_subtoken_pooling, WordTokenizer}; + +fn default_device() -> Device { + static PANIC_HOOK_LOCK: OnceLock> = OnceLock::new(); + + let lock = PANIC_HOOK_LOCK.get_or_init(|| Mutex::new(())); + let _guard = lock.lock().unwrap_or_else(|poisoned| poisoned.into_inner()); + + let hook = std::panic::take_hook(); + std::panic::set_hook(Box::new(|_| {})); + let result = std::panic::catch_unwind(Device::gpu_blocking); + std::panic::set_hook(hook); + + result.ok().and_then(Result::ok).unwrap_or_else(Device::cpu) +} + +/// Builder for constructing a [`Gliner`] model. +#[derive(Default)] +pub struct GlinerBuilder { + source: GlinerSource, + cache: Cache, + device: Option, + decoding_mode: DecodingMode, + threshold: f32, + max_width: Option, +} + +impl GlinerBuilder { + /// Set the model source. + pub fn with_source(mut self, source: GlinerSource) -> Self { + self.source = source; + self + } + + /// Set the decoding mode (Flat or Nested). + pub fn with_decoding_mode(mut self, mode: DecodingMode) -> Self { + self.decoding_mode = mode; + self + } + + /// Set the confidence threshold (default 0.5). + pub fn with_threshold(mut self, threshold: f32) -> Self { + self.threshold = threshold; + self + } + + /// Set the maximum span width (overrides config). + pub fn with_max_width(mut self, max_width: usize) -> Self { + self.max_width = Some(max_width); + self + } + + /// Set the device. + pub fn with_device(mut self, device: Device) -> Self { + self.device = Some(device); + self + } + + /// Set the cache location. + #[cfg(feature = "tokio")] + pub fn with_cache(mut self, cache: Cache) -> Self { + self.cache = cache; + self + } + + /// Build the model. + pub async fn build(self) -> Result { + self.build_with_loading_handler(ModelLoadingProgress::multi_bar_loading_indicator()) + .await + } + + /// Build the model with a loading handler. + pub async fn build_with_loading_handler( + self, + loading_handler: impl FnMut(ModelLoadingProgress) + Send + 'static, + ) -> Result { + Gliner::from_builder(self, loading_handler).await + } +} + +/// GLiNER Named Entity Recognition model. +/// +/// The bi-encoder architecture enables efficient NER with arbitrary entity types. +/// Labels are encoded independently from text, allowing pre-computation and caching. +pub struct Gliner { + text_encoder: TextEncoder, + label_encoder: LabelEncoder, + span_layer: SpanLayer, + tokenizer: Arc, + decoder: Decoder, + device: Device, + max_width: usize, + /// Cached label embeddings for repeated inference. + cached_labels: Option, +} + +impl Gliner { + /// Create a new builder. + pub fn builder() -> GlinerBuilder { + GlinerBuilder { + threshold: 0.5, + ..Default::default() + } + } + + /// Create with default settings (base model). + pub async fn new() -> Result { + Self::builder().build().await + } + + async fn from_builder( + builder: GlinerBuilder, + mut progress_handler: impl FnMut(ModelLoadingProgress) + Send + 'static, + ) -> Result { + let GlinerBuilder { + source, + cache, + device, + decoding_mode, + threshold, + max_width: max_width_override, + } = builder; + + // Download config file + let config_source = format!("Config ({})", source.config); + let mut create_progress = ModelLoadingProgress::downloading_progress(config_source); + let config_bytes = cache + .get_bytes(&source.config, |progress| { + progress_handler(create_progress(progress)) + }) + .await?; + + let config = + GlinerConfig::from_json(&config_bytes).map_err(GlinerLoadingError::LoadConfig)?; + + // Download tokenizer + let tokenizer_source = format!("Tokenizer ({})", source.tokenizer); + let mut create_progress = ModelLoadingProgress::downloading_progress(tokenizer_source); + let tokenizer_bytes = cache + .get_bytes(&source.tokenizer, |progress| { + progress_handler(create_progress(progress)) + }) + .await?; + + let tokenizer = + Tokenizer::from_bytes(&tokenizer_bytes).map_err(GlinerLoadingError::LoadTokenizer)?; + let word_tokenizer = WordTokenizer::new(tokenizer); + + // Download main model weights + let model_source = format!("Text Encoder ({})", source.model); + let mut create_progress = ModelLoadingProgress::downloading_progress(model_source); + let model_bytes = cache + .get_bytes(&source.model, |progress| { + progress_handler(create_progress(progress)) + }) + .await?; + + // Download label encoder weights + let label_source = format!("Label Encoder ({})", source.label_encoder); + let mut create_progress = ModelLoadingProgress::downloading_progress(label_source); + let _label_bytes = cache + .get_bytes(&source.label_encoder, |progress| { + progress_handler(create_progress(progress)) + }) + .await?; + + // Initialize device + let device = match device { + Some(device) => device, + None => default_device(), + }; + + // Load text encoder + let mut model_cursor = std::io::Cursor::new(&model_bytes); + let mut text_vb = VarBuilder::from_gguf(&mut model_cursor) + .map_err(|err| GlinerLoadingError::LoadModel(fusor::Error::from(err)))?; + + let text_encoder = TextEncoder::load(&device, &mut text_vb)?; + + // Load span layer from main model weights + let max_width = max_width_override.unwrap_or(config.max_width); + let span_layer = SpanLayer::load(&device, &mut text_vb, max_width)?; + + // Load label encoder + let label_encoder_source = BertSource::new() + .with_model(source.label_encoder.clone()) + .with_config(source.label_encoder_config.clone()) + .with_tokenizer(source.label_encoder_tokenizer.clone()); + + // Create projection VarBuilder from main model + let mut model_cursor2 = std::io::Cursor::new(&model_bytes); + let mut proj_vb = VarBuilder::from_gguf(&mut model_cursor2) + .map_err(|err| GlinerLoadingError::LoadModel(fusor::Error::from(err)))?; + + let label_encoder = LabelEncoder::load(&device, &mut proj_vb, label_encoder_source).await?; + + let decoder = Decoder::new(threshold, decoding_mode); + + Ok(Self { + text_encoder, + label_encoder, + span_layer, + tokenizer: Arc::new(word_tokenizer), + decoder, + device, + max_width, + cached_labels: None, + }) + } + + /// Cache label embeddings for repeated inference with the same labels. + /// + /// This significantly speeds up inference when using fixed label sets. + pub async fn cache_labels(&mut self, labels: &[&str]) -> Result<(), GlinerError> { + let label_embeddings = self + .label_encoder + .encode_labels(labels) + .await? + .to_concrete(); + self.cached_labels = Some(CachedLabels::new( + labels.iter().map(|s| s.to_string()).collect(), + label_embeddings, + )); + Ok(()) + } + + /// Clear cached label embeddings. + pub fn clear_label_cache(&mut self) { + self.cached_labels = None; + } + + /// Check if labels are cached. + pub fn has_cached_labels(&self) -> bool { + self.cached_labels.is_some() + } + + /// Extract named entities from text. + pub async fn extract( + &mut self, + text: &str, + labels: &[&str], + ) -> Result, GlinerError> { + let mut results = self.extract_batch(&[text], labels).await?; + Ok(results.pop().unwrap_or_default()) + } + + /// Extract named entities using cached labels. + /// + /// Panics if no labels are cached. + pub async fn extract_with_cached_labels( + &mut self, + text: &str, + ) -> Result, GlinerError> { + let labels: Vec = self + .cached_labels + .as_ref() + .expect("No labels cached. Call cache_labels first.") + .labels + .clone(); + let labels: Vec<&str> = labels.iter().map(|label| label.as_str()).collect(); + self.extract(text, &labels).await + } + + /// Extract named entities from a batch of texts. + pub async fn extract_batch( + &mut self, + texts: &[&str], + labels: &[&str], + ) -> Result>, GlinerError> { + if texts.is_empty() { + return Ok(Vec::new()); + } + + // Get label embeddings (compute if not cached or labels differ) + let label_embeddings = if let Some(ref cached) = self.cached_labels { + let cached_labels: Vec<&str> = cached.labels.iter().map(|s| s.as_str()).collect(); + if cached_labels == labels { + cached.embeddings.clone() + } else { + self.label_encoder.encode_labels(labels).await? + } + } else { + self.label_encoder.encode_labels(labels).await? + }; + + let mut results = Vec::with_capacity(texts.len()); + for text in texts { + let entities = self + .extract_internal(text, labels, &label_embeddings) + .await?; + results.push(entities); + } + + Ok(results) + } + + async fn extract_internal( + &self, + text: &str, + labels: &[&str], + label_embeddings: &Tensor<2, f32>, + ) -> Result, GlinerError> { + // 1. Tokenize text + let tokenized = self.tokenizer.tokenize(text)?; + + if tokenized.num_words == 0 { + return Ok(Vec::new()); + } + + // 2. Prepare input tensors + let token_ids = Tensor::new(&self.device, &tokenized.token_ids); + let token_ids: Tensor<2, u32> = token_ids.unsqueeze(0).to_concrete(); + + let attention_mask = Tensor::new(&self.device, &tokenized.attention_mask); + let attention_mask: Tensor<2, u32> = attention_mask.unsqueeze(0).to_concrete(); + + // 3. Encode text + let token_embeddings = self.text_encoder.forward(&token_ids, Some(&attention_mask)); + + // Python's bi-encoder span model pools transformer token embeddings + // directly to words; the checkpoint still contains LSTM weights, but + // that path is not used in BaseBiEncoderModel.get_representations(). + let (word_embeddings, _word_mask) = + first_subtoken_pooling(&token_embeddings, &[tokenized.clone()], &self.device); + + // 4. Generate span representations + let (span_embeddings, span_indices) = + self.span_layer.forward(&word_embeddings, &self.device); + + // 5. Score spans against labels + let scores = Scorer::forward(&span_embeddings, label_embeddings); + + // 6. Decode predictions + let shape = scores.shape(); + let num_spans = shape[1]; + let num_labels = shape[2]; + + // Get scores for first batch item and apply sigmoid + let flat_scores: Tensor<2, f32> = scores.squeeze(0).to_concrete(); + let tensor_slice = flat_scores.as_slice().await?; + let scores_data: Vec = tensor_slice + .as_slice() + .iter() + .map(|&x| 1.0 / (1.0 + (-x).exp())) // sigmoid + .collect(); + + let entities = self.decoder.decode( + &scores_data, + num_spans, + num_labels, + &span_indices, + &tokenized.word_offsets, + labels, + text, + ); + + Ok(entities) + } + + /// Get the maximum span width. + pub fn max_width(&self) -> usize { + self.max_width + } + + /// Get the device. + pub fn device(&self) -> &Device { + &self.device + } +} + +#[cfg(test)] +mod gpu_parity_tests { + use super::*; + use fusor::layers::{Embedding, LayerNorm}; + use std::path::Path; + + fn local_edge_source() -> Option { + let weights_dir = Path::new(env!("CARGO_MANIFEST_DIR")).join("weights"); + let model_path = weights_dir.join("gliner-edge.gguf"); + let label_encoder_path = weights_dir.join("gliner-edge-label-encoder.gguf"); + if model_path.exists() && label_encoder_path.exists() { + Some(GlinerSource::local(model_path, label_encoder_path)) + } else { + None + } + } + + async fn load_local_edge(device: Device) -> Result { + Gliner::builder() + .with_source(local_edge_source().expect("local edge checkpoint missing")) + .with_device(device) + .build() + .await + } + + async fn tensor_values(tensor: &Tensor) -> fusor::Result> { + let slice = tensor.clone().as_slice().await?; + Ok(slice.as_slice().iter().copied().collect()) + } + + fn load_qmatrix_tensor_from_gguf( + gguf_bytes: &[u8], + scope: &[&str], + tensor_name: &str, + device: &Device, + ) -> fusor::QMatrix { + let mut cursor = std::io::Cursor::new(gguf_bytes); + let mut vb = VarBuilder::from_gguf(&mut cursor).unwrap(); + let mut vb = if scope.is_empty() { + vb + } else { + vb.pp(scope.join(".")) + }; + vb.get(tensor_name, device).unwrap() + } + + fn load_qmatrix_from_gguf( + gguf_bytes: &[u8], + scope: &[&str], + device: &Device, + ) -> fusor::QMatrix { + load_qmatrix_tensor_from_gguf(gguf_bytes, scope, "weight", device) + } + + fn load_layer_norm_from_gguf( + gguf_bytes: &[u8], + scope: &[&str], + device: &Device, + eps: f32, + ) -> LayerNorm<1, f32> { + let mut cursor = std::io::Cursor::new(gguf_bytes); + let mut vb = VarBuilder::from_gguf(&mut cursor).unwrap(); + let scope = if scope.is_empty() { + None + } else { + Some(scope.join(".")) + }; + if let Some(scope) = scope { + let mut scoped = vb.pp(scope); + LayerNorm::load(device, &mut scoped, eps).unwrap() + } else { + LayerNorm::load(device, &mut vb, eps).unwrap() + } + } + + fn load_modern_bert_config_from_gguf(gguf_bytes: &[u8]) -> ModernBertConfig { + let mut cursor = std::io::Cursor::new(gguf_bytes); + let mut vb = VarBuilder::from_gguf(&mut cursor).unwrap(); + ModernBertConfig::from_gguf(&mut vb.pp("text")).unwrap() + } + + async fn print_diff( + name: &str, + cpu: &Tensor, + gpu: &Tensor, + ) -> fusor::Result { + assert_eq!(cpu.shape(), gpu.shape(), "{name} shape mismatch"); + + let cpu_values = tensor_values(cpu).await?; + let gpu_values = tensor_values(gpu).await?; + + let mut max_abs_diff = 0.0f32; + let mut mean_abs_diff = 0.0f32; + + for (cpu, gpu) in cpu_values.iter().zip(&gpu_values) { + let diff = (cpu - gpu).abs(); + max_abs_diff = max_abs_diff.max(diff); + mean_abs_diff += diff; + } + + mean_abs_diff /= cpu_values.len().max(1) as f32; + + println!( + "{name}: shape={:?}, max_abs_diff={max_abs_diff:.6}, mean_abs_diff={mean_abs_diff:.6}", + cpu.shape() + ); + + Ok(max_abs_diff) + } + + fn build_text_inputs( + tokenized: &crate::tokenization::TokenizedText, + device: &Device, + ) -> (Tensor<2, u32>, Tensor<2, u32>) { + let token_ids = Tensor::new(device, &tokenized.token_ids); + let token_ids: Tensor<2, u32> = token_ids.unsqueeze(0).to_concrete(); + + let attention_mask = Tensor::new(device, &tokenized.attention_mask); + let attention_mask: Tensor<2, u32> = attention_mask.unsqueeze(0).to_concrete(); + + (token_ids, attention_mask) + } + + #[tokio::test] + #[ignore = "requires local edge checkpoint files and a working GPU device"] + async fn debug_cpu_gpu_parity_for_local_edge_checkpoint() { + if local_edge_source().is_none() { + eprintln!("Skipping GPU parity test: local edge checkpoint files are missing."); + return; + } + + let gpu_device = match std::panic::catch_unwind(Device::gpu_blocking) { + Ok(Ok(device)) => device, + Ok(Err(err)) => { + eprintln!("Skipping GPU parity test: failed to create GPU device: {err}"); + return; + } + Err(_) => { + eprintln!("Skipping GPU parity test: GPU device creation panicked."); + return; + } + }; + + let cpu_device = Device::cpu(); + let weights_dir = Path::new(env!("CARGO_MANIFEST_DIR")).join("weights"); + let model_bytes = std::fs::read(weights_dir.join("gliner-edge.gguf")).unwrap(); + let label_encoder_bytes = + std::fs::read(weights_dir.join("gliner-edge-label-encoder.gguf")).unwrap(); + + let mut cpu = load_local_edge(cpu_device.clone()).await.unwrap(); + let mut gpu = load_local_edge(gpu_device.clone()).await.unwrap(); + + let labels = ["person", "organization", "location"]; + let text = "Google was founded in Mountain View."; + + let cpu_tokenized = cpu.tokenizer.tokenize(text).unwrap(); + let gpu_tokenized = gpu.tokenizer.tokenize(text).unwrap(); + assert_eq!(cpu_tokenized.token_ids, gpu_tokenized.token_ids); + assert_eq!(cpu_tokenized.attention_mask, gpu_tokenized.attention_mask); + assert_eq!(cpu_tokenized.word_offsets, gpu_tokenized.word_offsets); + + let cpu_text_embedding = Embedding::new(load_qmatrix_from_gguf( + &model_bytes, + &["text", "token_embd"], + &cpu_device, + )); + let gpu_text_embedding = Embedding::new(load_qmatrix_from_gguf( + &model_bytes, + &["text", "token_embd"], + &gpu_device, + )); + let cpu_text_ids = Tensor::new(&cpu_device, &cpu_tokenized.token_ids); + let gpu_text_ids = Tensor::new(&gpu_device, &gpu_tokenized.token_ids); + let cpu_text_embedding_lookup: Tensor<2, f32> = cpu_text_embedding.forward(&cpu_text_ids); + let gpu_text_embedding_lookup: Tensor<2, f32> = gpu_text_embedding.forward(&gpu_text_ids); + let _ = print_diff( + "raw_text_token_embedding_lookup", + &cpu_text_embedding_lookup, + &gpu_text_embedding_lookup, + ) + .await + .unwrap(); + + let cpu_text_embd_norm = + load_layer_norm_from_gguf(&model_bytes, &["text", "embd_norm"], &cpu_device, 1e-6); + let gpu_text_embd_norm = + load_layer_norm_from_gguf(&model_bytes, &["text", "embd_norm"], &gpu_device, 1e-6); + let cpu_text_embedding_lookup_3d: Tensor<3, f32> = + cpu_text_embedding_lookup.unsqueeze(0).to_concrete(); + let gpu_text_embedding_lookup_3d: Tensor<3, f32> = + gpu_text_embedding_lookup.unsqueeze(0).to_concrete(); + let cpu_text_after_embd_norm = cpu_text_embd_norm.forward(&cpu_text_embedding_lookup_3d); + let gpu_text_after_embd_norm = gpu_text_embd_norm.forward(&gpu_text_embedding_lookup_3d); + let _ = print_diff( + "text_after_embd_norm", + &cpu_text_after_embd_norm, + &gpu_text_after_embd_norm, + ) + .await + .unwrap(); + + let cpu_layer0_qkv = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.0"], + "attn_qkv.weight", + &cpu_device, + ); + let gpu_layer0_qkv = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.0"], + "attn_qkv.weight", + &gpu_device, + ); + let cpu_layer0_qkv_projection = cpu_text_after_embd_norm.q_mat_mul(&cpu_layer0_qkv); + let gpu_layer0_qkv_projection = gpu_text_after_embd_norm.q_mat_mul(&gpu_layer0_qkv); + let _ = print_diff( + "text_layer0_qkv_projection", + &cpu_layer0_qkv_projection, + &gpu_layer0_qkv_projection, + ) + .await + .unwrap(); + + let label_probe_indices = [0u32, 1, 2, 17, 101, 257, 1024]; + let cpu_label_token_embedding = Embedding::new(load_qmatrix_from_gguf( + &label_encoder_bytes, + &["token_embd"], + &cpu_device, + )); + let gpu_label_token_embedding = Embedding::new(load_qmatrix_from_gguf( + &label_encoder_bytes, + &["token_embd"], + &gpu_device, + )); + let cpu_label_ids = Tensor::new(&cpu_device, &label_probe_indices); + let gpu_label_ids = Tensor::new(&gpu_device, &label_probe_indices); + let cpu_label_embedding_lookup: Tensor<2, f32> = + cpu_label_token_embedding.forward(&cpu_label_ids); + let gpu_label_embedding_lookup: Tensor<2, f32> = + gpu_label_token_embedding.forward(&gpu_label_ids); + let _ = print_diff( + "raw_label_token_embedding_lookup", + &cpu_label_embedding_lookup, + &gpu_label_embedding_lookup, + ) + .await + .unwrap(); + + let single_label = ["organization"]; + let cpu_single_label_embeddings = cpu.label_encoder.encode_labels(&single_label).await.unwrap(); + let gpu_single_label_embeddings = gpu.label_encoder.encode_labels(&single_label).await.unwrap(); + let _ = print_diff( + "single_label_embeddings", + &cpu_single_label_embeddings, + &gpu_single_label_embeddings, + ) + .await + .unwrap(); + + let (cpu_label_token_embeddings, cpu_label_attention_mask) = cpu + .label_encoder + .debug_sentence_token_embeddings_and_mask(&labels) + .unwrap(); + let (gpu_label_token_embeddings, gpu_label_attention_mask) = gpu + .label_encoder + .debug_sentence_token_embeddings_and_mask(&labels) + .unwrap(); + let _ = print_diff( + "label_token_embeddings", + &cpu_label_token_embeddings, + &gpu_label_token_embeddings, + ) + .await + .unwrap(); + let _ = print_diff( + "label_attention_mask", + &cpu_label_attention_mask.cast(), + &gpu_label_attention_mask.cast(), + ) + .await + .unwrap(); + + let (cpu_label_states, _) = cpu.label_encoder.debug_sentence_hidden_states(&labels).unwrap(); + let (gpu_label_states, _) = gpu.label_encoder.debug_sentence_hidden_states(&labels).unwrap(); + for (idx, (cpu_state, gpu_state)) in cpu_label_states.iter().zip(&gpu_label_states).enumerate() + { + let name = if idx == 0 { + "label_post_embeddings".to_string() + } else { + format!("label_post_layer_{}", idx - 1) + }; + let _ = print_diff(&name, cpu_state, gpu_state).await.unwrap(); + } + + let (cpu_label_layer0_attention, cpu_label_layer0_intermediate, cpu_label_layer0_output) = + cpu.label_encoder.debug_sentence_first_layer(&labels).unwrap(); + let (gpu_label_layer0_attention, gpu_label_layer0_intermediate, gpu_label_layer0_output) = + gpu.label_encoder.debug_sentence_first_layer(&labels).unwrap(); + let _ = print_diff( + "label_layer0_attention_output", + &cpu_label_layer0_attention, + &gpu_label_layer0_attention, + ) + .await + .unwrap(); + let ( + cpu_label_layer0_query, + cpu_label_layer0_key, + cpu_label_layer0_value, + cpu_label_layer0_self_output, + cpu_label_layer0_attention_output_debug, + ) = cpu + .label_encoder + .debug_sentence_first_layer_attention(&labels) + .unwrap(); + let ( + gpu_label_layer0_query, + gpu_label_layer0_key, + gpu_label_layer0_value, + gpu_label_layer0_self_output, + gpu_label_layer0_attention_output_debug, + ) = gpu + .label_encoder + .debug_sentence_first_layer_attention(&labels) + .unwrap(); + let _ = print_diff( + "label_layer0_query", + &cpu_label_layer0_query, + &gpu_label_layer0_query, + ) + .await + .unwrap(); + let _ = print_diff( + "label_layer0_key", + &cpu_label_layer0_key, + &gpu_label_layer0_key, + ) + .await + .unwrap(); + let _ = print_diff( + "label_layer0_value", + &cpu_label_layer0_value, + &gpu_label_layer0_value, + ) + .await + .unwrap(); + let _ = print_diff( + "label_layer0_self_attention_output", + &cpu_label_layer0_self_output, + &gpu_label_layer0_self_output, + ) + .await + .unwrap(); + let _ = print_diff( + "label_layer0_attention_output_debug_split", + &cpu_label_layer0_attention_output_debug, + &gpu_label_layer0_attention_output_debug, + ) + .await + .unwrap(); + let _ = print_diff( + "label_layer0_intermediate_output", + &cpu_label_layer0_intermediate, + &gpu_label_layer0_intermediate, + ) + .await + .unwrap(); + let _ = print_diff( + "label_layer0_output_debug", + &cpu_label_layer0_output, + &gpu_label_layer0_output, + ) + .await + .unwrap(); + + let cpu_label_mean_pool = cpu.label_encoder.debug_sentence_mean_pool(&labels).unwrap(); + let gpu_label_mean_pool = gpu.label_encoder.debug_sentence_mean_pool(&labels).unwrap(); + let _ = print_diff( + "label_mean_pool", + &cpu_label_mean_pool, + &gpu_label_mean_pool, + ) + .await + .unwrap(); + + let cpu_label_sentence_embeddings = cpu + .label_encoder + .debug_sentence_embeddings(&labels) + .await + .unwrap(); + let gpu_label_sentence_embeddings = gpu + .label_encoder + .debug_sentence_embeddings(&labels) + .await + .unwrap(); + let _ = print_diff( + "label_sentence_embeddings", + &cpu_label_sentence_embeddings, + &gpu_label_sentence_embeddings, + ) + .await + .unwrap(); + + let cpu_projected_label_embeddings = cpu + .label_encoder + .debug_projection(&cpu_label_sentence_embeddings); + let gpu_projected_label_embeddings = gpu + .label_encoder + .debug_projection(&gpu_label_sentence_embeddings); + let _ = print_diff( + "label_projected_embeddings", + &cpu_projected_label_embeddings, + &gpu_projected_label_embeddings, + ) + .await + .unwrap(); + + let cpu_label_embeddings = cpu.label_encoder.encode_labels(&labels).await.unwrap(); + let gpu_label_embeddings = gpu.label_encoder.encode_labels(&labels).await.unwrap(); + let _ = print_diff("label_embeddings", &cpu_label_embeddings, &gpu_label_embeddings) + .await + .unwrap(); + + let (cpu_token_ids, cpu_attention_mask) = build_text_inputs(&cpu_tokenized, &cpu_device); + let (gpu_token_ids, gpu_attention_mask) = build_text_inputs(&gpu_tokenized, &gpu_device); + + let cpu_token_embeddings = cpu + .text_encoder + .forward(&cpu_token_ids, Some(&cpu_attention_mask)); + let gpu_token_embeddings = gpu + .text_encoder + .forward(&gpu_token_ids, Some(&gpu_attention_mask)); + let cpu_text_states = cpu + .text_encoder + .debug_hidden_states(&cpu_token_ids, Some(&cpu_attention_mask)); + let gpu_text_states = gpu + .text_encoder + .debug_hidden_states(&gpu_token_ids, Some(&gpu_attention_mask)); + + for (idx, (cpu_state, gpu_state)) in cpu_text_states.iter().zip(&gpu_text_states).enumerate() + { + let name = if idx + 1 == cpu_text_states.len() { + "text_final_norm_output".to_string() + } else if idx == 0 { + "text_post_embedding_norm".to_string() + } else { + format!("text_post_layer_{}", idx - 1) + }; + let _ = print_diff(&name, cpu_state, gpu_state).await.unwrap(); + } + + let text_config = load_modern_bert_config_from_gguf(&model_bytes); + let cpu_layer2_input = cpu_text_states[2].clone(); + let gpu_layer2_input = gpu_text_states[2].clone(); + let cpu_layer2_attn_norm = load_layer_norm_from_gguf( + &model_bytes, + &["text", "blk.2", "attn_norm"], + &cpu_device, + text_config.norm_eps, + ); + let gpu_layer2_attn_norm = load_layer_norm_from_gguf( + &model_bytes, + &["text", "blk.2", "attn_norm"], + &gpu_device, + text_config.norm_eps, + ); + let cpu_layer2_attn_input = cpu_layer2_attn_norm.forward(&cpu_layer2_input); + let gpu_layer2_attn_input = gpu_layer2_attn_norm.forward(&gpu_layer2_input); + let _ = print_diff( + "text_layer2_attn_norm_output", + &cpu_layer2_attn_input, + &gpu_layer2_attn_input, + ) + .await + .unwrap(); + + let cpu_layer2_qkv = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.2"], + "attn_qkv.weight", + &cpu_device, + ); + let gpu_layer2_qkv = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.2"], + "attn_qkv.weight", + &gpu_device, + ); + let cpu_layer2_qkv_projection = cpu_layer2_attn_input.q_mat_mul(&cpu_layer2_qkv); + let gpu_layer2_qkv_projection = gpu_layer2_attn_input.q_mat_mul(&gpu_layer2_qkv); + let _ = print_diff( + "text_layer2_qkv_projection", + &cpu_layer2_qkv_projection, + &gpu_layer2_qkv_projection, + ) + .await + .unwrap(); + let rope_cache_cpu = fusor::RopeCache::new( + text_config.head_dimension, + text_config.context_length, + text_config.rope_theta, + &cpu_device, + ) + .unwrap(); + let rope_cache_gpu = fusor::RopeCache::new( + text_config.head_dimension, + text_config.context_length, + text_config.rope_theta, + &gpu_device, + ) + .unwrap(); + + let hidden_size = text_config.num_heads * text_config.head_dimension; + let [b_sz, seq_len, _] = cpu_layer0_qkv_projection.shape(); + let cpu_query_states = cpu_layer0_qkv_projection + .narrow(2, 0, hidden_size) + .reshape([b_sz, seq_len, text_config.num_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let cpu_key_states = cpu_layer0_qkv_projection + .narrow(2, hidden_size, hidden_size) + .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let cpu_value_states = cpu_layer0_qkv_projection + .narrow(2, 2 * hidden_size, hidden_size) + .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let gpu_query_states = gpu_layer0_qkv_projection + .narrow(2, 0, hidden_size) + .reshape([b_sz, seq_len, text_config.num_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let gpu_key_states = gpu_layer0_qkv_projection + .narrow(2, hidden_size, hidden_size) + .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let gpu_value_states = gpu_layer0_qkv_projection + .narrow(2, 2 * hidden_size, hidden_size) + .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + + let (cpu_query_after_rope, cpu_key_after_rope) = + rope_cache_cpu.forward(&cpu_query_states, &cpu_key_states, 0); + let (gpu_query_after_rope, gpu_key_after_rope) = + rope_cache_gpu.forward(&gpu_query_states, &gpu_key_states, 0); + let _ = print_diff( + "text_layer0_query_after_rope", + &cpu_query_after_rope, + &gpu_query_after_rope, + ) + .await + .unwrap(); + let _ = print_diff( + "text_layer0_key_after_rope", + &cpu_key_after_rope, + &gpu_key_after_rope, + ) + .await + .unwrap(); + + let cpu_attention_scores = + cpu_query_after_rope.mat_mul(&cpu_key_after_rope.transpose(2, 3)); + let gpu_attention_scores = + gpu_query_after_rope.mat_mul(&gpu_key_after_rope.transpose(2, 3)); + let _ = print_diff( + "text_layer0_attention_scores", + &cpu_attention_scores, + &gpu_attention_scores, + ) + .await + .unwrap(); + + let scale = 1.0 / (text_config.head_dimension as f32).sqrt(); + let cpu_attention_probs = cpu_attention_scores + .mul_scalar(scale) + .softmax_last_dim::<3>(); + let gpu_attention_probs = gpu_attention_scores + .mul_scalar(scale) + .softmax_last_dim::<3>(); + let _ = print_diff( + "text_layer0_attention_probs", + &cpu_attention_probs, + &gpu_attention_probs, + ) + .await + .unwrap(); + + let cpu_attention_context = cpu_attention_probs.mat_mul(&cpu_value_states); + let gpu_attention_context = gpu_attention_probs.mat_mul(&gpu_value_states); + let _ = print_diff( + "text_layer0_attention_context", + &cpu_attention_context, + &gpu_attention_context, + ) + .await + .unwrap(); + + let cpu_attn_output_weight = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.0"], + "attn_output.weight", + &cpu_device, + ); + let gpu_attn_output_weight = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.0"], + "attn_output.weight", + &gpu_device, + ); + let cpu_attention_context_flat = cpu_attention_context + .transpose(1, 2) + .to_concrete() + .reshape([b_sz, seq_len, hidden_size]) + .to_concrete(); + let gpu_attention_context_flat = gpu_attention_context + .transpose(1, 2) + .to_concrete() + .reshape([b_sz, seq_len, hidden_size]) + .to_concrete(); + let cpu_attention_output = cpu_attention_context_flat.q_mat_mul(&cpu_attn_output_weight); + let gpu_attention_output = gpu_attention_context_flat.q_mat_mul(&gpu_attn_output_weight); + let _ = print_diff( + "text_layer0_attention_output_projection", + &cpu_attention_output, + &gpu_attention_output, + ) + .await + .unwrap(); + + let cpu_after_attention = cpu_text_after_embd_norm.add_(&cpu_attention_output); + let gpu_after_attention = gpu_text_after_embd_norm.add_(&gpu_attention_output); + let _ = print_diff( + "text_layer0_after_attention_residual", + &cpu_after_attention, + &gpu_after_attention, + ) + .await + .unwrap(); + + let cpu_ffn_norm = load_layer_norm_from_gguf( + &model_bytes, + &["text", "blk.0", "ffn_norm"], + &cpu_device, + text_config.norm_eps, + ); + let gpu_ffn_norm = load_layer_norm_from_gguf( + &model_bytes, + &["text", "blk.0", "ffn_norm"], + &gpu_device, + text_config.norm_eps, + ); + let cpu_ffn_input = cpu_ffn_norm.forward(&cpu_after_attention); + let gpu_ffn_input = gpu_ffn_norm.forward(&gpu_after_attention); + let _ = print_diff("text_layer0_ffn_input", &cpu_ffn_input, &gpu_ffn_input) + .await + .unwrap(); + + let cpu_ffn_gate_up = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.0"], + "ffn_gate_up.weight", + &cpu_device, + ); + let gpu_ffn_gate_up = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.0"], + "ffn_gate_up.weight", + &gpu_device, + ); + let cpu_ffn_gate_up_proj = cpu_ffn_input.q_mat_mul(&cpu_ffn_gate_up).to_concrete(); + let gpu_ffn_gate_up_proj = gpu_ffn_input.q_mat_mul(&gpu_ffn_gate_up).to_concrete(); + let _ = print_diff( + "text_layer0_ffn_gate_up_projection", + &cpu_ffn_gate_up_proj, + &gpu_ffn_gate_up_proj, + ) + .await + .unwrap(); + + let intermediate_size = cpu_ffn_gate_up.shape()[0] / 2; + let cpu_gate = cpu_ffn_gate_up_proj + .narrow(2, 0, intermediate_size) + .to_concrete(); + let cpu_up = cpu_ffn_gate_up_proj.narrow(2, intermediate_size, intermediate_size); + let gpu_gate = gpu_ffn_gate_up_proj + .narrow(2, 0, intermediate_size) + .to_concrete(); + let gpu_up = gpu_ffn_gate_up_proj.narrow(2, intermediate_size, intermediate_size); + let cpu_ffn_activated = cpu_gate.gelu().mul_(&cpu_up); + let gpu_ffn_activated = gpu_gate.gelu().mul_(&gpu_up); + let _ = print_diff( + "text_layer0_ffn_activated", + &cpu_ffn_activated, + &gpu_ffn_activated, + ) + .await + .unwrap(); + + let cpu_ffn_down = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.0"], + "ffn_down.weight", + &cpu_device, + ); + let gpu_ffn_down = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.0"], + "ffn_down.weight", + &gpu_device, + ); + let cpu_ffn_output = cpu_ffn_activated.q_mat_mul(&cpu_ffn_down); + let gpu_ffn_output = gpu_ffn_activated.q_mat_mul(&gpu_ffn_down); + let _ = print_diff("text_layer0_ffn_output", &cpu_ffn_output, &gpu_ffn_output) + .await + .unwrap(); + + let cpu_layer0_output = cpu_after_attention.add_(&cpu_ffn_output); + let gpu_layer0_output = gpu_after_attention.add_(&gpu_ffn_output); + let _ = print_diff("text_layer0_output", &cpu_layer0_output, &gpu_layer0_output) + .await + .unwrap(); + + let cpu_layer1_attn_norm = load_layer_norm_from_gguf( + &model_bytes, + &["text", "blk.1", "attn_norm"], + &cpu_device, + text_config.norm_eps, + ); + let gpu_layer1_attn_norm = load_layer_norm_from_gguf( + &model_bytes, + &["text", "blk.1", "attn_norm"], + &gpu_device, + text_config.norm_eps, + ); + let cpu_layer1_attn_input = cpu_layer1_attn_norm.forward(&cpu_layer0_output); + let gpu_layer1_attn_input = gpu_layer1_attn_norm.forward(&gpu_layer0_output); + let _ = print_diff( + "text_layer1_attn_norm_output", + &cpu_layer1_attn_input, + &gpu_layer1_attn_input, + ) + .await + .unwrap(); + + let cpu_layer2_query_states = cpu_layer2_qkv_projection + .narrow(2, 0, hidden_size) + .reshape([b_sz, seq_len, text_config.num_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let cpu_layer2_key_states = cpu_layer2_qkv_projection + .narrow(2, hidden_size, hidden_size) + .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let cpu_layer2_value_states = cpu_layer2_qkv_projection + .narrow(2, 2 * hidden_size, hidden_size) + .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let gpu_layer2_query_states = gpu_layer2_qkv_projection + .narrow(2, 0, hidden_size) + .reshape([b_sz, seq_len, text_config.num_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let gpu_layer2_key_states = gpu_layer2_qkv_projection + .narrow(2, hidden_size, hidden_size) + .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + let gpu_layer2_value_states = gpu_layer2_qkv_projection + .narrow(2, 2 * hidden_size, hidden_size) + .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .transpose(1, 2) + .to_concrete(); + + let (cpu_layer2_query_after_rope, cpu_layer2_key_after_rope) = + rope_cache_cpu.forward(&cpu_layer2_query_states, &cpu_layer2_key_states, 0); + let (gpu_layer2_query_after_rope, gpu_layer2_key_after_rope) = + rope_cache_gpu.forward(&gpu_layer2_query_states, &gpu_layer2_key_states, 0); + let _ = print_diff( + "text_layer2_query_after_rope", + &cpu_layer2_query_after_rope, + &gpu_layer2_query_after_rope, + ) + .await + .unwrap(); + let _ = print_diff( + "text_layer2_key_after_rope", + &cpu_layer2_key_after_rope, + &gpu_layer2_key_after_rope, + ) + .await + .unwrap(); + + let cpu_layer2_attention_scores = + cpu_layer2_query_after_rope.mat_mul(&cpu_layer2_key_after_rope.transpose(2, 3)); + let gpu_layer2_attention_scores = + gpu_layer2_query_after_rope.mat_mul(&gpu_layer2_key_after_rope.transpose(2, 3)); + let _ = print_diff( + "text_layer2_attention_scores", + &cpu_layer2_attention_scores, + &gpu_layer2_attention_scores, + ) + .await + .unwrap(); + + let cpu_layer2_attention_probs = cpu_layer2_attention_scores + .mul_scalar(scale) + .softmax_last_dim::<3>(); + let gpu_layer2_attention_probs = gpu_layer2_attention_scores + .mul_scalar(scale) + .softmax_last_dim::<3>(); + let _ = print_diff( + "text_layer2_attention_probs", + &cpu_layer2_attention_probs, + &gpu_layer2_attention_probs, + ) + .await + .unwrap(); + + let cpu_layer2_attention_context = + cpu_layer2_attention_probs.mat_mul(&cpu_layer2_value_states); + let gpu_layer2_attention_context = + gpu_layer2_attention_probs.mat_mul(&gpu_layer2_value_states); + let _ = print_diff( + "text_layer2_attention_context", + &cpu_layer2_attention_context, + &gpu_layer2_attention_context, + ) + .await + .unwrap(); + + let cpu_layer2_attn_output_weight = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.2"], + "attn_output.weight", + &cpu_device, + ); + let gpu_layer2_attn_output_weight = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.2"], + "attn_output.weight", + &gpu_device, + ); + let cpu_layer2_attention_context_flat = cpu_layer2_attention_context + .transpose(1, 2) + .to_concrete() + .reshape([b_sz, seq_len, hidden_size]) + .to_concrete(); + let gpu_layer2_attention_context_flat = gpu_layer2_attention_context + .transpose(1, 2) + .to_concrete() + .reshape([b_sz, seq_len, hidden_size]) + .to_concrete(); + let cpu_layer2_attention_output = + cpu_layer2_attention_context_flat.q_mat_mul(&cpu_layer2_attn_output_weight); + let gpu_layer2_attention_output = + gpu_layer2_attention_context_flat.q_mat_mul(&gpu_layer2_attn_output_weight); + let _ = print_diff( + "text_layer2_attention_output_projection", + &cpu_layer2_attention_output, + &gpu_layer2_attention_output, + ) + .await + .unwrap(); + + let cpu_layer2_after_attention = cpu_layer2_input.add_(&cpu_layer2_attention_output); + let gpu_layer2_after_attention = gpu_layer2_input.add_(&gpu_layer2_attention_output); + let _ = print_diff( + "text_layer2_after_attention_residual", + &cpu_layer2_after_attention, + &gpu_layer2_after_attention, + ) + .await + .unwrap(); + + let cpu_layer2_ffn_norm = load_layer_norm_from_gguf( + &model_bytes, + &["text", "blk.2", "ffn_norm"], + &cpu_device, + text_config.norm_eps, + ); + let gpu_layer2_ffn_norm = load_layer_norm_from_gguf( + &model_bytes, + &["text", "blk.2", "ffn_norm"], + &gpu_device, + text_config.norm_eps, + ); + let cpu_layer2_ffn_input = cpu_layer2_ffn_norm.forward(&cpu_layer2_after_attention); + let gpu_layer2_ffn_input = gpu_layer2_ffn_norm.forward(&gpu_layer2_after_attention); + let _ = print_diff( + "text_layer2_ffn_input", + &cpu_layer2_ffn_input, + &gpu_layer2_ffn_input, + ) + .await + .unwrap(); + + let cpu_layer2_ffn_gate_up = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.2"], + "ffn_gate_up.weight", + &cpu_device, + ); + let gpu_layer2_ffn_gate_up = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.2"], + "ffn_gate_up.weight", + &gpu_device, + ); + let cpu_layer2_ffn_gate_up_proj = + cpu_layer2_ffn_input.q_mat_mul(&cpu_layer2_ffn_gate_up).to_concrete(); + let gpu_layer2_ffn_gate_up_proj = + gpu_layer2_ffn_input.q_mat_mul(&gpu_layer2_ffn_gate_up).to_concrete(); + let _ = print_diff( + "text_layer2_ffn_gate_up_projection", + &cpu_layer2_ffn_gate_up_proj, + &gpu_layer2_ffn_gate_up_proj, + ) + .await + .unwrap(); + + let layer2_intermediate_size = cpu_layer2_ffn_gate_up.shape()[0] / 2; + let cpu_layer2_gate = cpu_layer2_ffn_gate_up_proj + .narrow(2, 0, layer2_intermediate_size) + .to_concrete(); + let cpu_layer2_up = cpu_layer2_ffn_gate_up_proj + .narrow(2, layer2_intermediate_size, layer2_intermediate_size); + let gpu_layer2_gate = gpu_layer2_ffn_gate_up_proj + .narrow(2, 0, layer2_intermediate_size) + .to_concrete(); + let gpu_layer2_up = gpu_layer2_ffn_gate_up_proj + .narrow(2, layer2_intermediate_size, layer2_intermediate_size); + let cpu_layer2_ffn_activated = cpu_layer2_gate.gelu().mul_(&cpu_layer2_up); + let gpu_layer2_ffn_activated = gpu_layer2_gate.gelu().mul_(&gpu_layer2_up); + let _ = print_diff( + "text_layer2_ffn_activated", + &cpu_layer2_ffn_activated, + &gpu_layer2_ffn_activated, + ) + .await + .unwrap(); + + let cpu_layer2_ffn_down = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.2"], + "ffn_down.weight", + &cpu_device, + ); + let gpu_layer2_ffn_down = load_qmatrix_tensor_from_gguf( + &model_bytes, + &["text", "blk.2"], + "ffn_down.weight", + &gpu_device, + ); + let cpu_layer2_ffn_output = cpu_layer2_ffn_activated.q_mat_mul(&cpu_layer2_ffn_down); + let gpu_layer2_ffn_output = gpu_layer2_ffn_activated.q_mat_mul(&gpu_layer2_ffn_down); + let _ = print_diff( + "text_layer2_ffn_output", + &cpu_layer2_ffn_output, + &gpu_layer2_ffn_output, + ) + .await + .unwrap(); + + let cpu_layer2_output = cpu_layer2_after_attention.add_(&cpu_layer2_ffn_output); + let gpu_layer2_output = gpu_layer2_after_attention.add_(&gpu_layer2_ffn_output); + let _ = print_diff("text_layer2_output", &cpu_layer2_output, &gpu_layer2_output) + .await + .unwrap(); + + let _ = print_diff("token_embeddings", &cpu_token_embeddings, &gpu_token_embeddings) + .await + .unwrap(); + + let (cpu_word_embeddings, _) = + first_subtoken_pooling(&cpu_token_embeddings, &[cpu_tokenized.clone()], &cpu_device); + let (gpu_word_embeddings, _) = + first_subtoken_pooling(&gpu_token_embeddings, &[gpu_tokenized.clone()], &gpu_device); + let _ = print_diff("word_embeddings", &cpu_word_embeddings, &gpu_word_embeddings) + .await + .unwrap(); + + let (cpu_span_embeddings, cpu_span_indices) = + cpu.span_layer.forward(&cpu_word_embeddings, &cpu_device); + let (gpu_span_embeddings, gpu_span_indices) = + gpu.span_layer.forward(&gpu_word_embeddings, &gpu_device); + assert_eq!(cpu_span_indices, gpu_span_indices); + let _ = print_diff("span_embeddings", &cpu_span_embeddings, &gpu_span_embeddings) + .await + .unwrap(); + + let cpu_scores = Scorer::forward(&cpu_span_embeddings, &cpu_label_embeddings); + let gpu_scores = Scorer::forward(&gpu_span_embeddings, &gpu_label_embeddings); + let max_score_diff = print_diff("span_scores", &cpu_scores, &gpu_scores) + .await + .unwrap(); + + let cpu_entities = cpu.extract(text, &labels).await.unwrap(); + let gpu_entities = gpu.extract(text, &labels).await.unwrap(); + println!("cpu_entities={cpu_entities:?}"); + println!("gpu_entities={gpu_entities:?}"); + + let cpu_entities: Vec<_> = cpu_entities + .iter() + .map(|entity| (entity.label.as_str(), entity.text.as_str())) + .collect(); + let gpu_entities: Vec<_> = gpu_entities + .iter() + .map(|entity| (entity.label.as_str(), entity.text.as_str())) + .collect(); + + assert!( + max_score_diff < 0.05, + "CPU/GPU score drift is too large: max_abs_diff={max_score_diff:.6}" + ); + assert_eq!(gpu_entities, cpu_entities); + } +} diff --git a/models/rgliner/src/raw/label_encoder.rs b/models/rgliner/src/raw/label_encoder.rs new file mode 100644 index 000000000..3f96eaabe --- /dev/null +++ b/models/rgliner/src/raw/label_encoder.rs @@ -0,0 +1,308 @@ +//! Label encoder using sentence transformers. + +use fusor::{Device, Result, Tensor, VarBuilder}; +use kalosm_language_model::Embedding; +use rbert::{Bert, BertSource, Pooling}; +use std::sync::Arc; + +use crate::error::GlinerError; + +/// Projection FFN for aligning label embeddings to text encoder dimension. +/// +/// Architecture: Linear(hidden, hidden*4) -> ReLU -> Linear(hidden*4, hidden) +/// This matches the Python create_projection_layer() function in GLiNER. +pub struct ProjectionFFN { + weight1: Tensor<2, f32>, + bias1: Tensor<1, f32>, + weight2: Tensor<2, f32>, + bias2: Tensor<1, f32>, +} + +impl ProjectionFFN { + /// Load projection FFN from weights with proper transposition. + pub fn load(device: &Device, vb: &mut VarBuilder<'_>) -> Result { + // Try different naming conventions + let (weight1, bias1) = + Self::load_layer(device, vb, &["label_fnn.0", "label_ffn.0", "label_proj.0"])?; + let (weight2, bias2) = + Self::load_layer(device, vb, &["label_fnn.2", "label_ffn.2", "label_proj.2"])?; + + Ok(Self { + weight1, + bias1, + weight2, + bias2, + }) + } + + fn load_layer( + device: &Device, + vb: &mut VarBuilder, + prefixes: &[&str], + ) -> Result<(Tensor<2, f32>, Tensor<1, f32>)> { + for prefix in prefixes { + let mut layer_vb = vb.pp(prefix); + if let Ok(weight_q) = layer_vb.get("weight", device) { + let weight: Tensor<2, f32> = weight_q.dequantize(); + + // PyTorch nn.Linear stores weights as [out_features, in_features] + // GGUF stores the same way. Fusor loads as-is. + // For x @ W where x is [B, in], we need W to be [in, out] + // So we transpose [out, in] -> [in, out] + let weight_t = weight.t().to_concrete(); + + if let Ok(bias_q) = layer_vb.get("bias", device) { + let bias: Tensor<1, f32> = bias_q.dequantize(); + return Ok((weight_t, bias)); + } + } + } + Err(fusor::Error::msg(format!( + "Could not load projection layer with prefixes {:?}", + prefixes + ))) + } + + /// Get output dimension. + pub fn out_features(&self) -> usize { + // After transpose, weight2 is [in=1536, out=384], so output dim is shape[1] + self.weight2.shape()[1] + } + + /// Forward pass through projection. + /// Computes: ReLU(x @ W1 + b1) @ W2 + b2 + pub fn forward(&self, x: &Tensor<2, f32>) -> Tensor<2, f32> { + // Layer 1: x @ W1 + b1 + // x is [num_labels, 384], W1 is [384, 1536] after transpose + // So x @ W1 = [num_labels, 384] @ [384, 1536] = [num_labels, 1536] + let h1 = x.mat_mul(&self.weight1); + + let [num_labels, hidden_dim] = h1.shape(); + let bias1_broadcast: Tensor<2, f32> = self + .bias1 + .unsqueeze(0) + .to_concrete() + .broadcast_as([num_labels, hidden_dim]) + .to_concrete(); + let h1_biased = (h1 + bias1_broadcast).to_concrete(); + + // ReLU activation (Python GLiNER uses ReLU, not GELU) + let h1_relu = h1_biased.relu(); + + // Layer 2: h @ W2 + b2 + // h is [num_labels, 1536], W2 is [1536, 384] after transpose + // So h @ W2 = [num_labels, 1536] @ [1536, 384] = [num_labels, 384] + let out = h1_relu.mat_mul(&self.weight2); + let [num_labels2, out_dim] = out.shape(); + let bias2_broadcast: Tensor<2, f32> = self + .bias2 + .unsqueeze(0) + .to_concrete() + .broadcast_as([num_labels2, out_dim]) + .to_concrete(); + (out + bias2_broadcast).to_concrete() + } +} + +/// Label encoder: sentence transformer + projection FFN. +pub struct LabelEncoder { + /// Sentence transformer model (reuses rbert). + sentence_encoder: Arc, + /// Projection FFN to align dimensions. + projection: ProjectionFFN, + /// Output dimension. + output_dim: usize, + /// Device for creating tensors. + device: Device, +} + +impl LabelEncoder { + /// Load label encoder from separate GGUF file. + pub async fn load( + device: &Device, + projection_vb: &mut VarBuilder<'_>, + sentence_encoder_source: BertSource, + ) -> std::result::Result { + // Load sentence encoder from separate model + let sentence_encoder = Bert::builder() + .with_source(sentence_encoder_source) + .with_device(device.clone()) + .build() + .await?; + + let projection = ProjectionFFN::load(device, projection_vb)?; + let output_dim = projection.out_features(); + + Ok(Self { + sentence_encoder: Arc::new(sentence_encoder), + projection, + output_dim, + device: device.clone(), + }) + } + + /// Get the output dimension. + pub fn output_dim(&self) -> usize { + self.output_dim + } + + #[cfg(test)] + pub async fn debug_sentence_embeddings( + &self, + labels: &[&str], + ) -> std::result::Result, GlinerError> { + let embeddings = self + .sentence_encoder + .embed_batch_with_pooling_and_normalization(labels.to_vec(), Pooling::Mean, false) + .await?; + Ok(self.embeddings_to_tensor(&embeddings)) + } + + #[cfg(test)] + pub fn debug_projection(&self, x: &Tensor<2, f32>) -> Tensor<2, f32> { + self.projection.forward(x) + } + + #[cfg(test)] + pub fn debug_sentence_token_embeddings_and_mask( + &self, + labels: &[&str], + ) -> std::result::Result<(Tensor<3, f32>, Tensor<2, u32>), GlinerError> { + self.sentence_encoder + .debug_batch_forward(labels.iter().map(|s| (*s).to_string()).collect()) + .map_err(Into::into) + } + + #[cfg(test)] + pub fn debug_sentence_mean_pool( + &self, + labels: &[&str], + ) -> std::result::Result, GlinerError> { + self.sentence_encoder + .debug_batch_mean_pool(labels.iter().map(|s| (*s).to_string()).collect(), false) + .map_err(Into::into) + } + + #[cfg(test)] + pub fn debug_sentence_hidden_states( + &self, + labels: &[&str], + ) -> std::result::Result<(Vec>, Tensor<2, u32>), GlinerError> { + self.sentence_encoder + .debug_batch_hidden_states(labels.iter().map(|s| (*s).to_string()).collect()) + .map_err(Into::into) + } + + #[cfg(test)] + pub fn debug_sentence_first_layer( + &self, + labels: &[&str], + ) -> std::result::Result<(Tensor<3, f32>, Tensor<3, f32>, Tensor<3, f32>), GlinerError> { + self.sentence_encoder + .debug_batch_first_layer(labels.iter().map(|s| (*s).to_string()).collect()) + .map_err(Into::into) + } + + #[cfg(test)] + pub fn debug_sentence_first_layer_attention( + &self, + labels: &[&str], + ) -> std::result::Result< + ( + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<3, f32>, + Tensor<3, f32>, + ), + GlinerError, + > { + self.sentence_encoder + .debug_batch_first_layer_attention(labels.iter().map(|s| (*s).to_string()).collect()) + .map_err(Into::into) + } + + /// Encode labels to embeddings. + /// + /// # Arguments + /// * `labels` - Label strings to encode + /// + /// # Returns + /// Label embeddings [num_labels, output_dim] + pub async fn encode_labels( + &self, + labels: &[&str], + ) -> std::result::Result, GlinerError> { + if labels.is_empty() { + return Ok(Tensor::zeros(&self.device, [0, self.output_dim])); + } + + // Python GLiNER mean-pools label tokens without the L2 normalization that + // rbert applies in its default embedding API. + let embeddings = self + .sentence_encoder + .embed_batch_with_pooling_and_normalization(labels.to_vec(), Pooling::Mean, false) + .await?; + + // Convert Embeddings to tensor + let label_tensor = self.embeddings_to_tensor(&embeddings); + + // Project to text encoder dimension using the label_fnn + let projected = self.projection.forward(&label_tensor); + + // Return projected embeddings without normalization + // The model was trained end-to-end with this projection + Ok(projected) + } + + /// Convert Vec to Tensor<2, f32> + fn embeddings_to_tensor(&self, embeddings: &[Embedding]) -> Tensor<2, f32> { + if embeddings.is_empty() { + return Tensor::zeros(&self.device, [0, self.output_dim]); + } + + let num_labels = embeddings.len(); + let embed_dim = embeddings[0].vector().len(); + + // Flatten all embeddings into a single Vec + let mut data: Vec = Vec::with_capacity(num_labels * embed_dim); + for emb in embeddings { + data.extend_from_slice(emb.vector()); + } + + // Create tensor from flat data + Tensor::new(&self.device, &data) + .reshape([num_labels, embed_dim]) + .to_concrete() + } +} + +/// Cached label embeddings for efficient repeated inference. +pub struct CachedLabels { + /// Original label strings. + pub labels: Vec, + /// Precomputed label embeddings [num_labels, hidden_dim]. + pub embeddings: Tensor<2, f32>, +} + +impl CachedLabels { + /// Create cached labels from precomputed embeddings. + pub fn new(labels: Vec, embeddings: Tensor<2, f32>) -> Self { + Self { labels, embeddings } + } + + /// Get the number of labels. + pub fn len(&self) -> usize { + self.labels.len() + } + + /// Check if empty. + pub fn is_empty(&self) -> bool { + self.labels.is_empty() + } + + /// Get label at index. + pub fn get_label(&self, idx: usize) -> Option<&str> { + self.labels.get(idx).map(|s| s.as_str()) + } +} diff --git a/models/rgliner/src/raw/mod.rs b/models/rgliner/src/raw/mod.rs new file mode 100644 index 000000000..abd4617e2 --- /dev/null +++ b/models/rgliner/src/raw/mod.rs @@ -0,0 +1,13 @@ +//! Raw model implementations for GLiNER. + +pub mod modern_bert; + +mod label_encoder; +mod scorer; +mod span_layer; +mod text_encoder; + +pub use label_encoder::{CachedLabels, LabelEncoder}; +pub use scorer::Scorer; +pub use span_layer::SpanLayer; +pub use text_encoder::TextEncoder; diff --git a/models/rgliner/src/raw/modern_bert/attention.rs b/models/rgliner/src/raw/modern_bert/attention.rs new file mode 100644 index 000000000..b2c6b2e2e --- /dev/null +++ b/models/rgliner/src/raw/modern_bert/attention.rs @@ -0,0 +1,99 @@ +//! ModernBERT self-attention with RoPE and fused QKV. + +use fusor::{Device, QMatrix, Result, RopeCache, Tensor, VarBuilder}; + +/// ModernBERT self-attention with fused QKV projection and RoPE. +pub struct ModernBertAttention { + /// Fused QKV projection: [3 * hidden_size, hidden_size] + wqkv: QMatrix, + wo: QMatrix, + num_heads: usize, + num_kv_heads: usize, + head_dim: usize, +} + +impl ModernBertAttention { + pub fn load( + device: &Device, + vb: &mut VarBuilder, + num_heads: usize, + num_kv_heads: usize, + head_dim: usize, + _eps: f32, + ) -> Result { + // Fused QKV weight + let wqkv = vb.get("attn_qkv.weight", device)?; + let wo = vb.get("attn_output.weight", device)?; + + Ok(Self { + wqkv, + wo, + num_heads, + num_kv_heads, + head_dim, + }) + } + + pub fn forward( + &self, + hidden_states: &Tensor<3, f32>, + rope_cache: &RopeCache, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + let [b_sz, seq_len, _hidden_size] = hidden_states.shape(); + let hidden_size = self.num_heads * self.head_dim; + + // Compute fused QKV projection: [batch, seq_len, 3 * hidden_size] + let qkv = hidden_states.q_mat_mul(&self.wqkv).to_concrete(); + + // Split into Q, K, V - each [batch, seq_len, hidden_size] + let query_states = qkv + .narrow(2, 0, hidden_size) + .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + + let key_states = qkv + .narrow(2, hidden_size, hidden_size) + .reshape([b_sz, seq_len, self.num_kv_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + + let value_states = qkv + .narrow(2, 2 * hidden_size, hidden_size) + .reshape([b_sz, seq_len, self.num_kv_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + + // Apply RoPE to Q and K + let (query_states, key_states) = rope_cache.forward(&query_states, &key_states, 0); + + // Scaled dot-product attention + let scale = 1.0 / (self.head_dim as f32).sqrt(); + + // Convert attention mask for flash attention if provided + const MASK_NEG_VALUE: f32 = -10000.0; + let mask: Option> = attention_mask.map(|m| { + let mask_f32: Tensor<2, f32> = m.cast(); + let zeros = mask_f32.zeros_like(); + let ones = (zeros + 1.0f32).to_concrete(); + ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete() + }); + + let attn_output = query_states.flash_attention( + &key_states, + &value_states, + scale, + mask.as_ref().map(|m| (m, fusor::MaskKind::BatchKeyMask)), + ); + + // Reshape and project output + let attn_output = attn_output.transpose(1, 2); + let attn_output = attn_output + .to_concrete() + .reshape([b_sz, seq_len, hidden_size]) + .to_concrete(); + + attn_output.q_mat_mul(&self.wo) + } +} diff --git a/models/rgliner/src/raw/modern_bert/config.rs b/models/rgliner/src/raw/modern_bert/config.rs new file mode 100644 index 000000000..5c6960bf5 --- /dev/null +++ b/models/rgliner/src/raw/modern_bert/config.rs @@ -0,0 +1,101 @@ +//! ModernBERT configuration from GGUF metadata. + +use fusor::{Result, VarBuilder}; + +/// Configuration for ModernBERT loaded from GGUF metadata. +#[derive(Debug, Clone)] +pub struct ModernBertConfig { + /// Number of attention heads. + pub num_heads: usize, + /// Number of key-value heads (for GQA). + pub num_kv_heads: usize, + /// Number of transformer layers. + pub num_layers: usize, + /// Hidden size (embedding dimension). + pub hidden_size: usize, + /// Dimension per attention head. + pub head_dimension: usize, + /// Intermediate size for FFN. + pub intermediate_size: usize, + /// Maximum context length. + pub context_length: usize, + /// RoPE base frequency. + pub rope_theta: f32, + /// LayerNorm epsilon. + pub norm_eps: f32, +} + +impl ModernBertConfig { + /// Load configuration from GGUF metadata. + pub fn from_gguf(vb: &VarBuilder) -> Result { + let num_heads = vb + .get_metadata(".attention.head_count") + .and_then(|v| v.to_u32().ok()) + .ok_or_else(|| { + fusor::Error::msg("Missing required GGUF metadata: .attention.head_count") + })? as usize; + + let num_kv_heads = vb + .get_metadata(".attention.head_count_kv") + .and_then(|v| v.to_u32().ok()) + .unwrap_or(num_heads as u32) as usize; + + let num_layers = vb + .get_metadata(".block_count") + .and_then(|v| v.to_u32().ok()) + .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: .block_count"))? + as usize; + + let hidden_size = vb + .get_metadata(".embedding_length") + .and_then(|v| v.to_u32().ok()) + .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: .embedding_length"))? + as usize; + + if hidden_size % num_heads != 0 { + return Err(fusor::Error::msg(format!( + "hidden_size ({hidden_size}) must be divisible by num_heads ({num_heads})" + ))); + } + + let intermediate_size = vb + .get_metadata(".feed_forward_length") + .and_then(|v| v.to_u32().ok()) + .unwrap_or((hidden_size * 4) as u32) as usize; + + let context_length = vb + .get_metadata(".context_length") + .and_then(|v| v.to_u32().ok()) + .unwrap_or(8192) as usize; + + let rope_theta = vb + .get_metadata(".rope.freq_base") + .and_then(|v| v.to_f32().ok()) + .unwrap_or(10000.0); + + let norm_eps = vb + .get_metadata(".attention.layer_norm_rms_epsilon") + .and_then(|v| v.to_f32().ok()) + .unwrap_or(1e-6); + + // Use attention.key_length for head dimension + // Fall back to hidden_size / num_heads if not present + let head_dimension = vb + .get_metadata(".attention.key_length") + .and_then(|v| v.to_u32().ok()) + .map(|x| x as usize) + .unwrap_or_else(|| hidden_size / num_heads); + + Ok(Self { + num_heads, + num_kv_heads, + num_layers, + hidden_size, + head_dimension, + intermediate_size, + context_length, + rope_theta, + norm_eps, + }) + } +} diff --git a/models/rgliner/src/raw/modern_bert/feed_forward.rs b/models/rgliner/src/raw/modern_bert/feed_forward.rs new file mode 100644 index 000000000..e9cea0d1d --- /dev/null +++ b/models/rgliner/src/raw/modern_bert/feed_forward.rs @@ -0,0 +1,45 @@ +//! ModernBERT GeGLU Feed Forward Network. + +use fusor::{Device, QMatrix, Result, Tensor, VarBuilder}; + +/// GeGLU Feed Forward Network with fused gate+up projection. +/// +/// Formula: GeGLU(x) = GELU(gate) * up @ down +/// where [gate, up] = x @ fused_gate_up +/// +/// This differs from Qwen's SiLU-gated FFN by using GELU instead of SiLU. +pub struct GeGluFeedForward { + /// Fused gate+up projection: [2 * intermediate_size, hidden_size] + gate_up: QMatrix, + down: QMatrix, + intermediate_size: usize, +} + +impl GeGluFeedForward { + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let gate_up = vb.get("ffn_gate_up.weight", device)?; + let down = vb.get("ffn_down.weight", device)?; + + // Determine intermediate size from fused weight dimensions + // gate_up is [2 * intermediate_size, hidden_size] + let intermediate_size = gate_up.shape()[0] / 2; + + Ok(Self { + gate_up, + down, + intermediate_size, + }) + } + + pub fn forward(&self, x: &Tensor<3, f32>) -> Tensor<3, f32> { + // Compute fused gate+up: [batch, seq_len, 2 * intermediate_size] + let gate_up = x.q_mat_mul(&self.gate_up).to_concrete(); + + // Split into gate and up + let gate = gate_up.narrow(2, 0, self.intermediate_size).to_concrete(); + let up = gate_up.narrow(2, self.intermediate_size, self.intermediate_size); + + // GeGLU: GELU(gate) * up, then project down + gate.gelu().mul_(&up).q_mat_mul(&self.down) + } +} diff --git a/models/rgliner/src/raw/modern_bert/layer.rs b/models/rgliner/src/raw/modern_bert/layer.rs new file mode 100644 index 000000000..280edda1a --- /dev/null +++ b/models/rgliner/src/raw/modern_bert/layer.rs @@ -0,0 +1,77 @@ +//! ModernBERT transformer layer with pre-norm architecture. + +use fusor::layers::LayerNorm; +use fusor::{Device, Result, RopeCache, Tensor, VarBuilder}; + +use super::attention::ModernBertAttention; +use super::feed_forward::GeGluFeedForward; + +/// A single ModernBERT transformer layer with pre-norm architecture. +/// +/// Note: The first layer (index 0) doesn't have its own attention_norm because +/// the embedding norm serves that purpose. This is handled by making attention_norm +/// optional and passing in the pre-normalized input for layer 0. +pub struct ModernBertLayer { + /// Pre-attention RMSNorm (None for layer 0, which uses embedding norm) + attention_norm: Option>, + attention: ModernBertAttention, + ffn_norm: LayerNorm<1, f32>, + feed_forward: GeGluFeedForward, +} + +impl ModernBertLayer { + pub fn load( + device: &Device, + vb: &mut VarBuilder, + num_heads: usize, + num_kv_heads: usize, + head_dim: usize, + eps: f32, + layer_idx: usize, + ) -> Result { + // Layer 0 doesn't have attn_norm - it uses embedding norm instead + let attention_norm = if layer_idx > 0 { + Some(LayerNorm::load(device, &mut vb.pp("attn_norm"), eps)?) + } else { + None + }; + + let attention = + ModernBertAttention::load(device, vb, num_heads, num_kv_heads, head_dim, eps)?; + let ffn_norm = LayerNorm::load(device, &mut vb.pp("ffn_norm"), eps)?; + let feed_forward = GeGluFeedForward::load(device, vb)?; + + Ok(Self { + attention_norm, + attention, + ffn_norm, + feed_forward, + }) + } + + pub fn forward( + &self, + hidden_states: &Tensor<3, f32>, + rope_cache: &RopeCache, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + // Pre-norm + attention + residual + // For layer 0, hidden_states is already normalized by embedding norm + let residual = hidden_states; + let hidden_states = if let Some(ref norm) = self.attention_norm { + norm.forward(hidden_states) + } else { + hidden_states.clone() + }; + let hidden_states = self + .attention + .forward(&hidden_states, rope_cache, attention_mask); + let hidden_states = residual.add_(&hidden_states); + + // Pre-norm + FFN + residual + let residual = &hidden_states; + let ffn_input = self.ffn_norm.forward(&hidden_states); + let ffn_output = self.feed_forward.forward(&ffn_input); + residual.add_(&ffn_output) + } +} diff --git a/models/rgliner/src/raw/modern_bert/mod.rs b/models/rgliner/src/raw/modern_bert/mod.rs new file mode 100644 index 000000000..fe408bdca --- /dev/null +++ b/models/rgliner/src/raw/modern_bert/mod.rs @@ -0,0 +1,16 @@ +//! ModernBERT/Ettin encoder implementation. +//! +//! ModernBERT uses: +//! - RoPE (Rotary Position Embeddings) +//! - Pre-normalization with RMSNorm +//! - GeGLU activation in FFN +//! - No token type IDs + +mod attention; +mod config; +mod feed_forward; +mod layer; +mod model; + +pub use config::ModernBertConfig; +pub use model::ModernBertModel; diff --git a/models/rgliner/src/raw/modern_bert/model.rs b/models/rgliner/src/raw/modern_bert/model.rs new file mode 100644 index 000000000..5d1e88a71 --- /dev/null +++ b/models/rgliner/src/raw/modern_bert/model.rs @@ -0,0 +1,127 @@ +//! ModernBERT encoder model. + +use fusor::layers::{Embedding, LayerNorm}; +use fusor::{Device, Result, RopeCache, Tensor, VarBuilder}; + +use super::config::ModernBertConfig; +use super::layer::ModernBertLayer; + +/// ModernBERT encoder model (text encoder for GLiNER). +pub struct ModernBertModel { + token_embeddings: Embedding, + /// Embedding norm applied after token embeddings, before first layer + embedding_norm: LayerNorm<1, f32>, + layers: Vec, + final_norm: LayerNorm<1, f32>, + rope_cache: RopeCache, + pub(crate) device: Device, + config: ModernBertConfig, +} + +impl ModernBertModel { + /// Load ModernBERT from GGUF weights. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let config = ModernBertConfig::from_gguf(vb)?; + + // Load token embeddings + let token_embeddings = Embedding::load(device, &mut vb.pp("token_embd"))?; + + // Load embedding norm (applied before first layer) + let embedding_norm = LayerNorm::load(device, &mut vb.pp("embd_norm"), config.norm_eps)?; + + // Create RoPE cache + let rope_cache = RopeCache::new( + config.head_dimension, + config.context_length, + config.rope_theta, + device, + )?; + + // Load transformer layers + let mut layers = Vec::with_capacity(config.num_layers); + for i in 0..config.num_layers { + let layer = ModernBertLayer::load( + device, + &mut vb.pp(format!("blk.{i}")), + config.num_heads, + config.num_kv_heads, + config.head_dimension, + config.norm_eps, + i, + )?; + layers.push(layer); + } + + // Load final layer norm + let final_norm = LayerNorm::load(device, &mut vb.pp("output_norm"), config.norm_eps)?; + + Ok(Self { + token_embeddings, + embedding_norm, + layers, + final_norm, + rope_cache, + device: device.clone(), + config, + }) + } + + /// Forward pass through the model. + /// + /// Returns: [batch_size, seq_len, hidden_size] + pub fn forward( + &self, + input_ids: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + // Get token embeddings + let hidden_states = self.token_embeddings.forward(input_ids); + + // Apply embedding norm (serves as pre-norm for layer 0) + let mut hidden_states = self.embedding_norm.forward(&hidden_states); + + // Pass through transformer layers + for layer in &self.layers { + hidden_states = layer.forward(&hidden_states, &self.rope_cache, attention_mask); + } + + // Apply final layer norm + self.final_norm.forward(&hidden_states) + } + + /// Get the maximum sequence length. + pub fn max_seq_len(&self) -> usize { + self.config.context_length + } + + /// Get the embedding dimension. + pub fn embedding_dim(&self) -> usize { + self.config.hidden_size + } + + /// Get the device. + pub fn device(&self) -> &Device { + &self.device + } + + #[cfg(test)] + pub fn debug_hidden_states( + &self, + input_ids: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Vec> { + let mut states = Vec::with_capacity(self.layers.len() + 2); + + let hidden_states = self.token_embeddings.forward(input_ids); + let mut hidden_states = self.embedding_norm.forward(&hidden_states); + states.push(hidden_states.clone()); + + for layer in &self.layers { + hidden_states = layer.forward(&hidden_states, &self.rope_cache, attention_mask); + states.push(hidden_states.clone()); + } + + states.push(self.final_norm.forward(&hidden_states)); + states + } +} diff --git a/models/rgliner/src/raw/scorer.rs b/models/rgliner/src/raw/scorer.rs new file mode 100644 index 000000000..966649c9c --- /dev/null +++ b/models/rgliner/src/raw/scorer.rs @@ -0,0 +1,62 @@ +//! Scoring layer for span-label matching. + +use fusor::Tensor; + +/// Scorer for computing span-label similarity. +pub struct Scorer; + +impl Scorer { + /// Compute entity scores using dot product similarity. + /// + /// GLiNER uses raw dot product (not cosine similarity). + /// The output logits are passed through sigmoid externally. + /// + /// # Arguments + /// * `span_embeddings` - Span embeddings [batch, num_spans, hidden_dim] + /// * `label_embeddings` - Label embeddings [num_labels, hidden_dim] + /// + /// # Returns + /// Raw dot product scores [batch, num_spans, num_labels] + pub fn forward( + span_embeddings: &Tensor<3, f32>, + label_embeddings: &Tensor<2, f32>, + ) -> Tensor<3, f32> { + let [batch_size, num_spans, hidden_dim] = span_embeddings.shape(); + let [num_labels, _] = label_embeddings.shape(); + + // Flatten batch dimension for matmul + let span_concrete = span_embeddings.to_concrete(); + let flat_spans = span_concrete + .reshape([batch_size * num_spans, hidden_dim]) + .to_concrete(); + + // Transpose labels: [hidden_dim, num_labels] + let labels_t = label_embeddings.t(); + + // Matmul: [batch * num_spans, hidden_dim] @ [hidden_dim, num_labels] + // = [batch * num_spans, num_labels] + let flat_logits = flat_spans.mat_mul(&labels_t); + + // Reshape to [batch, num_spans, num_labels] + flat_logits + .reshape([batch_size, num_spans, num_labels]) + .to_concrete() + } +} + +/// Apply sigmoid to raw scores in Rust (not on GPU). +/// +/// # Arguments +/// * `logits` - Raw logit scores +/// +/// # Returns +/// Probability scores (0.0 to 1.0) +#[inline] +pub fn sigmoid(x: f32) -> f32 { + 1.0 / (1.0 + (-x).exp()) +} + +/// Apply sigmoid to a slice of logits. +pub fn apply_sigmoid(logits: &[f32]) -> Vec { + logits.iter().map(|&x| sigmoid(x)).collect() +} diff --git a/models/rgliner/src/raw/span_layer.rs b/models/rgliner/src/raw/span_layer.rs new file mode 100644 index 000000000..93418904c --- /dev/null +++ b/models/rgliner/src/raw/span_layer.rs @@ -0,0 +1,221 @@ +//! Span representation layer. +//! +//! The actual GLiNER architecture uses: +//! - project_start: 2-layer FFN for start word +//! - project_end: 2-layer FFN for end word +//! - out_project: 2-layer FFN for combined (start + end) representation + +use fusor::layers::Linear; +use fusor::{Device, Result, Tensor, VarBuilder}; + +/// Span representation layer. +/// +/// Creates span embeddings by projecting start and end word embeddings +/// separately, then combining them through an output projection. +pub struct SpanLayer { + /// Project start word: [hidden_dim] -> [hidden_dim] + start_fc1: Linear, + start_fc2: Linear, + /// Project end word: [hidden_dim] -> [hidden_dim] + end_fc1: Linear, + end_fc2: Linear, + /// Output projection: [2 * hidden_dim] -> [hidden_dim] + out_fc1: Linear, + out_fc2: Linear, + /// Maximum span width + max_width: usize, + /// Hidden dimension + hidden_dim: usize, +} + +impl SpanLayer { + /// Load span layer from GGUF weights. + pub fn load(device: &Device, vb: &mut VarBuilder, max_width: usize) -> Result { + // Try different weight naming conventions + let start_fc1 = Linear::load(device, &mut vb.pp("span.start_fc1")).or_else(|_| { + Linear::load( + device, + &mut vb.pp("span_rep_layer.span_rep_layer.project_start.0"), + ) + })?; + let start_fc2 = Linear::load(device, &mut vb.pp("span.start_fc2")).or_else(|_| { + Linear::load( + device, + &mut vb.pp("span_rep_layer.span_rep_layer.project_start.3"), + ) + })?; + + let end_fc1 = Linear::load(device, &mut vb.pp("span.end_fc1")).or_else(|_| { + Linear::load( + device, + &mut vb.pp("span_rep_layer.span_rep_layer.project_end.0"), + ) + })?; + let end_fc2 = Linear::load(device, &mut vb.pp("span.end_fc2")).or_else(|_| { + Linear::load( + device, + &mut vb.pp("span_rep_layer.span_rep_layer.project_end.3"), + ) + })?; + + let out_fc1 = Linear::load(device, &mut vb.pp("span.out_fc1")).or_else(|_| { + Linear::load( + device, + &mut vb.pp("span_rep_layer.span_rep_layer.out_project.0"), + ) + })?; + let out_fc2 = Linear::load(device, &mut vb.pp("span.out_fc2")).or_else(|_| { + Linear::load( + device, + &mut vb.pp("span_rep_layer.span_rep_layer.out_project.3"), + ) + })?; + + let hidden_dim = out_fc2.out_features(); + + Ok(Self { + start_fc1, + start_fc2, + end_fc1, + end_fc2, + out_fc1, + out_fc2, + max_width, + hidden_dim, + }) + } + + /// Get the maximum span width. + pub fn max_width(&self) -> usize { + self.max_width + } + + /// Enumerate all valid spans up to max_width. + /// + /// Returns Vec of (start_word, end_word) pairs. + pub fn enumerate_spans(&self, num_words: usize) -> Vec<(usize, usize)> { + let mut spans = Vec::new(); + for start in 0..num_words { + for width in 1..=self.max_width.min(num_words - start) { + let end = start + width - 1; + spans.push((start, end)); + } + } + spans + } + + /// Generate span representations from word embeddings. + /// + /// # Arguments + /// * `word_embeddings` - Word embeddings [batch, num_words, hidden_dim] + /// * `device` - Device for output tensors + /// + /// # Returns + /// * Span embeddings [batch, num_spans, hidden_dim] + /// * Span indices (start_word, end_word) for each span + pub fn forward( + &self, + word_embeddings: &Tensor<3, f32>, + device: &Device, + ) -> (Tensor<3, f32>, Vec<(usize, usize)>) { + let shape = word_embeddings.shape(); + let batch_size = shape[0]; + let num_words = shape[1]; + let hidden_dim = shape[2]; + + // Enumerate all valid spans + let span_indices = self.enumerate_spans(num_words); + let num_spans = span_indices.len(); + + if num_spans == 0 { + // Return empty tensor if no spans + let empty = Tensor::zeros(device, [batch_size, 1, hidden_dim]); + return (empty, vec![(0, 0)]); + } + + // Build start and end embeddings for all spans + let (start_emb, end_emb) = + self.gather_span_embeddings(word_embeddings, &span_indices, device); + + // Project start embeddings: [batch, num_spans, hidden_dim] + // create_projection_layer uses: Linear -> ReLU -> Dropout -> Linear + let start_projected = self + .start_fc2 + .forward(&self.start_fc1.forward(&start_emb).relu()); + + // Project end embeddings: [batch, num_spans, hidden_dim] + let end_projected = self.end_fc2.forward(&self.end_fc1.forward(&end_emb).relu()); + + // Concatenate: [batch, num_spans, 2 * hidden_dim] + // Python does: cat([start, end]).relu() before out_project + let combined = Tensor::cat( + [start_projected.to_concrete(), end_projected.to_concrete()], + 2, + ) + .relu(); + + // Output projection: [batch, num_spans, hidden_dim] + let span_embeddings = self + .out_fc2 + .forward(&self.out_fc1.forward(&combined).relu()); + + (span_embeddings, span_indices) + } + + fn gather_span_embeddings( + &self, + word_embeddings: &Tensor<3, f32>, + span_indices: &[(usize, usize)], + device: &Device, + ) -> (Tensor<3, f32>, Tensor<3, f32>) { + let shape = word_embeddings.shape(); + let batch_size = shape[0]; + let num_words = shape[1]; + let hidden_dim = shape[2]; + let num_spans = span_indices.len(); + + // Create index tensors for gathering + let start_indices: Vec = span_indices.iter().map(|(s, _)| *s as u32).collect(); + let end_indices: Vec = span_indices.iter().map(|(_, e)| *e as u32).collect(); + + // Flatten word_embeddings to [batch * num_words, hidden_dim] + let word_embeddings_concrete = word_embeddings.to_concrete(); + let flat_embeddings = word_embeddings_concrete + .reshape([batch_size * num_words, hidden_dim]) + .to_concrete(); + + // Build offset indices for batch processing + let mut start_offset_indices: Vec = Vec::with_capacity(batch_size * num_spans); + let mut end_offset_indices: Vec = Vec::with_capacity(batch_size * num_spans); + + for batch_idx in 0..batch_size { + let offset = (batch_idx * num_words) as u32; + for &start in &start_indices { + start_offset_indices.push(start + offset); + } + } + for batch_idx in 0..batch_size { + let offset = (batch_idx * num_words) as u32; + for &end in &end_indices { + end_offset_indices.push(end + offset); + } + } + + let start_idx_tensor = Tensor::new(device, &start_offset_indices); + let end_idx_tensor = Tensor::new(device, &end_offset_indices); + + // Gather start and end embeddings + let start_emb = flat_embeddings.index_select(0, &start_idx_tensor); + let end_emb = flat_embeddings.index_select(0, &end_idx_tensor); + + // Reshape to [batch, num_spans, hidden_dim] + let start_emb = start_emb + .reshape([batch_size, num_spans, hidden_dim]) + .to_concrete(); + let end_emb = end_emb + .reshape([batch_size, num_spans, hidden_dim]) + .to_concrete(); + + (start_emb, end_emb) + } +} diff --git a/models/rgliner/src/raw/text_encoder.rs b/models/rgliner/src/raw/text_encoder.rs new file mode 100644 index 000000000..19fd945a3 --- /dev/null +++ b/models/rgliner/src/raw/text_encoder.rs @@ -0,0 +1,54 @@ +//! Text encoder wrapper for GLiNER. + +use fusor::{Device, Result, Tensor, VarBuilder}; + +use super::modern_bert::ModernBertModel; + +/// Text encoder for GLiNER (ModernBERT/Ettin). +pub struct TextEncoder { + model: ModernBertModel, +} + +impl TextEncoder { + /// Load text encoder from GGUF weights. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + // GLiNER GGUF uses "text." prefix for text encoder weights + let model = ModernBertModel::load(device, &mut vb.pp("text"))?; + Ok(Self { model }) + } + + /// Forward pass returning per-token embeddings. + /// + /// Returns: [batch_size, seq_len, hidden_size] + pub fn forward( + &self, + input_ids: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + self.model.forward(input_ids, attention_mask) + } + + /// Get the maximum sequence length. + pub fn max_seq_len(&self) -> usize { + self.model.max_seq_len() + } + + /// Get the embedding dimension. + pub fn embedding_dim(&self) -> usize { + self.model.embedding_dim() + } + + /// Get the device. + pub fn device(&self) -> &Device { + self.model.device() + } + + #[cfg(test)] + pub fn debug_hidden_states( + &self, + input_ids: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Vec> { + self.model.debug_hidden_states(input_ids, attention_mask) + } +} diff --git a/models/rgliner/src/source.rs b/models/rgliner/src/source.rs new file mode 100644 index 000000000..864945677 --- /dev/null +++ b/models/rgliner/src/source.rs @@ -0,0 +1,310 @@ +//! Model source configuration for GLiNER variants. + +use kalosm_model_types::FileSource; +use std::path::{Path, PathBuf}; + +/// Source configuration for GLiNER models. +/// +/// Specifies where to download the model files from. +pub struct GlinerSource { + /// Main model GGUF file (text encoder + span layer weights) + pub(crate) model: FileSource, + /// Label encoder GGUF file (sentence transformer) + pub(crate) label_encoder: FileSource, + /// Label encoder config JSON file + pub(crate) label_encoder_config: FileSource, + /// Label encoder tokenizer JSON file + pub(crate) label_encoder_tokenizer: FileSource, + /// Tokenizer JSON file (for text encoder) + pub(crate) tokenizer: FileSource, + /// GLiNER config JSON file + pub(crate) config: FileSource, +} + +impl GlinerSource { + fn huggingface_or_cached(model_id: &str, revision: &str, file: &str) -> FileSource { + if let Some(path) = Self::find_cached_hf_file(model_id, revision, file) { + FileSource::local(path) + } else { + FileSource::huggingface(model_id.to_string(), revision.to_string(), file.to_string()) + } + } + + fn find_cached_hf_file(model_id: &str, revision: &str, file: &str) -> Option { + let snapshots_dir = Self::huggingface_cache_dir()? + .join("hub") + .join(format!("models--{}", model_id.replace('/', "--"))) + .join("snapshots"); + + if !snapshots_dir.exists() { + return None; + } + + let file = Path::new(file); + + if revision != "main" { + let candidate = snapshots_dir.join(revision).join(file); + if candidate.exists() { + return Some(candidate); + } + } + + let refs_path = snapshots_dir + .parent() + .map(|parent| parent.join("refs").join(revision)); + if let Some(refs_path) = refs_path { + if let Ok(snapshot) = std::fs::read_to_string(&refs_path) { + let candidate = snapshots_dir.join(snapshot.trim()).join(file); + if candidate.exists() { + return Some(candidate); + } + } + } + + let entries = std::fs::read_dir(&snapshots_dir).ok()?; + for entry in entries.flatten() { + let candidate = entry.path().join(file); + if candidate.exists() { + return Some(candidate); + } + } + + None + } + + fn huggingface_cache_dir() -> Option { + if let Some(hf_home) = std::env::var_os("HF_HOME") { + return Some(PathBuf::from(hf_home)); + } + + if let Some(xdg_cache) = std::env::var_os("XDG_CACHE_HOME") { + return Some(PathBuf::from(xdg_cache).join("huggingface")); + } + + std::env::var_os("HOME").map(|home| PathBuf::from(home).join(".cache").join("huggingface")) + } + + /// GLiNER bi-encoder v2.0 Edge variant (60M parameters). + /// + /// The smallest and fastest variant, using: + /// - Text encoder: ettin-encoder-32m + /// - Label encoder: all-MiniLM-L6-v2 + pub fn edge() -> Self { + Self { + model: FileSource::huggingface( + "knowledgator/gliner-bi-edge-v2.0-gguf".to_string(), + "main".to_string(), + "gliner-bi-edge-v2.0-Q8_0.gguf".to_string(), + ), + label_encoder: FileSource::huggingface( + "knowledgator/gliner-bi-edge-v2.0-gguf".to_string(), + "main".to_string(), + "label-encoder-Q8_0.gguf".to_string(), + ), + label_encoder_config: FileSource::huggingface( + "sentence-transformers/all-MiniLM-L6-v2".to_string(), + "main".to_string(), + "config.json".to_string(), + ), + label_encoder_tokenizer: FileSource::huggingface( + "sentence-transformers/all-MiniLM-L6-v2".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + tokenizer: FileSource::huggingface( + "knowledgator/gliner-bi-edge-v2.0".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + config: FileSource::huggingface( + "knowledgator/gliner-bi-edge-v2.0".to_string(), + "main".to_string(), + "gliner_config.json".to_string(), + ), + } + } + + /// GLiNER bi-encoder v2.0 Small variant (108M parameters). + /// + /// Good balance of speed and accuracy, using: + /// - Text encoder: ettin-encoder-68m + /// - Label encoder: all-MiniLM-L12-v2 + pub fn small() -> Self { + Self { + model: FileSource::huggingface( + "knowledgator/gliner-bi-small-v2.0-gguf".to_string(), + "main".to_string(), + "gliner-bi-small-v2.0-Q8_0.gguf".to_string(), + ), + label_encoder: FileSource::huggingface( + "knowledgator/gliner-bi-small-v2.0-gguf".to_string(), + "main".to_string(), + "label-encoder-Q8_0.gguf".to_string(), + ), + label_encoder_config: FileSource::huggingface( + "sentence-transformers/all-MiniLM-L12-v2".to_string(), + "main".to_string(), + "config.json".to_string(), + ), + label_encoder_tokenizer: FileSource::huggingface( + "sentence-transformers/all-MiniLM-L12-v2".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + tokenizer: FileSource::huggingface( + "knowledgator/gliner-bi-small-v2.0".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + config: FileSource::huggingface( + "knowledgator/gliner-bi-small-v2.0".to_string(), + "main".to_string(), + "gliner_config.json".to_string(), + ), + } + } + + /// GLiNER bi-encoder v2.0 Base variant (194M parameters). + /// + /// Default variant with good accuracy, using: + /// - Text encoder: ettin-encoder-150m + /// - Label encoder: bge-small-en-v1.5 + pub fn base() -> Self { + Self { + model: FileSource::huggingface( + "knowledgator/gliner-bi-base-v2.0-gguf".to_string(), + "main".to_string(), + "gliner-bi-base-v2.0-Q8_0.gguf".to_string(), + ), + label_encoder: FileSource::huggingface( + "knowledgator/gliner-bi-base-v2.0-gguf".to_string(), + "main".to_string(), + "label-encoder-Q8_0.gguf".to_string(), + ), + label_encoder_config: FileSource::huggingface( + "BAAI/bge-small-en-v1.5".to_string(), + "main".to_string(), + "config.json".to_string(), + ), + label_encoder_tokenizer: FileSource::huggingface( + "BAAI/bge-small-en-v1.5".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + tokenizer: FileSource::huggingface( + "knowledgator/gliner-bi-base-v2.0".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + config: FileSource::huggingface( + "knowledgator/gliner-bi-base-v2.0".to_string(), + "main".to_string(), + "gliner_config.json".to_string(), + ), + } + } + + /// GLiNER bi-encoder v2.0 Large variant (530M parameters). + /// + /// Highest accuracy variant, using: + /// - Text encoder: ettin-encoder-400m + /// - Label encoder: bge-base-en-v1.5 + pub fn large() -> Self { + Self { + model: FileSource::huggingface( + "knowledgator/gliner-bi-large-v2.0-gguf".to_string(), + "main".to_string(), + "gliner-bi-large-v2.0-Q8_0.gguf".to_string(), + ), + label_encoder: FileSource::huggingface( + "knowledgator/gliner-bi-large-v2.0-gguf".to_string(), + "main".to_string(), + "label-encoder-Q8_0.gguf".to_string(), + ), + label_encoder_config: FileSource::huggingface( + "BAAI/bge-base-en-v1.5".to_string(), + "main".to_string(), + "config.json".to_string(), + ), + label_encoder_tokenizer: FileSource::huggingface( + "BAAI/bge-base-en-v1.5".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + tokenizer: FileSource::huggingface( + "knowledgator/gliner-bi-large-v2.0".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + config: FileSource::huggingface( + "knowledgator/gliner-bi-large-v2.0".to_string(), + "main".to_string(), + "gliner_config.json".to_string(), + ), + } + } + + /// Create a custom source with specific file locations. + pub fn custom( + model: FileSource, + label_encoder: FileSource, + label_encoder_config: FileSource, + label_encoder_tokenizer: FileSource, + tokenizer: FileSource, + config: FileSource, + ) -> Self { + Self { + model, + label_encoder, + label_encoder_config, + label_encoder_tokenizer, + tokenizer, + config, + } + } +} + +impl Default for GlinerSource { + fn default() -> Self { + Self::base() + } +} + +impl GlinerSource { + /// Create a source from local GGUF files (for testing converted models). + /// + /// # Arguments + /// * `model_path` - Path to main model GGUF (text encoder + span layer + projection) + /// * `label_encoder_path` - Path to label encoder GGUF (BERT/MiniLM) + pub fn local( + model_path: impl Into, + label_encoder_path: impl Into, + ) -> Self { + let model_path = model_path.into(); + let label_encoder_path = label_encoder_path.into(); + Self { + model: FileSource::local(model_path), + label_encoder: FileSource::local(label_encoder_path), + label_encoder_config: Self::huggingface_or_cached( + "sentence-transformers/all-MiniLM-L6-v2", + "main", + "config.json", + ), + label_encoder_tokenizer: Self::huggingface_or_cached( + "sentence-transformers/all-MiniLM-L6-v2", + "main", + "tokenizer.json", + ), + tokenizer: Self::huggingface_or_cached( + "knowledgator/gliner-bi-edge-v2.0", + "main", + "tokenizer.json", + ), + config: Self::huggingface_or_cached( + "knowledgator/gliner-bi-edge-v2.0", + "main", + "gliner_config.json", + ), + } + } +} diff --git a/models/rgliner/src/tokenization.rs b/models/rgliner/src/tokenization.rs new file mode 100644 index 000000000..1d11ca38a --- /dev/null +++ b/models/rgliner/src/tokenization.rs @@ -0,0 +1,278 @@ +//! Word-level tokenization with subtoken-to-word mapping. + +use fusor::{Device, Tensor}; +use tokenizers::Tokenizer; + +use crate::error::GlinerError; + +/// Tokenization result with word-level alignment. +#[derive(Debug, Clone)] +pub struct TokenizedText { + /// Token IDs for the model. + pub token_ids: Vec, + /// Attention mask (1 for real tokens, 0 for padding). + pub attention_mask: Vec, + /// Maps each token position to its word index (-1 for special tokens). + pub token_to_word: Vec, + /// Index of the first token for each word. + pub word_first_token: Vec, + /// Number of words in the input. + pub num_words: usize, + /// Character offsets for each word: (start_char, end_char). + pub word_offsets: Vec<(usize, usize)>, +} + +/// Word-level tokenizer wrapper. +pub struct WordTokenizer { + tokenizer: Tokenizer, +} + +impl WordTokenizer { + /// Create a new word tokenizer from a HuggingFace tokenizer. + pub fn new(tokenizer: Tokenizer) -> Self { + Self { tokenizer } + } + + /// Load tokenizer from JSON bytes. + pub fn from_bytes(bytes: &[u8]) -> Result { + let tokenizer = Tokenizer::from_bytes(bytes)?; + Ok(Self::new(tokenizer)) + } + + /// Tokenize text and track word boundaries. + pub fn tokenize(&self, text: &str) -> Result { + let split_words = split_words(text); + let words: Vec = split_words.iter().map(|(word, _)| word.clone()).collect(); + let word_offsets: Vec<(usize, usize)> = + split_words.iter().map(|(_, offsets)| *offsets).collect(); + + let encoding = self + .tokenizer + .encode(words, true) + .map_err(GlinerError::Tokenizer)?; + + let token_ids = encoding.get_ids().to_vec(); + let attention_mask = encoding + .get_attention_mask() + .iter() + .map(|&x| x as u32) + .collect(); + + // Build token-to-word mapping + // word_ids() returns Option for each token + let token_to_word: Vec = encoding + .get_word_ids() + .iter() + .map(|opt| opt.map(|w| w as i32).unwrap_or(-1)) + .collect(); + + let num_words = word_offsets.len(); + + // Find first token index for each word + let mut word_first_token = vec![0usize; num_words]; + let mut seen_words = vec![false; num_words]; + for (token_idx, &word_id) in token_to_word.iter().enumerate() { + if word_id >= 0 { + let word_id = word_id as usize; + if !seen_words[word_id] { + word_first_token[word_id] = token_idx; + seen_words[word_id] = true; + } + } + } + + Ok(TokenizedText { + token_ids, + attention_mask, + token_to_word, + word_first_token, + num_words, + word_offsets, + }) + } + + /// Tokenize a batch of texts. + pub fn tokenize_batch(&self, texts: &[&str]) -> Result, GlinerError> { + texts.iter().map(|text| self.tokenize(text)).collect() + } +} + +fn split_words(text: &str) -> Vec<(String, (usize, usize))> { + let mut words = Vec::new(); + let mut chars = text.char_indices().peekable(); + + while let Some((start, ch)) = chars.peek().copied() { + if ch.is_whitespace() { + chars.next(); + continue; + } + + if is_word_char(ch) { + chars.next(); + let mut end = start + ch.len_utf8(); + + while let Some((idx, next_ch)) = chars.peek().copied() { + if is_word_char(next_ch) { + end = idx + next_ch.len_utf8(); + chars.next(); + continue; + } + + if matches!(next_ch, '-' | '_') { + let mut lookahead = chars.clone(); + lookahead.next(); + if let Some((_, after_delimiter)) = lookahead.peek().copied() { + if is_word_char(after_delimiter) { + chars.next(); + end = idx + next_ch.len_utf8(); + + while let Some((word_idx, word_ch)) = chars.peek().copied() { + if !is_word_char(word_ch) { + break; + } + end = word_idx + word_ch.len_utf8(); + chars.next(); + } + continue; + } + } + } + + break; + } + + words.push((text[start..end].to_string(), (start, end))); + continue; + } + + chars.next(); + let end = start + ch.len_utf8(); + words.push((text[start..end].to_string(), (start, end))); + } + + words +} + +fn is_word_char(ch: char) -> bool { + ch.is_alphanumeric() || ch == '_' +} + +/// Pool token embeddings to word embeddings using first-subtoken strategy. +/// +/// # Arguments +/// * `token_embeddings` - Token embeddings [batch, seq_len, hidden_dim] +/// * `tokenized` - Tokenization results for each batch item +/// * `device` - Device to create output tensor on +/// +/// # Returns +/// * Word embeddings [batch, max_words, hidden_dim] +/// * Word mask [batch, max_words] - 1 for valid words, 0 for padding +pub fn first_subtoken_pooling( + token_embeddings: &Tensor<3, f32>, + tokenized: &[TokenizedText], + device: &Device, +) -> (Tensor<3, f32>, Tensor<2, u32>) { + let shape = token_embeddings.shape(); + let batch_size = shape[0]; + let hidden_dim = shape[2]; + + // Find max words across batch + let max_words = tokenized.iter().map(|t| t.num_words).max().unwrap_or(0); + + if max_words == 0 { + // Return empty tensors if no words + let word_emb = Tensor::zeros(device, [batch_size, 1, hidden_dim]); + let word_mask = Tensor::zeros(device, [batch_size, 1]); + return (word_emb, word_mask); + } + + // Build gather indices for each batch item + // For each batch, we need to gather word_first_token[w] for each word w + let mut all_indices: Vec = Vec::with_capacity(batch_size * max_words); + let mut mask_data: Vec = Vec::with_capacity(batch_size * max_words); + + for t in tokenized { + for word_idx in 0..max_words { + if word_idx < t.num_words { + all_indices.push(t.word_first_token[word_idx] as u32); + mask_data.push(1); + } else { + // Padding - use index 0 (will be masked out) + all_indices.push(0); + mask_data.push(0); + } + } + } + + // Create index tensor [batch_size * max_words] + let _indices = Tensor::new(device, &all_indices); + + // Reshape token_embeddings to [batch_size * seq_len, hidden_dim] for gathering + let seq_len = shape[1]; + let token_embeddings_concrete = token_embeddings.to_concrete(); + let flat_embeddings = token_embeddings_concrete + .reshape([batch_size * seq_len, hidden_dim]) + .to_concrete(); + + // For each batch, we need to offset the indices by batch_idx * seq_len + let mut offset_indices: Vec = Vec::with_capacity(batch_size * max_words); + for batch_idx in 0..batch_size { + let offset = (batch_idx * seq_len) as u32; + for word_idx in 0..max_words { + let idx = all_indices[batch_idx * max_words + word_idx]; + offset_indices.push(idx + offset); + } + } + let offset_indices_tensor = Tensor::new(device, &offset_indices); + + // Gather word embeddings + let gathered = flat_embeddings.index_select(0, &offset_indices_tensor); + + // Reshape to [batch_size, max_words, hidden_dim] + let word_embeddings = gathered + .reshape([batch_size, max_words, hidden_dim]) + .to_concrete(); + + // Create word mask + let word_mask = Tensor::new(device, &mask_data) + .reshape([batch_size, max_words]) + .to_concrete(); + + (word_embeddings, word_mask) +} + +#[cfg(test)] +mod tests { + use super::split_words; + + #[test] + fn split_words_matches_gliner_word_regex() { + let words = split_words("all-MiniLM_L6-v2 rocks."); + + assert_eq!( + words, + vec![ + ("all-MiniLM_L6-v2".to_string(), (0, 16)), + ("rocks".to_string(), (17, 22)), + (".".to_string(), (22, 23)), + ] + ); + } + + #[test] + fn split_words_keeps_punctuation_as_separate_words() { + let words = split_words("Apple Inc. was founded."); + + assert_eq!( + words, + vec![ + ("Apple".to_string(), (0, 5)), + ("Inc".to_string(), (6, 9)), + (".".to_string(), (9, 10)), + ("was".to_string(), (11, 14)), + ("founded".to_string(), (15, 22)), + (".".to_string(), (22, 23)), + ] + ); + } +} diff --git a/models/rgliner/tests/example_regression.rs b/models/rgliner/tests/example_regression.rs new file mode 100644 index 000000000..6e5182512 --- /dev/null +++ b/models/rgliner/tests/example_regression.rs @@ -0,0 +1,86 @@ +use std::path::PathBuf; + +use fusor::Device; +use rgliner::{Gliner, GlinerSource}; + +fn local_edge_source() -> GlinerSource { + let manifest_dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")); + let weights_dir = manifest_dir.join("weights"); + + GlinerSource::local( + weights_dir.join("gliner-edge.gguf"), + weights_dir.join("gliner-edge-label-encoder.gguf"), + ) +} + +#[test] +fn edge_example_sentences_regression() -> anyhow::Result<()> { + tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build()? + .block_on(async { + // Keep the regression on an explicit CPU backend; libtest's panic-hook + // environment can interfere with the auto-device probe, while the plain + // example covers the user-facing default path separately. + let mut gliner = Gliner::builder() + .with_source(local_edge_source()) + .with_device(Device::cpu()) + .build() + .await?; + + let labels = ["person", "organization", "location"]; + let cases = [ + ( + "Apple Inc. was founded by Steve Jobs in California.", + vec![ + ("organization", "Apple Inc."), + ("person", "Steve Jobs"), + ("location", "California"), + ], + ), + ( + "Microsoft Corporation is headquartered in Seattle.", + vec![ + ("organization", "Microsoft Corporation"), + ("location", "Seattle"), + ], + ), + ( + "Elon Musk is the CEO of Tesla.", + vec![("person", "Elon Musk"), ("organization", "Tesla")], + ), + ( + "Google was founded in Mountain View.", + vec![("organization", "Google"), ("location", "Mountain View")], + ), + ]; + + for (text, expected) in cases { + let uncached_entities = gliner.extract(text, &labels).await?; + let uncached: Vec<(&str, &str)> = uncached_entities + .iter() + .map(|entity| (entity.label.as_str(), entity.text.as_str())) + .collect(); + + assert_eq!( + uncached, expected, + "unexpected uncached entities for input: {text}" + ); + + gliner.cache_labels(&labels).await?; + let entities = gliner.extract_with_cached_labels(text).await?; + let actual: Vec<(&str, &str)> = entities + .iter() + .map(|entity| (entity.label.as_str(), entity.text.as_str())) + .collect(); + + assert_eq!(actual, expected, "unexpected entities for input: {text}"); + assert!( + entities.iter().all(|entity| entity.score >= 0.5), + "all expected entities should remain above the default threshold for input: {text}" + ); + } + + Ok(()) + }) +} From 7867d9a3f973d48fceea919edbf28619498f0685 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Sun, 12 Apr 2026 21:38:28 -0500 Subject: [PATCH 02/34] remote --- models/rgliner/src/source.rs | 39 ++++++++++++++++++++++++++++++++++++ 1 file changed, 39 insertions(+) diff --git a/models/rgliner/src/source.rs b/models/rgliner/src/source.rs index 864945677..8cedaaf4b 100644 --- a/models/rgliner/src/source.rs +++ b/models/rgliner/src/source.rs @@ -124,6 +124,45 @@ impl GlinerSource { } } + /// Demonthos GLiNER GGUF edge upload. + /// + /// Uses the GGUF weights and sidecar tokenizer/config files from + /// `Demonthos/gliner-gguf`. + pub fn demonthos_edge() -> Self { + Self { + model: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "gliner-edge.gguf", + ), + label_encoder: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "gliner-edge-label-encoder.gguf", + ), + label_encoder_config: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "label-encoder-config.json", + ), + label_encoder_tokenizer: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "label-encoder-tokenizer.json", + ), + tokenizer: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "text-tokenizer.json", + ), + config: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "text-gliner-config.json", + ), + } + } + /// GLiNER bi-encoder v2.0 Small variant (108M parameters). /// /// Good balance of speed and accuracy, using: From 08f2a3d39ace6767284e28c5e51bb4ed52453230 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 19:17:04 -0500 Subject: [PATCH 03/34] relex working! --- .claude/settings.local.json | 10 +- Cargo.lock | 1 + models/rgliner/Cargo.toml | 1 + models/rgliner/examples/relex.rs | 76 ++ .../rgliner/scripts/convert_relex_to_gguf.py | 390 +++++++++ models/rgliner/scripts/debug_compare.py | 173 ++++ models/rgliner/src/error.rs | 3 + models/rgliner/src/lib.rs | 25 + models/rgliner/src/raw/bilstm.rs | 241 ++++++ models/rgliner/src/raw/joint_scorer.rs | 282 +++++++ models/rgliner/src/raw/mdeberta/attention.rs | 483 +++++++++++ models/rgliner/src/raw/mdeberta/config.rs | 138 +++ .../rgliner/src/raw/mdeberta/feed_forward.rs | 31 + models/rgliner/src/raw/mdeberta/layer.rs | 87 ++ models/rgliner/src/raw/mdeberta/mod.rs | 12 + models/rgliner/src/raw/mdeberta/model.rs | 181 ++++ models/rgliner/src/raw/mod.rs | 9 + models/rgliner/src/raw/pair_projector.rs | 141 ++++ models/rgliner/src/raw/relations_layer.rs | 120 +++ models/rgliner/src/raw/span_layer.rs | 61 ++ models/rgliner/src/relation_decoding.rs | 292 +++++++ models/rgliner/src/relex.rs | 790 ++++++++++++++++++ models/rgliner/src/relex_tokenization.rs | 278 ++++++ models/rgliner/src/source.rs | 125 ++- 24 files changed, 3945 insertions(+), 5 deletions(-) create mode 100644 models/rgliner/examples/relex.rs create mode 100644 models/rgliner/scripts/convert_relex_to_gguf.py create mode 100644 models/rgliner/scripts/debug_compare.py create mode 100644 models/rgliner/src/raw/bilstm.rs create mode 100644 models/rgliner/src/raw/joint_scorer.rs create mode 100644 models/rgliner/src/raw/mdeberta/attention.rs create mode 100644 models/rgliner/src/raw/mdeberta/config.rs create mode 100644 models/rgliner/src/raw/mdeberta/feed_forward.rs create mode 100644 models/rgliner/src/raw/mdeberta/layer.rs create mode 100644 models/rgliner/src/raw/mdeberta/mod.rs create mode 100644 models/rgliner/src/raw/mdeberta/model.rs create mode 100644 models/rgliner/src/raw/pair_projector.rs create mode 100644 models/rgliner/src/raw/relations_layer.rs create mode 100644 models/rgliner/src/relation_decoding.rs create mode 100644 models/rgliner/src/relex.rs create mode 100644 models/rgliner/src/relex_tokenization.rs diff --git a/.claude/settings.local.json b/.claude/settings.local.json index 076b30c7c..5eaa6721c 100644 --- a/.claude/settings.local.json +++ b/.claude/settings.local.json @@ -10,7 +10,15 @@ "Bash(GLINER_MODEL=/Users/evanalmloff/Desktop/Github/ner/models/rgliner/weights/gliner-edge.gguf cargo run:*)", "Bash(RUST_BACKTRACE=1 GLINER_MODEL=/Users/evanalmloff/Desktop/Github/ner/models/rgliner/weights/gliner-edge.gguf cargo run:*)", "Bash(git -C /Users/evanalmloff/Desktop/Github/ner checkout -- fusor-ml/core/src/matmul/mod.rs)", - "WebFetch(domain:api.github.com)" + "WebFetch(domain:api.github.com)", + "WebFetch(domain:arxiv.org)", + "WebFetch(domain:urchade.github.io)", + "Bash(python3:*)", + "Bash(RUST_BACKTRACE=1 cargo run:*)", + "Bash(GLINER_MODEL=/Users/evanalmloff/Desktop/Github/ner/models/rgliner/weights/gliner-relex-multi-v1.0.gguf cargo run:*)", + "Bash(GLINER_MODEL=./weights/gliner-relex-multi-v1.0.gguf cargo run:*)", + "Bash(cargo add:*)", + "Bash(pip install:*)" ] } } diff --git a/Cargo.lock b/Cargo.lock index dff86fcb6..4fc1787ba 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7681,6 +7681,7 @@ dependencies = [ "kalosm-common", "kalosm-language-model", "kalosm-model-types", + "pollster", "rbert", "serde", "serde_json", diff --git a/models/rgliner/Cargo.toml b/models/rgliner/Cargo.toml index aaaad55e6..c4a8a63d0 100644 --- a/models/rgliner/Cargo.toml +++ b/models/rgliner/Cargo.toml @@ -23,6 +23,7 @@ serde = { version = "1", features = ["derive"] } kalosm-common = { workspace = true } kalosm-model-types.workspace = true kalosm-language-model.workspace = true +pollster = "0.4.0" [dev-dependencies] anyhow.workspace = true diff --git a/models/rgliner/examples/relex.rs b/models/rgliner/examples/relex.rs new file mode 100644 index 000000000..799e41899 --- /dev/null +++ b/models/rgliner/examples/relex.rs @@ -0,0 +1,76 @@ +//! Example of using GlinerRelEx for joint NER and relation extraction. +//! +//! Run with: +//! ``` +//! cargo run --example relex -p rgliner +//! ``` +//! +//! The GGUF file has the tokenizer and GLiNER config baked in as metadata, +//! so only the model file path is needed. + +use rgliner::relex::{GlinerRelEx, GlinerRelExSource}; +use std::env; +use std::path::PathBuf; + +fn get_model_path() -> Option { + let weights_dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("weights"); + let default_path = weights_dir.join("gliner-relex-multi-v1.0.gguf"); + let model_path = env::var("GLINER_MODEL") + .map(PathBuf::from) + .unwrap_or(default_path); + + if !model_path.exists() { + eprintln!("Model file not found: {:?}", model_path); + return None; + } + Some(model_path) +} + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + let model_path = get_model_path().ok_or_else(|| { + anyhow::anyhow!( + "Model file not found. Convert with:\n \ + python scripts/convert_relex_to_gguf.py -m knowledgator/gliner-relex-multi-v1.0 \ + -o weights/gliner-relex-multi-v1.0.gguf" + ) + })?; + + println!("Loading model from: {:?}", model_path); + + let source = GlinerRelExSource::local(model_path); + + let text = "Perfect! Now I have all the information I need to create a comprehensive GGUF inference framework from scratch. Let me create a single-file implementation with full SIMD support targeting nightly Rust"; + let entity_labels = ["technology", "language", "file format"]; + let relation_labels = ["supported by", "implemented with"]; + + println!("\nText: {}", text); + println!("Entity labels: {:?}", entity_labels); + println!("Relation labels: {:?}", relation_labels); + + let relex = GlinerRelEx::builder() + .with_source(source) + .with_entity_threshold(0.1) + .build() + .await?; + + let (entities, relations) = relex.extract(text, &entity_labels, &relation_labels).await?; + + println!("\nEntities found ({}):", entities.len()); + for entity in &entities { + println!( + " {} [{}] (score: {:.3})", + entity.text, entity.label, entity.score + ); + } + + println!("\nRelations found ({}):", relations.len()); + for relation in &relations { + println!( + " {} --[{}]--> {} (score: {:.3})", + relation.head.text, relation.relation, relation.tail.text, relation.score + ); + } + + Ok(()) +} diff --git a/models/rgliner/scripts/convert_relex_to_gguf.py b/models/rgliner/scripts/convert_relex_to_gguf.py new file mode 100644 index 000000000..e60535bf0 --- /dev/null +++ b/models/rgliner/scripts/convert_relex_to_gguf.py @@ -0,0 +1,390 @@ +#!/usr/bin/env python3 +""" +Convert GLiNER-RelEx PyTorch models to GGUF format. + +Usage: + python convert_relex_to_gguf.py --model knowledgator/gliner-relex-multi-v1.0 --output gliner-relex-multi-v1.0.gguf + +This script converts the GLiNER-RelEx model including: +1. mDeBERTa-v3 encoder +2. Span representation layer +3. Relations representation layer (adjacency scoring) +4. Pair projector +5. Entity/relation label FFNs +""" + +import argparse +import json +import os +import struct +import sys +from pathlib import Path +from typing import Any, Dict, List, Tuple + +import numpy as np +import torch +from huggingface_hub import snapshot_download + +# GGUF constants (same as convert_to_gguf.py) +GGUF_MAGIC = 0x46554747 +GGUF_VERSION = 3 + +GGUF_TYPE_UINT8 = 0 +GGUF_TYPE_INT8 = 1 +GGUF_TYPE_UINT16 = 2 +GGUF_TYPE_INT16 = 3 +GGUF_TYPE_UINT32 = 4 +GGUF_TYPE_INT32 = 5 +GGUF_TYPE_FLOAT32 = 6 +GGUF_TYPE_BOOL = 7 +GGUF_TYPE_STRING = 8 +GGUF_TYPE_ARRAY = 9 +GGUF_TYPE_UINT64 = 10 +GGUF_TYPE_INT64 = 11 +GGUF_TYPE_FLOAT64 = 12 + +GGML_TYPE_F32 = 0 +GGML_TYPE_F16 = 1 +GGML_TYPE_Q8_0 = 8 +GGML_TYPE_BF16 = 30 + + +class GGUFWriter: + """Simple GGUF file writer.""" + + def __init__(self, path: str): + self.path = path + self.metadata: Dict[str, Any] = {} + self.tensors: List[Tuple[str, np.ndarray, int]] = [] + + def add_metadata(self, key: str, value: Any): + self.metadata[key] = value + + def add_tensor(self, name: str, data: np.ndarray, ggml_type: int = GGML_TYPE_F32): + self.tensors.append((name, data, ggml_type)) + + def _write_string(self, f, s: str): + encoded = s.encode('utf-8') + f.write(struct.pack('> 16) & 0xFFFF).astype(np.uint16) + + self._write_string(f, name) + f.write(struct.pack(' Tuple[Dict[str, torch.Tensor], Dict, str, str]: + """Load PyTorch model from HuggingFace. + + Returns: + state_dict, config (parsed), gliner_config_json (raw text), tokenizer_json (raw text) + """ + print(f"Downloading model: {model_id}") + + model_dir = snapshot_download( + model_id, + cache_dir=cache_dir, + allow_patterns=["*.bin", "*.json", "*.safetensors"] + ) + + config_path = os.path.join(model_dir, "gliner_config.json") + with open(config_path, 'r') as f: + gliner_config_json = f.read() + config = json.loads(gliner_config_json) + + tokenizer_path = os.path.join(model_dir, "tokenizer.json") + with open(tokenizer_path, 'r') as f: + tokenizer_json = f.read() + + weights_path = os.path.join(model_dir, "pytorch_model.bin") + if os.path.exists(weights_path): + print(f"Loading weights from: {weights_path}") + state_dict = torch.load(weights_path, map_location='cpu', weights_only=True) + else: + from safetensors.torch import load_file + weights_path = os.path.join(model_dir, "model.safetensors") + print(f"Loading weights from: {weights_path}") + state_dict = load_file(weights_path) + + return state_dict, config, gliner_config_json, tokenizer_json + + +def map_relex_weight_name(pytorch_name: str) -> str: + """Map PyTorch weight names to GGUF conventions for GLiNER-RelEx.""" + name = pytorch_name + + # ===== mDeBERTa Encoder (token_rep_layer.bert_layer.model.*) ===== + name = name.replace("token_rep_layer.bert_layer.model.embeddings.word_embeddings.weight", "text.token_embd.weight") + name = name.replace("token_rep_layer.bert_layer.model.embeddings.LayerNorm.weight", "text.embd_norm.weight") + name = name.replace("token_rep_layer.bert_layer.model.embeddings.LayerNorm.bias", "text.embd_norm.bias") + + # Relative position embeddings + name = name.replace("token_rep_layer.bert_layer.model.encoder.rel_embeddings.weight", "text.rel_pos_embd.weight") + name = name.replace("token_rep_layer.bert_layer.model.encoder.LayerNorm.weight", "text.output_norm.weight") + name = name.replace("token_rep_layer.bert_layer.model.encoder.LayerNorm.bias", "text.output_norm.bias") + + # DeBERTa layers + name = name.replace("token_rep_layer.bert_layer.model.encoder.layer.", "text.blk.") + + # DeBERTa attention + name = name.replace(".attention.self.query_proj.", ".attention.query.") + name = name.replace(".attention.self.key_proj.", ".attention.key.") + name = name.replace(".attention.self.value_proj.", ".attention.value.") + name = name.replace(".attention.self.pos_proj.", ".attention.pos_proj.") + name = name.replace(".attention.self.pos_q_proj.", ".attention.pos_q_proj.") + name = name.replace(".attention.output.dense.", ".attention.output.") + name = name.replace(".attention.output.LayerNorm.", ".attention_norm.") + + # DeBERTa FFN + name = name.replace(".intermediate.dense.", ".ffn.intermediate.") + name = name.replace(".output.dense.", ".ffn.output.") + name = name.replace(".output.LayerNorm.", ".output_norm.") + + # ===== Span Representation Layer ===== + name = name.replace("span_rep_layer.span_rep_layer.project_start.0.", "span.start_fc1.") + name = name.replace("span_rep_layer.span_rep_layer.project_start.3.", "span.start_fc2.") + name = name.replace("span_rep_layer.span_rep_layer.project_end.0.", "span.end_fc1.") + name = name.replace("span_rep_layer.span_rep_layer.project_end.3.", "span.end_fc2.") + name = name.replace("span_rep_layer.span_rep_layer.out_project.0.", "span.out_fc1.") + name = name.replace("span_rep_layer.span_rep_layer.out_project.3.", "span.out_fc2.") + + # ===== BiLSTM (rnn.lstm.*) ===== + name = name.replace("rnn.lstm.", "rnn.") + + # ===== Scorer (scorer.*) ===== + # Keep scorer.* as-is: scorer.proj_token, scorer.proj_label, scorer.out_mlp + + # ===== Pair Representation Layer (pair_rep_layer.*) ===== + name = name.replace("pair_rep_layer.0.", "pair_proj.0.") + name = name.replace("pair_rep_layer.3.", "pair_proj.3.") + + # ===== Prompt Representation Layer (prompt_rep_layer.*) ===== + # Keep prompt_rep_layer.* as-is + + return name + + +def convert_relex_to_gguf( + model_id: str, + output_path: str, + quantize: str = "f32", + cache_dir: str = None +): + """Convert GLiNER-RelEx model to GGUF format.""" + + state_dict, config, gliner_config_json, tokenizer_json = load_pytorch_model(model_id, cache_dir) + + print("\nModel weights:") + for name, tensor in state_dict.items(): + print(f" {name}: {tensor.shape} {tensor.dtype}") + + if quantize == "f32": + ggml_type = GGML_TYPE_F32 + elif quantize == "f16": + ggml_type = GGML_TYPE_F16 + elif quantize == "bf16": + ggml_type = GGML_TYPE_BF16 + else: + raise ValueError(f"Unsupported quantization: {quantize}") + + writer = GGUFWriter(output_path) + + # Add metadata + writer.add_metadata("general.architecture", "gliner-relex") + writer.add_metadata("general.name", model_id.split("/")[-1]) + writer.add_metadata("general.quantization_version", 2) + + # GLiNER-RelEx specific metadata + writer.add_metadata("gliner.max_width", config.get("max_width", 12)) + writer.add_metadata("gliner.span_mode", config.get("span_mode", "markerV0")) + writer.add_metadata("gliner.subtoken_pooling", config.get("subtoken_pooling", "first")) + + # mDeBERTa config + encoder_config = config.get("encoder_config", {}) + hidden_size = encoder_config.get("hidden_size", 768) + num_heads = encoder_config.get("num_attention_heads", 12) + num_layers = encoder_config.get("num_hidden_layers", 12) + intermediate_size = encoder_config.get("intermediate_size", 3072) + vocab_size = encoder_config.get("vocab_size", 250105) + context_length = encoder_config.get("max_position_embeddings", 512) + max_relative_positions = encoder_config.get("max_relative_positions", 512) + + # Handle -1 which means "use full context" + if max_relative_positions <= 0: + # Derive from rel_embeddings shape if available, otherwise use context_length // 2 + rel_emb_key = "token_rep_layer.bert_layer.model.encoder.rel_embeddings.weight" + if rel_emb_key in state_dict: + num_positions = state_dict[rel_emb_key].shape[0] + max_relative_positions = num_positions // 2 + print(f"Derived max_relative_positions from rel_embeddings: {max_relative_positions}") + else: + max_relative_positions = context_length // 2 + print(f"Using default max_relative_positions: {max_relative_positions}") + + writer.add_metadata("gliner.attention.head_count", num_heads) + writer.add_metadata("gliner.block_count", num_layers) + writer.add_metadata("gliner.embedding_length", hidden_size) + writer.add_metadata("gliner.feed_forward_length", intermediate_size) + writer.add_metadata("gliner.context_length", context_length) + writer.add_metadata("gliner.attention.max_relative_positions", max_relative_positions) + writer.add_metadata("gliner.attention.layer_norm_epsilon", 1e-7) + writer.add_metadata("gliner.vocab_size", vocab_size) + + # RelEx-specific metadata + writer.add_metadata("gliner.relex.ent_token_id", 250102) # <> + writer.add_metadata("gliner.relex.rel_token_id", 250104) # <> + + # Embed tokenizer.json and gliner_config.json as string metadata so the + # GGUF is self-contained (no separate files needed at inference time). + writer.add_metadata("gliner.tokenizer_json", tokenizer_json) + writer.add_metadata("gliner.config_json", gliner_config_json) + + # Convert tensors + print(f"\nConverting {len(state_dict)} tensors to GGUF...") + + for pytorch_name, tensor in state_dict.items(): + gguf_name = map_relex_weight_name(pytorch_name) + + try: + data = tensor.detach().float().cpu().numpy() + except RuntimeError: + import array + t = tensor.detach().float().cpu().contiguous() + data = np.frombuffer( + array.array('f', t.flatten().tolist()), + dtype=np.float32 + ).reshape(t.shape) + + print(f" {pytorch_name} -> {gguf_name} {data.shape}") + writer.add_tensor(gguf_name, data, ggml_type) + + writer.write() + print(f"\nOutput: {output_path}") + print(f"Size: {os.path.getsize(output_path) / 1024 / 1024:.2f} MB") + print("\nConversion complete!") + + +def main(): + parser = argparse.ArgumentParser(description="Convert GLiNER-RelEx PyTorch models to GGUF") + parser.add_argument( + "--model", "-m", + type=str, + required=True, + help="HuggingFace model ID (e.g., knowledgator/gliner-relex-multi-v1.0)" + ) + parser.add_argument( + "--output", "-o", + type=str, + required=True, + help="Output GGUF file path" + ) + parser.add_argument( + "--quantize", "-q", + type=str, + default="f32", + choices=["f32", "f16", "bf16"], + help="Quantization type (default: f32)" + ) + parser.add_argument( + "--cache-dir", + type=str, + default=None, + help="HuggingFace cache directory" + ) + + args = parser.parse_args() + + convert_relex_to_gguf( + model_id=args.model, + output_path=args.output, + quantize=args.quantize, + cache_dir=args.cache_dir + ) + + +if __name__ == "__main__": + main() diff --git a/models/rgliner/scripts/debug_compare.py b/models/rgliner/scripts/debug_compare.py new file mode 100644 index 000000000..342b1a98b --- /dev/null +++ b/models/rgliner/scripts/debug_compare.py @@ -0,0 +1,173 @@ +#!/usr/bin/env python3 +"""Debug script to compare encoder outputs between Python and Rust implementations.""" + +import torch +import numpy as np +from transformers import AutoModel, AutoTokenizer, DebertaV2Model +from gliner import GLiNER + +# Load the model +model = GLiNER.from_pretrained("knowledgator/gliner-relex-multi-v1.0") +tokenizer = model.data_processor.transformer_tokenizer + +# Text and labels +text = "Apple was founded by Steve Jobs in California." +entity_labels = ["person", "organization", "location"] +relation_labels = ["founded by", "located in"] + +# Build the same prompt as our Rust code +# IMPORTANT: Don't add [CLS] - tokenizer adds it automatically +def build_prompt(): + """Build the prompt in the same format as Rust (without [CLS]).""" + parts = [] + for label in entity_labels: + parts.append("<>") + parts.append(label) + parts.append("[SEP]") + for label in relation_labels: + parts.append("<>") + parts.append(label) + parts.append("[SEP]") + parts.append(text) + return " ".join(parts) + +prompt = build_prompt() +print(f"Prompt: {prompt}") + +# Tokenize (tokenizer should add [CLS] automatically) +encoding = tokenizer( + prompt, + return_tensors="pt", + padding=False, + truncation=True, + max_length=512, + add_special_tokens=True, +) +input_ids = encoding["input_ids"] +attention_mask = encoding["attention_mask"] + +print(f"\nToken IDs (first 30): {input_ids[0, :30].tolist()}") +print(f"Token count: {input_ids.shape[1]}") + +# Decode tokens to see what they are +tokens = tokenizer.convert_ids_to_tokens(input_ids[0].tolist()) +print(f"\nFirst 30 tokens: {tokens[:30]}") + +# Find <> token positions +ent_token_id = tokenizer.convert_tokens_to_ids("<>") +print(f"\n<> token ID: {ent_token_id}") + +ent_positions = [] +for i, tok in enumerate(input_ids[0]): + if tok.item() == ent_token_id: + ent_positions.append(i) +print(f"<> positions: {ent_positions}") + +# Find the encoder - try different attribute names +print(f"\nModel type: {type(model)}") +print(f"Model attributes: {[a for a in dir(model) if not a.startswith('_')]}") + +# The encoder is typically model.model in GLiNER +encoder = None +if hasattr(model, 'model'): + encoder = model.model + print(f"Found encoder at model.model: {type(encoder)}") +elif hasattr(model, 'token_rep_layer'): + encoder = model.token_rep_layer + print(f"Found encoder at model.token_rep_layer: {type(encoder)}") + +# Check what encoder contains +if encoder is not None: + print(f"Encoder attributes: {[a for a in dir(encoder) if not a.startswith('_')]}") + + # Try to find the actual DeBERTa model + if hasattr(encoder, 'deberta'): + deberta = encoder.deberta + print(f"Found DeBERTa at encoder.deberta: {type(deberta)}") + elif hasattr(encoder, 'model'): + deberta = encoder.model + print(f"Found DeBERTa at encoder.model: {type(deberta)}") + else: + deberta = encoder + +# Run the encoder +with torch.no_grad(): + # Get the token_rep_layer (DeBERTa encoder) + token_rep_layer = model.model.token_rep_layer + print(f"\ntoken_rep_layer type: {type(token_rep_layer)}") + print(f"token_rep_layer children: {list(token_rep_layer.named_children())}") + + # Call the token_rep_layer + outputs = token_rep_layer(input_ids, attention_mask=attention_mask) + if hasattr(outputs, 'last_hidden_state'): + hidden_states = outputs.last_hidden_state + elif isinstance(outputs, tuple): + hidden_states = outputs[0] + else: + hidden_states = outputs + +print(f"\nEncoder (token_rep_layer) output shape: {hidden_states.shape}") + +# Stats +hs = hidden_states[0].numpy() +print(f"Encoder output stats: mean={hs.mean():.6f}, std={hs.std():.6f}, min={hs.min():.6f}, max={hs.max():.6f}") + +# Print values at <> positions +print(f"\nEncoder output at <> positions (first 5 values):") +for pos in ent_positions: + vals = hs[pos, :5] + print(f" pos {pos}: [{', '.join(f'{v:.4f}' for v in vals)}]") + +# Also check other positions for comparison +print(f"\nEncoder output at other positions:") +for pos in [0, 2, 4, 10, 17]: + if pos < hs.shape[0]: + vals = hs[pos, :5] + print(f" pos {pos}: [{', '.join(f'{v:.4f}' for v in vals)}]") + +# Check prompt_rep_layer +print("\n--- Prompt Rep Layer ---") +ent_embs = hidden_states[0, ent_positions, :] # [n_labels, hidden] +print(f"Entity embeddings shape: {ent_embs.shape}") +print(f"Entity embeddings stats: mean={ent_embs.mean():.6f}, std={ent_embs.std():.6f}") + +# Apply prompt_rep_layer from model.model +prompt_rep = model.model.prompt_rep_layer # This should be a Sequential or MLP +projected = prompt_rep(ent_embs) +print(f"After prompt_rep_layer: shape={projected.shape}") +print(f"After prompt_rep_layer stats: mean={projected.mean():.6f}, std={projected.std():.6f}") + +# Check what prompt_rep_layer consists of +print(f"\nPrompt rep layer structure:") +for name, module in prompt_rep.named_modules(): + if name: + print(f" {name}: {module}") + +# Print projected values for first 5 dims +print(f"\nProjected entity embeddings (first 5 values):") +for i, label in enumerate(entity_labels): + vals = projected[i, :5].detach().numpy() + print(f" {label}: [{', '.join(f'{v:.4f}' for v in vals)}]") + +# Now let's check the raw token embeddings (before any attention) +print("\n--- Raw Token Embeddings (before transformer) ---") +bert_layer = token_rep_layer.bert_layer +deberta_model = bert_layer.model + +# Get raw embeddings +word_embs = deberta_model.embeddings(input_ids) +print(f"Raw embeddings shape: {word_embs.shape}") + +word_embs_np = word_embs[0].detach().numpy() +print(f"Raw embeddings stats: mean={word_embs_np.mean():.6f}, std={word_embs_np.std():.6f}") + +print(f"\nRaw embeddings at <> positions (first 5 values):") +for pos in ent_positions: + vals = word_embs_np[pos, :5] + print(f" pos {pos}: [{', '.join(f'{v:.4f}' for v in vals)}]") + +print(f"\nRaw embeddings at other positions (first 5 values):") +for pos in [0, 2, 4, 10, 17]: + if pos < word_embs_np.shape[0]: + vals = word_embs_np[pos, :5] + print(f" pos {pos}: [{', '.join(f'{v:.4f}' for v in vals)}]") diff --git a/models/rgliner/src/error.rs b/models/rgliner/src/error.rs index 0804efec6..eb004523f 100644 --- a/models/rgliner/src/error.rs +++ b/models/rgliner/src/error.rs @@ -37,6 +37,9 @@ pub enum GlinerError { /// An error that can occur when tokenizing text. #[error("Tokenization error: {0}")] Tokenizer(tokenizers::Error), + /// A tokenization error with a string message. + #[error("Tokenization error: {0}")] + TokenizationError(String), /// An error that can occur with the label encoder. #[error("Label encoder error: {0}")] LabelEncoder(#[from] rbert::BertError), diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs index e4002b833..692dd7cce 100644 --- a/models/rgliner/src/lib.rs +++ b/models/rgliner/src/lib.rs @@ -48,6 +48,28 @@ //! # Ok(()) //! # } //! ``` +//! +//! ## Relation Extraction (GLiNER-RelEx) +//! +//! For joint NER and relation extraction, use the `relex` module: +//! +//! ```rust, no_run +//! use rgliner::relex::*; +//! +//! # async fn example() -> anyhow::Result<()> { +//! let relex = GlinerRelEx::builder() +//! .with_source(GlinerRelExSource::relex_multi()) +//! .build() +//! .await?; +//! +//! let (entities, relations) = relex.extract( +//! "Apple was founded by Steve Jobs.", +//! &["person", "organization"], +//! &["founded by"], +//! ).await?; +//! # Ok(()) +//! # } +//! ``` #![warn(missing_docs)] @@ -55,6 +77,9 @@ mod config; mod decoding; mod error; mod raw; +pub mod relation_decoding; +pub mod relex; +pub mod relex_tokenization; mod source; mod tokenization; diff --git a/models/rgliner/src/raw/bilstm.rs b/models/rgliner/src/raw/bilstm.rs new file mode 100644 index 000000000..f7283956e --- /dev/null +++ b/models/rgliner/src/raw/bilstm.rs @@ -0,0 +1,241 @@ +//! Bidirectional LSTM implementation for GLiNER token representation. +//! +//! The BiLSTM processes encoder output to capture bidirectional context +//! before span representation computation. + +use fusor::{Device, Result, Tensor, VarBuilder}; + +/// Bidirectional LSTM layer for token representation. +/// +/// Processes transformer encoder output through forward and backward LSTMs +/// and concatenates the outputs. +pub struct BiLstm { + // Forward LSTM weights + weight_ih_f: Tensor<2, f32>, // [4*hidden, input_size] + weight_hh_f: Tensor<2, f32>, // [4*hidden, hidden_size] + bias_ih_f: Tensor<1, f32>, // [4*hidden] + bias_hh_f: Tensor<1, f32>, // [4*hidden] + // Backward LSTM weights + weight_ih_b: Tensor<2, f32>, + weight_hh_b: Tensor<2, f32>, + bias_ih_b: Tensor<1, f32>, + bias_hh_b: Tensor<1, f32>, + hidden_size: usize, +} + +impl BiLstm { + /// Load BiLSTM weights from GGUF. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let weight_ih_f: Tensor<2, f32> = vb.get("weight_ih_l0", device)?.dequantize(); + let weight_hh_f: Tensor<2, f32> = vb.get("weight_hh_l0", device)?.dequantize(); + let bias_ih_f: Tensor<1, f32> = vb.get("bias_ih_l0", device)?.dequantize(); + let bias_hh_f: Tensor<1, f32> = vb.get("bias_hh_l0", device)?.dequantize(); + + let weight_ih_b: Tensor<2, f32> = vb.get("weight_ih_l0_reverse", device)?.dequantize(); + let weight_hh_b: Tensor<2, f32> = vb.get("weight_hh_l0_reverse", device)?.dequantize(); + let bias_ih_b: Tensor<1, f32> = vb.get("bias_ih_l0_reverse", device)?.dequantize(); + let bias_hh_b: Tensor<1, f32> = vb.get("bias_hh_l0_reverse", device)?.dequantize(); + + // hidden_size is 4*hidden (for i,f,g,o gates), so actual hidden = shape[0]/4 + let hidden_size = weight_ih_f.shape()[0] / 4; + + #[cfg(debug_assertions)] + { + eprintln!("[DEBUG] BiLstm loaded:"); + eprintln!(" weight_ih_f shape: {:?}", weight_ih_f.shape()); + eprintln!(" weight_hh_f shape: {:?}", weight_hh_f.shape()); + eprintln!(" bias_ih_f shape: {:?}", bias_ih_f.shape()); + eprintln!(" computed hidden_size: {}", hidden_size); + eprintln!(" output_dim: {}", 2 * hidden_size); + } + + Ok(Self { + weight_ih_f, + weight_hh_f, + bias_ih_f, + bias_hh_f, + weight_ih_b, + weight_hh_b, + bias_ih_b, + bias_hh_b, + hidden_size, + }) + } + + /// Forward pass through BiLSTM. + /// + /// # Arguments + /// * `input` - Input tensor [batch, seq_len, input_size] + /// + /// # Returns + /// Output tensor [batch, seq_len, 2*hidden_size] + pub async fn forward(&self, input: &Tensor<3, f32>) -> Tensor<3, f32> { + let [batch_size, seq_len, input_size] = input.shape(); + let device = input.device(); + let output_size = 2 * self.hidden_size; + + // Get all weight data upfront + let input_data = input.clone().as_slice().await.unwrap(); + let w_ih_f = self.weight_ih_f.clone().as_slice().await.unwrap(); + let w_hh_f = self.weight_hh_f.clone().as_slice().await.unwrap(); + let b_ih_f = self.bias_ih_f.clone().as_slice().await.unwrap(); + let b_hh_f = self.bias_hh_f.clone().as_slice().await.unwrap(); + let w_ih_b = self.weight_ih_b.clone().as_slice().await.unwrap(); + let w_hh_b = self.weight_hh_b.clone().as_slice().await.unwrap(); + let b_ih_b = self.bias_ih_b.clone().as_slice().await.unwrap(); + let b_hh_b = self.bias_hh_b.clone().as_slice().await.unwrap(); + + let mut output_data = vec![0.0f32; batch_size * seq_len * output_size]; + + for b in 0..batch_size { + // Forward LSTM + let forward_out = self.lstm_direction( + input_data.as_slice(), + b, + seq_len, + input_size, + w_ih_f.as_slice(), + w_hh_f.as_slice(), + b_ih_f.as_slice(), + b_hh_f.as_slice(), + false, + ); + + // Backward LSTM + let backward_out = self.lstm_direction( + input_data.as_slice(), + b, + seq_len, + input_size, + w_ih_b.as_slice(), + w_hh_b.as_slice(), + b_ih_b.as_slice(), + b_hh_b.as_slice(), + true, + ); + + // Concatenate forward and backward outputs + for t in 0..seq_len { + for i in 0..self.hidden_size { + let out_idx = b * seq_len * output_size + t * output_size; + output_data[out_idx + i] = forward_out[t * self.hidden_size + i]; + output_data[out_idx + self.hidden_size + i] = + backward_out[t * self.hidden_size + i]; + } + } + } + + Tensor::new(&device, &output_data) + .reshape([batch_size, seq_len, output_size]) + .to_concrete() + } + + /// Single direction LSTM pass. + fn lstm_direction( + &self, + input_data: &[f32], + batch_idx: usize, + seq_len: usize, + input_size: usize, + w_ih: &[f32], + w_hh: &[f32], + b_ih: &[f32], + b_hh: &[f32], + reverse: bool, + ) -> Vec { + let hidden_size = self.hidden_size; + let mut h = vec![0.0f32; hidden_size]; + let mut c = vec![0.0f32; hidden_size]; + let mut outputs = vec![0.0f32; seq_len * hidden_size]; + + // Process sequence in order (or reverse) + let indices: Vec = if reverse { + (0..seq_len).rev().collect() + } else { + (0..seq_len).collect() + }; + + for (out_idx, &t) in indices.iter().enumerate() { + // Get input at time t for this batch + let x_start = batch_idx * seq_len * input_size + t * input_size; + let x = &input_data[x_start..x_start + input_size]; + + // Compute gates: i, f, g, o + let mut gates = vec![0.0f32; 4 * hidden_size]; + + for g in 0..(4 * hidden_size) { + let mut sum = b_ih[g] + b_hh[g]; + + // Input contribution: x @ W_ih^T + for i in 0..input_size { + sum += x[i] * w_ih[g * input_size + i]; + } + + // Hidden contribution: h @ W_hh^T + for j in 0..hidden_size { + sum += h[j] * w_hh[g * hidden_size + j]; + } + + gates[g] = sum; + } + + // Apply activations and compute new h, c + for i in 0..hidden_size { + let i_gate = sigmoid(gates[i]); + let f_gate = sigmoid(gates[hidden_size + i]); + let g_gate = tanh(gates[2 * hidden_size + i]); + let o_gate = sigmoid(gates[3 * hidden_size + i]); + + c[i] = f_gate * c[i] + i_gate * g_gate; + h[i] = o_gate * tanh(c[i]); + } + + // Store output in correct position + let store_pos = if reverse { seq_len - 1 - out_idx } else { out_idx }; + for i in 0..hidden_size { + outputs[store_pos * hidden_size + i] = h[i]; + } + } + + outputs + } + + /// Get output dimension (2 * hidden_size for bidirectional). + pub fn output_dim(&self) -> usize { + 2 * self.hidden_size + } + + /// Get the hidden size of a single direction. + pub fn hidden_size(&self) -> usize { + self.hidden_size + } +} + +#[inline] +fn sigmoid(x: f32) -> f32 { + 1.0 / (1.0 + (-x).exp()) +} + +#[inline] +fn tanh(x: f32) -> f32 { + x.tanh() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_sigmoid() { + assert!((sigmoid(0.0) - 0.5).abs() < 1e-6); + assert!(sigmoid(10.0) > 0.99); + assert!(sigmoid(-10.0) < 0.01); + } + + #[test] + fn test_tanh() { + assert!(tanh(0.0).abs() < 1e-6); + assert!(tanh(10.0) > 0.99); + assert!(tanh(-10.0) < -0.99); + } +} diff --git a/models/rgliner/src/raw/joint_scorer.rs b/models/rgliner/src/raw/joint_scorer.rs new file mode 100644 index 000000000..6ff22ac90 --- /dev/null +++ b/models/rgliner/src/raw/joint_scorer.rs @@ -0,0 +1,282 @@ +//! Joint scorer for GLiNER-RelEx. +//! +//! The joint scorer projects token and label embeddings, then uses an MLP +//! to score (token, label) pairs. Outputs 3 classes per pair. + +use fusor::layers::Linear; +use fusor::{Device, Result, Tensor, VarBuilder}; + +/// Joint scorer for token-label pair classification. +/// +/// Architecture (GLiNER token-level scoring): +/// 1. Project label embeddings: proj_label(label_embs) -> [n_labels, proj_dim] +/// 2. Concatenate token embeddings (hidden) with projected labels (proj_dim) +/// 3. MLP: concat(token, proj_label) -> fc1 -> GELU -> fc2 -> scores +/// +/// Note: proj_token exists in weights but the actual forward pass concatenates +/// raw token embeddings with projected labels for the MLP input. +pub struct JointScorer { + #[allow(dead_code)] + proj_token: Linear, // Not used in main scoring path + proj_label: Linear, + out_fc1: Linear, + out_fc2: Linear, +} + +impl JointScorer { + /// Load joint scorer from GGUF. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let proj_token = Linear::load(device, &mut vb.pp("proj_token"))?; + let proj_label = Linear::load(device, &mut vb.pp("proj_label"))?; + let out_fc1 = Linear::load(device, &mut vb.pp("out_mlp.0"))?; + let out_fc2 = Linear::load(device, &mut vb.pp("out_mlp.3"))?; + + #[cfg(debug_assertions)] + { + eprintln!("[DEBUG] JointScorer loaded:"); + eprintln!(" proj_label: in={}, out={}", proj_label.in_features(), proj_label.out_features()); + eprintln!(" out_fc1: in={}, out={}", out_fc1.in_features(), out_fc1.out_features()); + eprintln!(" out_fc2: in={}, out={}", out_fc2.in_features(), out_fc2.out_features()); + // Print fc2 bias values (these are the biases for O, B, I classes) + if let Some(bias) = out_fc2.bias() { + let bias_data = pollster::block_on(bias.clone().as_slice()).unwrap(); + let b = bias_data.as_slice(); + eprintln!(" out_fc2 bias: O={:.6}, B={:.6}, I={:.6}", b[0], b[1], b[2]); + } + } + + Ok(Self { + proj_token, + proj_label, + out_fc1, + out_fc2, + }) + } + + /// Score token-label pairs using bilinear interaction. + /// + /// # Arguments + /// * `token_embs` - Token embeddings [batch, seq_len, hidden_dim] + /// * `label_embs` - Label embeddings [n_labels, hidden_dim] + /// + /// # Returns + /// Scores [batch, seq_len, n_labels, 3] (3 classes: O, B, I) + /// + /// # Architecture + /// The scorer uses bilinear interaction following GLiNER's design: + /// 1. Project both tokens and labels: hidden_dim -> hidden_dim * 2 + /// 2. Split each projection into two halves (first, second) + /// 3. MLP input = concat(token_first, label_first, token_second * label_second) + /// 4. This enables complex token-label interactions through the element-wise product + pub async fn forward( + &self, + token_embs: &Tensor<3, f32>, + label_embs: &Tensor<2, f32>, + ) -> Tensor<4, f32> { + let [batch_size, seq_len, hidden_dim] = token_embs.shape(); + let [n_labels, _] = label_embs.shape(); + + #[cfg(debug_assertions)] + eprintln!("[DEBUG] scorer.forward: batch={}, seq_len={}, hidden_dim={}, n_labels={}", + batch_size, seq_len, hidden_dim, n_labels); + + // Project both token and label embeddings + // token: [batch, seq, hidden] -> [batch, seq, hidden*2] + let proj_tokens = self.proj_token.forward(token_embs); + let [_, _, proj_dim] = proj_tokens.shape(); + let half_proj = proj_dim / 2; + + #[cfg(debug_assertions)] + { + // Verify proj_token computation + let input_data = token_embs.clone().as_slice().await.unwrap(); + let input_slice = input_data.as_slice(); + let output_data = proj_tokens.clone().as_slice().await.unwrap(); + let output_slice = output_data.as_slice(); + eprintln!("[DEBUG] proj_token input[0,0,:5]: {:?}", &input_slice[0..5]); + eprintln!("[DEBUG] proj_token output[0,0,:5]: {:?}", &output_slice[0..5]); + eprintln!("[DEBUG] proj_token output[0,0,768:773]: {:?}", &output_slice[768..773]); + } + + // label: [n_labels, hidden] -> [n_labels, hidden*2] + let label_embs_3d: Tensor<3, f32> = label_embs.unsqueeze(0).to_concrete(); + let proj_labels = self.proj_label.forward(&label_embs_3d); + let proj_labels: Tensor<2, f32> = proj_labels.squeeze(0).to_concrete(); + + #[cfg(debug_assertions)] + eprintln!("[DEBUG] proj_tokens shape: [{}, {}, {}], proj_labels shape: [{}, {}], half_proj={}", + batch_size, seq_len, proj_dim, n_labels, proj_dim, half_proj); + + // Split and combine: token_first + label_first + (token_second * label_second) + // MLP input dimension = half_proj + half_proj + half_proj = 3 * half_proj + let mlp_input_dim = 3 * half_proj; + + #[cfg(debug_assertions)] + eprintln!("[DEBUG] mlp_input_dim={} (3 * {})", mlp_input_dim, half_proj); + + // Get raw data slices (without expansion - we'll handle broadcast manually) + // proj_tokens shape: [batch, seq, proj_dim] + // proj_labels shape: [n_labels, proj_dim] + let tokens_data = proj_tokens.clone().as_slice().await.unwrap(); + let labels_data = proj_labels.clone().as_slice().await.unwrap(); + + let tokens_slice = tokens_data.as_slice(); // [batch * seq * proj_dim] + let labels_slice = labels_data.as_slice(); // [n_labels * proj_dim] + + #[cfg(debug_assertions)] + { + // Check if label projections are different for each label + eprintln!("[DEBUG] Label projection check (first 5 values per label):"); + for l in 0..n_labels { + let start = l * proj_dim; + let vals: Vec = (0..5).map(|i| labels_slice[start + i]).collect(); + eprintln!(" label {}: {:?}", l, vals); + } + + // Check token projections for different tokens + eprintln!("[DEBUG] Token projection check (first 5 tokens, first 5 values):"); + for t in 0..5.min(seq_len) { + let start = t * proj_dim; + let vals: Vec = (0..5).map(|i| tokens_slice[start + i]).collect(); + let vals_second: Vec = (0..5).map(|i| tokens_slice[start + half_proj + i]).collect(); + eprintln!(" token {}: first={:?}, second={:?}", t, vals, vals_second); + } + } + + // Build combined features with manual broadcasting + // Output: [batch, seq, n_labels, mlp_input_dim] + let total_elements = batch_size * seq_len * n_labels; + let mut combined_data = vec![0.0f32; total_elements * mlp_input_dim]; + + for b in 0..batch_size { + for s in 0..seq_len { + for l in 0..n_labels { + // Token features for (b, s): at index (b * seq_len + s) * proj_dim + let tok_base = (b * seq_len + s) * proj_dim; + // Label features for l: at index l * proj_dim + let lab_base = l * proj_dim; + // Output index for (b, s, l) + let out_idx = (b * seq_len * n_labels + s * n_labels + l) * mlp_input_dim; + + // token_first (first half of token projection) + for i in 0..half_proj { + combined_data[out_idx + i] = tokens_slice[tok_base + i]; + } + // label_first (first half of label projection) + for i in 0..half_proj { + combined_data[out_idx + half_proj + i] = labels_slice[lab_base + i]; + } + // element-wise product of second halves + for i in 0..half_proj { + let tok_second = tokens_slice[tok_base + half_proj + i]; + let lab_second = labels_slice[lab_base + half_proj + i]; + combined_data[out_idx + 2 * half_proj + i] = tok_second * lab_second; + } + } + } + } + + let device = token_embs.device(); + let combined: Tensor<3, f32> = Tensor::new(&device, &combined_data) + .reshape([1, total_elements, mlp_input_dim]) + .to_concrete(); + + // Apply MLP: fc1 -> ReLU -> fc2 + let hidden = self.out_fc1.forward(&combined); + let hidden = hidden.relu(); + let output = self.out_fc2.forward(&hidden); + + // Reshape back: [batch, seq, n_labels, 3] + output + .reshape([batch_size, seq_len, n_labels, 3]) + .to_concrete() + } + + /// Score with sigmoid for entity predictions. + /// + /// Returns the 3 per-class sigmoid scores (start, end, inside) for each + /// (token, label) pair, shape [batch, seq_len, n_labels, 3]. + /// + /// The 3 channels are: [start, end, inside] (NOT OBI). + /// Each channel is passed through independent sigmoid. + pub async fn forward_entity_scores( + &self, + token_embs: &Tensor<3, f32>, + label_embs: &Tensor<2, f32>, + ) -> Tensor<4, f32> { + let logits = self.forward(token_embs, label_embs).await; + let [_batch_size, seq_len, n_labels, num_classes] = logits.shape(); + + let logits_data = logits.clone().as_slice().await.unwrap(); + + #[cfg(debug_assertions)] + { + let data = logits_data.as_slice(); + eprintln!("[DEBUG] Raw logits (first 3 tokens, all labels) [start, end, inside]:"); + for s in 0..3.min(seq_len) { + for l in 0..n_labels { + let idx = s * n_labels * num_classes + l * num_classes; + eprintln!(" token {} label {}: start={:.4}, end={:.4}, inside={:.4}", + s, l, data[idx], data[idx+1], data[idx+2]); + } + } + } + + // Apply sigmoid to each value independently (NOT softmax). + let data = logits_data.as_slice(); + let sigmoid_data: Vec = data.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect(); + + let device = logits.device(); + Tensor::new(&device, &sigmoid_data) + .reshape(logits.shape()) + .to_concrete() + } +} + +/// Prompt representation layer for entity/relation labels. +/// +/// Projects label embeddings through a 2-layer FFN. +pub struct PromptRepLayer { + fc1: Linear, + fc2: Linear, +} + +impl PromptRepLayer { + /// Load from GGUF. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let fc1 = Linear::load(device, &mut vb.pp("0"))?; + let fc2 = Linear::load(device, &mut vb.pp("3"))?; + + #[cfg(debug_assertions)] + { + eprintln!("[DEBUG] PromptRepLayer loaded:"); + eprintln!(" fc1: in={}, out={}", fc1.in_features(), fc1.out_features()); + eprintln!(" fc2: in={}, out={}", fc2.in_features(), fc2.out_features()); + } + + Ok(Self { fc1, fc2 }) + } + + /// Project label embeddings. + /// + /// # Arguments + /// * `label_embs` - Label embeddings from encoder [n_labels, hidden] + /// + /// # Returns + /// Projected embeddings [n_labels, hidden] + pub fn forward(&self, label_embs: &Tensor<2, f32>) -> Tensor<2, f32> { + // Wrap as 3D for Linear::forward + let label_3d: Tensor<3, f32> = label_embs.unsqueeze(0).to_concrete(); + let hidden = self.fc1.forward(&label_3d); + let hidden = hidden.relu(); + let output = self.fc2.forward(&hidden); + output.squeeze(0).to_concrete() + } + + /// Forward for 3D tensor [batch, n_labels, hidden]. + pub fn forward_3d(&self, label_embs: &Tensor<3, f32>) -> Tensor<3, f32> { + let hidden = self.fc1.forward(label_embs); + let hidden = hidden.relu(); + self.fc2.forward(&hidden) + } +} diff --git a/models/rgliner/src/raw/mdeberta/attention.rs b/models/rgliner/src/raw/mdeberta/attention.rs new file mode 100644 index 000000000..b698e2d1c --- /dev/null +++ b/models/rgliner/src/raw/mdeberta/attention.rs @@ -0,0 +1,483 @@ +//! mDeBERTa disentangled self-attention. +//! +//! DeBERTa uses disentangled attention with three components: +//! - Content-to-Content (c2c): Standard attention between content vectors +//! - Content-to-Position (c2p): Attention from content to relative positions +//! - Position-to-Content (p2c): Attention from relative positions to content +//! +//! The attention score is: A = c2c + c2p + p2c + +use fusor::layers::{LayerNorm, Linear}; +use fusor::{Device, Result, Tensor, VarBuilder}; + +/// Relative position embeddings for disentangled attention. +pub struct RelativePositionEmbedding { + /// Relative position embedding table [2*max_pos, hidden_size] + embeddings: Tensor<2, f32>, + /// LayerNorm applied to embeddings (norm_rel_ebd = "layer_norm" in DeBERTa) + layer_norm: Option>, + /// Maximum relative positions (e.g., 256) + max_relative_positions: usize, +} + +impl RelativePositionEmbedding { + /// Load with an already-loaded LayerNorm (avoids borrow issues) + pub fn load_with_norm( + device: &Device, + vb: &mut VarBuilder, + layer_norm: Option>, + max_relative_positions: usize, + ) -> Result { + let weight = vb.get("weight", device)?; + let embeddings_raw: Tensor<2, f32> = weight.dequantize(); + + // GGUF stores shape as [hidden_size, 2*max_pos] but we need [2*max_pos, hidden_size] + let [dim0, dim1] = embeddings_raw.shape(); + let embeddings = if dim0 > dim1 { + // Shape is [hidden_size, positions] - need to transpose + #[cfg(debug_assertions)] + eprintln!("[DEBUG] Transposing rel_pos_embd from [{}, {}] to [{}, {}]", dim0, dim1, dim1, dim0); + embeddings_raw.transpose(0, 1).to_concrete() + } else { + embeddings_raw + }; + + #[cfg(debug_assertions)] + eprintln!("[DEBUG] RelativePositionEmbedding loaded: shape={:?}, max_relative_positions={}, has_layer_norm={}", + embeddings.shape(), max_relative_positions, layer_norm.is_some()); + + Ok(Self { + embeddings, + layer_norm, + max_relative_positions, + }) + } + + /// Apply log-bucket position encoding (matches Python make_log_bucket_position). + /// positions close to 0 use linear indexing, far positions are log-bucketed. + fn make_log_bucket_position(rel_pos: i32, bucket_size: i32, max_position: i32) -> i32 { + let sign = if rel_pos > 0 { 1 } else if rel_pos < 0 { -1 } else { 0 }; + let mid = bucket_size / 2; + let abs_pos = if rel_pos < mid && rel_pos > -mid { + mid - 1 + } else { + rel_pos.abs() + }; + if abs_pos <= mid { + rel_pos + } else { + // log_pos = ceil(log(abs_pos/mid) / log((max-1)/mid) * (mid-1)) + mid + let ratio = (abs_pos as f32) / (mid as f32); + let max_ratio = ((max_position - 1) as f32) / (mid as f32); + let log_pos = (ratio.ln() / max_ratio.ln() * ((mid - 1) as f32)).ceil() as i32 + mid; + log_pos * sign + } + } + + /// Compute relative position indices for a sequence. + /// Returns indices [seq_len, seq_len] where each entry is the relative position + /// index into the embedding table. + /// + /// Matches Python: rel_pos_ids = q_ids[:,None] - k_ids[None,:] = i - j + /// Then applies log bucketing with bucket_size=2*max_relative_positions (pos_ebd_size*2), + /// max_position = 2*max_relative_positions... actually: + /// - bucket_size = position_buckets = 256 (pos_ebd_size) + /// - max_position = max_relative_positions = 512 + pub fn compute_relative_indices(&self, seq_len: usize, device: &Device) -> Tensor<2, u32> { + // Python: bucket_size = position_buckets = 256, max_position = max_relative_positions = 512 + // att_span = pos_ebd_size = 256 (= bucket_size) + // The position embedding table has 2*pos_ebd_size = 512 entries + // After bucketing, rel_pos ranges in [-(pos_ebd_size), pos_ebd_size-1] approximately + // c2p_pos = clamp(rel_pos + att_span, 0, 2*att_span-1) -> [0, 2*pos_ebd_size-1] + let bucket_size = self.max_relative_positions as i32; // 256 (pos_ebd_size) + let max_position = 2 * bucket_size; // 512 (2*pos_ebd_size = max_relative_positions) + let att_span = bucket_size; // 256 + let num_positions = (2 * att_span) as i32; // 512 + + let mut indices = vec![0u32; seq_len * seq_len]; + + for i in 0..seq_len { + for j in 0..seq_len { + // Python: rel_pos = q - k = i - j + let rel_pos = i as i32 - j as i32; + // Apply log bucketing + let bucketed = Self::make_log_bucket_position(rel_pos, bucket_size, max_position); + // Shift to positive index: c2p_pos = clamp(bucketed + att_span, 0, 2*att_span-1) + let idx = (bucketed + att_span).clamp(0, num_positions - 1) as u32; + indices[i * seq_len + j] = idx; + } + } + + Tensor::new(device, &indices).reshape([seq_len, seq_len]).to_concrete() + } + + /// Get the raw relative position embedding table (normalized). + /// Returns embeddings [2*max_pos, hidden_size] + pub fn get_embeddings(&self) -> Tensor<2, f32> { + // Apply LayerNorm to embeddings (like HuggingFace get_rel_embedding) + if let Some(ref ln) = self.layer_norm { + // Add batch dimension for LayerNorm: [num_positions, hidden_size] -> [1, num_positions, hidden_size] + let emb_3d: Tensor<3, f32> = self.embeddings.unsqueeze(0).to_concrete(); + let normed = ln.forward(&emb_3d); + // Remove batch dimension + normed.squeeze(0).to_concrete() + } else { + self.embeddings.to_concrete() + } + } + + /// Get relative position embeddings for the given indices (legacy method). + /// Input: indices [seq_len, seq_len] + /// Output: embeddings [seq_len, seq_len, hidden_size] + pub fn forward(&self, indices: &Tensor<2, u32>) -> Tensor<3, f32> { + let [seq_len, _] = indices.shape(); + let [_num_positions, hidden_size] = self.embeddings.shape(); + + // Get normalized embeddings + let normalized_embeddings = self.get_embeddings(); + + // Flatten indices and gather + let flat_indices = indices.reshape([seq_len * seq_len]).to_concrete(); + let gathered = normalized_embeddings.index_select(0, &flat_indices); + + // Reshape back to [seq_len, seq_len, hidden_size] + gathered.reshape([seq_len, seq_len, hidden_size]).to_concrete() + } + + /// Get the maximum relative positions setting. + pub fn max_relative_positions(&self) -> usize { + self.max_relative_positions + } +} + +/// mDeBERTa disentangled self-attention with shared key attention (share_att_key=True). +pub struct MDebertaAttention { + query: Linear, + key: Linear, + value: Linear, + output: Linear, + num_heads: usize, + head_dim: usize, + scale: f32, +} + +impl MDebertaAttention { + pub fn load( + device: &Device, + vb: &mut VarBuilder, + num_heads: usize, + head_dim: usize, + ) -> Result { + let query = Linear::load(device, &mut vb.pp("query"))?; + let key = Linear::load(device, &mut vb.pp("key"))?; + let value = Linear::load(device, &mut vb.pp("value"))?; + let output = Linear::load(device, &mut vb.pp("output"))?; + + // Scale factor for disentangled attention (3 components: c2c, c2p, p2c) + // Python: scale = scaled_size_sqrt(query_layer, scale_factor) where scale_factor=3 + // This means sqrt(head_dim * 3) + let scale = 1.0 / ((head_dim as f32) * 3.0).sqrt(); + + Ok(Self { + query, + key, + value, + output, + num_heads, + head_dim, + scale, + }) + } + + /// Forward pass with disentangled attention. + /// + /// # Arguments + /// * `hidden_states` - Input [batch, seq_len, hidden_size] + /// * `rel_pos_emb` - Relative position embedding table [2*max_pos, hidden_size] + /// * `rel_pos_indices` - Relative position indices [seq_len, seq_len] + /// * `attention_mask` - Optional attention mask [batch, seq_len] + pub fn forward_with_indices( + &self, + hidden_states: &Tensor<3, f32>, + rel_pos_emb: &Tensor<2, f32>, + rel_pos_indices: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + let [b_sz, seq_len, _] = hidden_states.shape(); + let hidden_size = self.num_heads * self.head_dim; + let [num_positions, _] = rel_pos_emb.shape(); + + // Compute Q, K, V projections for content + let query = self.query.forward(hidden_states); + let key = self.key.forward(hidden_states); + let value = self.value.forward(hidden_states); + + // Reshape to [batch, num_heads, seq_len, head_dim] + let query = query + .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + let key = key + .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + let value = value + .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + + // === Content-to-Content attention === + // c2c = Q @ K^T + let c2c_scores = query.mat_mul(&key.transpose(2, 3)); + + // === Position attention with shared Q/K projections === + // Project position embeddings using the same Q and K projections + // rel_pos_emb: [2*max_pos, hidden_size] -> [1, 2*max_pos, hidden_size] + let rel_emb_3d: Tensor<3, f32> = rel_pos_emb.unsqueeze(0).to_concrete(); + + // pos_query = query_proj(rel_emb): [1, 2*max_pos, hidden] -> [1, heads, 2*max_pos, head_dim] + let pos_query = self.query.forward(&rel_emb_3d); + let pos_query = pos_query + .reshape([1, num_positions, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + + // pos_key = key_proj(rel_emb): [1, 2*max_pos, hidden] -> [1, heads, 2*max_pos, head_dim] + let pos_key = self.key.forward(&rel_emb_3d); + let pos_key = pos_key + .reshape([1, num_positions, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + + // === Content-to-Position attention === + // c2p = Q @ pos_key^T -> [batch, heads, seq, 2*max_pos] + // Then gather based on relative positions + let c2p_all = query.mat_mul(&pos_key.transpose(2, 3)); + let c2p_scores = self.gather_c2p(&c2p_all, rel_pos_indices); + + // === Position-to-Content attention === + // p2c = K @ pos_query^T -> [batch, heads, seq, 2*max_pos] + // Then gather based on transposed relative positions + let p2c_all = key.mat_mul(&pos_query.transpose(2, 3)); + let p2c_scores = self.gather_p2c(&p2c_all, rel_pos_indices); + + // Combine: attention = (c2c + c2p + p2c) * scale + let attn_scores = c2c_scores + .add_(&c2p_scores) + .add_(&p2c_scores) + .mul_scalar(self.scale); + + // Apply attention mask + let attn_scores = if let Some(mask) = attention_mask { + const MASK_NEG_VALUE: f32 = -10000.0; + let mask_f32: Tensor<2, f32> = mask.cast(); + let zeros = mask_f32.zeros_like(); + let ones = (zeros + 1.0f32).to_concrete(); + let mask_bias = ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete(); + // Broadcast mask to [batch, 1, 1, seq_len] + let mask_bias_3d: Tensor<3, f32> = mask_bias.unsqueeze(1).to_concrete(); + let mask_bias_4d: Tensor<4, f32> = mask_bias_3d.unsqueeze(1).to_concrete(); + attn_scores.add_(&mask_bias_4d) + } else { + attn_scores + }; + + // Softmax + let attn_probs = attn_scores.softmax_last_dim::<3>(); + + // Apply attention to values + let context = attn_probs.mat_mul(&value); + + // Reshape back to [batch, seq_len, hidden_size] + let context = context + .transpose(1, 2) + .to_concrete() + .reshape([b_sz, seq_len, hidden_size]) + .to_concrete(); + + // Output projection + self.output.forward(&context) + } + + /// Gather c2p attention scores based on relative position indices. + /// + /// Input: c2p_all [batch, heads, seq_len, 2*max_pos] - scores to all positions + /// rel_pos_indices: [seq_len, seq_len] - index into position embeddings + /// + /// Output: [batch, heads, seq_len, seq_len] - gathered scores + fn gather_c2p( + &self, + c2p_all: &Tensor<4, f32>, + rel_pos_indices: &Tensor<2, u32>, + ) -> Tensor<4, f32> { + let [b_sz, num_heads, seq_len, _num_pos] = c2p_all.shape(); + let device = c2p_all.device(); + + // Get data slices + let c2p_data = pollster::block_on(c2p_all.clone().as_slice()).unwrap(); + let indices_data = pollster::block_on(rel_pos_indices.clone().as_slice()).unwrap(); + let c2p = c2p_data.as_slice(); + let indices = indices_data.as_slice(); + let num_pos = _num_pos; + + let mut gathered = vec![0.0f32; b_sz * num_heads * seq_len * seq_len]; + + for b in 0..b_sz { + for h in 0..num_heads { + for i in 0..seq_len { + for j in 0..seq_len { + // Index into c2p_all: [b, h, i, rel_pos[i,j]] + let rel_idx = indices[i * seq_len + j] as usize; + let c2p_idx = b * num_heads * seq_len * num_pos + + h * seq_len * num_pos + + i * num_pos + + rel_idx; + let out_idx = b * num_heads * seq_len * seq_len + + h * seq_len * seq_len + + i * seq_len + + j; + gathered[out_idx] = c2p[c2p_idx]; + } + } + } + } + + Tensor::new(&device, &gathered) + .reshape([b_sz, num_heads, seq_len, seq_len]) + .to_concrete() + } + + /// Gather p2c attention scores. + /// + /// Python derivation: + /// - r_pos = relative_pos (since seq_q == seq_k) + /// - p2c_pos[i,j] = clamp(-r_pos[i,j] + att_span, 0, 2*att_span-1) + /// = clamp(-(i-j) + att_span) = clamp((j-i) + att_span) + /// - gather_out[b, m, n] = p2c_att[b, m, p2c_pos[m, n]] + /// - final[b, i, j] = gather_out[b, j, i] (after transpose) + /// = p2c_att[b, j, p2c_pos[j, i]] + /// = p2c_att[b, j, clamp(i - j + att_span)] + /// = p2c_all[b, j, indices[i, j]] (using indices[i,j] = bucketed(i-j) + att_span) + fn gather_p2c( + &self, + p2c_all: &Tensor<4, f32>, + rel_pos_indices: &Tensor<2, u32>, + ) -> Tensor<4, f32> { + let [b_sz, num_heads, seq_len, num_pos] = p2c_all.shape(); + let device = p2c_all.device(); + + let p2c_data = pollster::block_on(p2c_all.clone().as_slice()).unwrap(); + let indices_data = pollster::block_on(rel_pos_indices.clone().as_slice()).unwrap(); + let p2c = p2c_data.as_slice(); + let indices = indices_data.as_slice(); + + let mut gathered = vec![0.0f32; b_sz * num_heads * seq_len * seq_len]; + + for b in 0..b_sz { + for h in 0..num_heads { + for i in 0..seq_len { + for j in 0..seq_len { + // final[b, i, j] = p2c_all[b, j, indices[i, j]] + let rel_idx = (indices[i * seq_len + j] as usize).min(num_pos - 1); + let p2c_idx = b * num_heads * seq_len * num_pos + + h * seq_len * num_pos + + j * num_pos // key dim = j + + rel_idx; + let out_idx = b * num_heads * seq_len * seq_len + + h * seq_len * seq_len + + i * seq_len + + j; + gathered[out_idx] = p2c[p2c_idx]; + } + } + } + } + + Tensor::new(&device, &gathered) + .reshape([b_sz, num_heads, seq_len, seq_len]) + .to_concrete() + } + + /// Legacy forward pass (for compatibility). + pub fn forward( + &self, + hidden_states: &Tensor<3, f32>, + rel_pos_emb: Option<&Tensor<3, f32>>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + // This method is kept for backward compatibility but shouldn't be used + // with the new architecture + if rel_pos_emb.is_some() { + panic!("Use forward_with_indices for proper position attention"); + } + + let [b_sz, seq_len, _] = hidden_states.shape(); + let hidden_size = self.num_heads * self.head_dim; + + let query = self.query.forward(hidden_states); + let key = self.key.forward(hidden_states); + let value = self.value.forward(hidden_states); + + let query = query.reshape([b_sz, seq_len, self.num_heads, self.head_dim]).transpose(1, 2).to_concrete(); + let key = key.reshape([b_sz, seq_len, self.num_heads, self.head_dim]).transpose(1, 2).to_concrete(); + let value = value.reshape([b_sz, seq_len, self.num_heads, self.head_dim]).transpose(1, 2).to_concrete(); + + let c2c_scores = query.mat_mul(&key.transpose(2, 3)); + let attn_scores = c2c_scores.mul_scalar(1.0 / (self.head_dim as f32).sqrt()); + + let attn_scores = if let Some(mask) = attention_mask { + const MASK_NEG_VALUE: f32 = -10000.0; + let mask_f32: Tensor<2, f32> = mask.cast(); + let zeros = mask_f32.zeros_like(); + let ones = (zeros + 1.0f32).to_concrete(); + let mask_bias = ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete(); + let mask_bias_3d: Tensor<3, f32> = mask_bias.unsqueeze(1).to_concrete(); + let mask_bias_4d: Tensor<4, f32> = mask_bias_3d.unsqueeze(1).to_concrete(); + attn_scores.add_(&mask_bias_4d) + } else { + attn_scores + }; + + let attn_probs = attn_scores.softmax_last_dim::<3>(); + let context = attn_probs.mat_mul(&value); + let context = context.transpose(1, 2).to_concrete().reshape([b_sz, seq_len, hidden_size]).to_concrete(); + self.output.forward(&context) + } +} + +/// Shared relative position embedding layer (used across all layers in DeBERTa). +pub struct DisentangledSelfAttention { + attention: MDebertaAttention, +} + +impl DisentangledSelfAttention { + pub fn load( + device: &Device, + vb: &mut VarBuilder, + num_heads: usize, + head_dim: usize, + ) -> Result { + let attention = MDebertaAttention::load(device, vb, num_heads, head_dim)?; + Ok(Self { attention }) + } + + /// Forward with relative position indices and embedding table. + pub fn forward_with_rel( + &self, + hidden_states: &Tensor<3, f32>, + rel_pos_emb: &Tensor<2, f32>, + rel_pos_indices: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + self.attention.forward_with_indices(hidden_states, rel_pos_emb, rel_pos_indices, attention_mask) + } + + pub fn forward( + &self, + hidden_states: &Tensor<3, f32>, + rel_pos_emb: Option<&Tensor<3, f32>>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + self.attention.forward(hidden_states, rel_pos_emb, attention_mask) + } +} diff --git a/models/rgliner/src/raw/mdeberta/config.rs b/models/rgliner/src/raw/mdeberta/config.rs new file mode 100644 index 000000000..86578db39 --- /dev/null +++ b/models/rgliner/src/raw/mdeberta/config.rs @@ -0,0 +1,138 @@ +//! mDeBERTa-v3 configuration from GGUF metadata. + +use fusor::{Result, VarBuilder}; + +/// Configuration for mDeBERTa-v3 loaded from GGUF metadata. +#[derive(Debug, Clone)] +pub struct MDebertaConfig { + /// Number of attention heads. + pub num_heads: usize, + /// Number of transformer layers. + pub num_layers: usize, + /// Hidden size (embedding dimension). + pub hidden_size: usize, + /// Dimension per attention head. + pub head_dimension: usize, + /// Intermediate size for FFN. + pub intermediate_size: usize, + /// Maximum context length. + pub context_length: usize, + /// Maximum relative position distance for attention. + pub max_relative_positions: usize, + /// LayerNorm epsilon. + pub norm_eps: f32, + /// Vocabulary size. + pub vocab_size: usize, + /// Position buckets for relative position encoding. + pub position_buckets: usize, + /// Whether to share attention weights across layers. + pub share_att_key: bool, +} + +impl MDebertaConfig { + /// Load configuration from GGUF metadata. + /// + /// Note: GGUF metadata keys use "gliner." prefix regardless of VarBuilder scope, + /// since metadata is stored globally (not per-tensor). + pub fn from_gguf(vb: &VarBuilder) -> Result { + // Metadata keys use "gliner." prefix (not the tensor prefix) + let num_heads = vb + .get_metadata("gliner.attention.head_count") + .and_then(|v| v.to_u32().ok()) + .ok_or_else(|| { + fusor::Error::msg("Missing required GGUF metadata: gliner.attention.head_count") + })? as usize; + + let num_layers = vb + .get_metadata("gliner.block_count") + .and_then(|v| v.to_u32().ok()) + .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: gliner.block_count"))? + as usize; + + let hidden_size = vb + .get_metadata("gliner.embedding_length") + .and_then(|v| v.to_u32().ok()) + .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: gliner.embedding_length"))? + as usize; + + if hidden_size % num_heads != 0 { + return Err(fusor::Error::msg(format!( + "hidden_size ({hidden_size}) must be divisible by num_heads ({num_heads})" + ))); + } + + let head_dimension = vb + .get_metadata("gliner.attention.key_length") + .and_then(|v| v.to_u32().ok()) + .map(|x| x as usize) + .unwrap_or_else(|| hidden_size / num_heads); + + let intermediate_size = vb + .get_metadata("gliner.feed_forward_length") + .and_then(|v| v.to_u32().ok()) + .unwrap_or((hidden_size * 4) as u32) as usize; + + let context_length = vb + .get_metadata("gliner.context_length") + .and_then(|v| v.to_u32().ok()) + .unwrap_or(512) as usize; + + // DeBERTa-specific: maximum relative position distance + let max_relative_positions = vb + .get_metadata("gliner.attention.max_relative_positions") + .and_then(|v| v.to_u32().ok()) + .unwrap_or(512) as usize; + + let norm_eps = vb + .get_metadata("gliner.attention.layer_norm_epsilon") + .and_then(|v| v.to_f32().ok()) + .unwrap_or(1e-7); + + let vocab_size = vb + .get_metadata("gliner.vocab_size") + .and_then(|v| v.to_u32().ok()) + .unwrap_or(250105) as usize; + + // DeBERTa-v3 specific: position buckets for relative position encoding + let position_buckets = vb + .get_metadata("gliner.attention.position_buckets") + .and_then(|v| v.to_u32().ok()) + .unwrap_or(256) as usize; + + let share_att_key = vb + .get_metadata("gliner.attention.share_att_key") + .and_then(|v| v.to_bool().ok()) + .unwrap_or(true); + + Ok(Self { + num_heads, + num_layers, + hidden_size, + head_dimension, + intermediate_size, + context_length, + max_relative_positions, + norm_eps, + vocab_size, + position_buckets, + share_att_key, + }) + } + + /// Create a default config for mDeBERTa-v3-base. + pub fn mdeberta_v3_base() -> Self { + Self { + num_heads: 12, + num_layers: 12, + hidden_size: 768, + head_dimension: 64, + intermediate_size: 3072, + context_length: 512, + max_relative_positions: 512, + norm_eps: 1e-7, + vocab_size: 250105, + position_buckets: 256, + share_att_key: true, + } + } +} diff --git a/models/rgliner/src/raw/mdeberta/feed_forward.rs b/models/rgliner/src/raw/mdeberta/feed_forward.rs new file mode 100644 index 000000000..d9f39706c --- /dev/null +++ b/models/rgliner/src/raw/mdeberta/feed_forward.rs @@ -0,0 +1,31 @@ +//! mDeBERTa Feed Forward Network. + +use fusor::layers::Linear; +use fusor::{Device, Result, Tensor, VarBuilder}; + +/// Standard GELU Feed Forward Network for mDeBERTa. +/// +/// Formula: FFN(x) = GELU(x @ W1 + b1) @ W2 + b2 +pub struct MDebertaFeedForward { + intermediate: Linear, + output: Linear, +} + +impl MDebertaFeedForward { + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let intermediate = Linear::load(device, &mut vb.pp("intermediate"))?; + let output = Linear::load(device, &mut vb.pp("output"))?; + + Ok(Self { + intermediate, + output, + }) + } + + pub fn forward(&self, x: &Tensor<3, f32>) -> Tensor<3, f32> { + // Intermediate: x @ W1 + b1, then GELU + let hidden = self.intermediate.forward(x).gelu(); + // Output: hidden @ W2 + b2 + self.output.forward(&hidden) + } +} diff --git a/models/rgliner/src/raw/mdeberta/layer.rs b/models/rgliner/src/raw/mdeberta/layer.rs new file mode 100644 index 000000000..554bb934e --- /dev/null +++ b/models/rgliner/src/raw/mdeberta/layer.rs @@ -0,0 +1,87 @@ +//! mDeBERTa transformer layer. + +use fusor::layers::LayerNorm; +use fusor::{Device, Result, Tensor, VarBuilder}; + +use super::attention::DisentangledSelfAttention; +use super::feed_forward::MDebertaFeedForward; + +/// A single mDeBERTa transformer layer. +/// +/// Architecture: +/// 1. Self-attention with disentangled attention +/// 2. Add & LayerNorm +/// 3. Feed-forward network +/// 4. Add & LayerNorm +pub struct MDebertaLayer { + attention: DisentangledSelfAttention, + attention_norm: LayerNorm<1, f32>, + feed_forward: MDebertaFeedForward, + output_norm: LayerNorm<1, f32>, +} + +impl MDebertaLayer { + pub fn load( + device: &Device, + vb: &mut VarBuilder, + num_heads: usize, + head_dim: usize, + eps: f32, + ) -> Result { + let attention = DisentangledSelfAttention::load( + device, + &mut vb.pp("attention"), + num_heads, + head_dim, + )?; + let attention_norm = LayerNorm::load(device, &mut vb.pp("attention_norm"), eps)?; + let feed_forward = MDebertaFeedForward::load(device, &mut vb.pp("ffn"))?; + let output_norm = LayerNorm::load(device, &mut vb.pp("output_norm"), eps)?; + + Ok(Self { + attention, + attention_norm, + feed_forward, + output_norm, + }) + } + + /// Forward pass through the layer with proper position attention. + /// + /// # Arguments + /// * `hidden_states` - Input [batch, seq_len, hidden_size] + /// * `rel_pos_emb` - Relative position embedding table [2*max_pos, hidden_size] + /// * `rel_pos_indices` - Relative position indices [seq_len, seq_len] + /// * `attention_mask` - Optional attention mask [batch, seq_len] + pub fn forward_with_rel( + &self, + hidden_states: &Tensor<3, f32>, + rel_pos_emb: &Tensor<2, f32>, + rel_pos_indices: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + // Self-attention + residual + norm + let attn_output = self.attention.forward_with_rel(hidden_states, rel_pos_emb, rel_pos_indices, attention_mask); + let hidden_states = self.attention_norm.forward(&hidden_states.add_(&attn_output)); + + // FFN + residual + norm + let ffn_output = self.feed_forward.forward(&hidden_states); + self.output_norm.forward(&hidden_states.add_(&ffn_output)) + } + + /// Legacy forward pass (for compatibility). + pub fn forward( + &self, + hidden_states: &Tensor<3, f32>, + rel_pos_emb: Option<&Tensor<3, f32>>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + // Self-attention + residual + norm + let attn_output = self.attention.forward(hidden_states, rel_pos_emb, attention_mask); + let hidden_states = self.attention_norm.forward(&hidden_states.add_(&attn_output)); + + // FFN + residual + norm + let ffn_output = self.feed_forward.forward(&hidden_states); + self.output_norm.forward(&hidden_states.add_(&ffn_output)) + } +} diff --git a/models/rgliner/src/raw/mdeberta/mod.rs b/models/rgliner/src/raw/mdeberta/mod.rs new file mode 100644 index 000000000..18ee18d02 --- /dev/null +++ b/models/rgliner/src/raw/mdeberta/mod.rs @@ -0,0 +1,12 @@ +//! mDeBERTa-v3 encoder for GLiNER-RelEx. +//! +//! mDeBERTa uses disentangled attention with relative position embeddings, +//! which differs from ModernBERT's RoPE-based attention. + +mod attention; +mod config; +mod feed_forward; +mod layer; +mod model; + +pub use model::MDebertaModel; diff --git a/models/rgliner/src/raw/mdeberta/model.rs b/models/rgliner/src/raw/mdeberta/model.rs new file mode 100644 index 000000000..d9a3dd867 --- /dev/null +++ b/models/rgliner/src/raw/mdeberta/model.rs @@ -0,0 +1,181 @@ +//! mDeBERTa-v3 encoder model. + +#[cfg(debug_assertions)] +use pollster; + +use fusor::layers::{Embedding, LayerNorm}; +use fusor::{Device, Result, Tensor, VarBuilder}; + +use super::attention::RelativePositionEmbedding; +use super::config::MDebertaConfig; +use super::layer::MDebertaLayer; + +/// mDeBERTa-v3 encoder model for GLiNER-RelEx. +/// +/// This is a bidirectional transformer encoder using disentangled attention +/// with relative position embeddings. +pub struct MDebertaModel { + /// Token embeddings + token_embeddings: Embedding, + /// Embedding LayerNorm + embedding_norm: LayerNorm<1, f32>, + /// Relative position embeddings (shared across layers, with LayerNorm) + rel_pos_embedding: RelativePositionEmbedding, + /// Transformer layers + layers: Vec, + /// Device + device: Device, + /// Configuration + config: MDebertaConfig, +} + +impl MDebertaModel { + /// Load mDeBERTa from GGUF weights. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let config = MDebertaConfig::from_gguf(vb)?; + + // Load token embeddings + let token_embeddings = Embedding::load(device, &mut vb.pp("token_embd"))?; + + // Load embedding LayerNorm + let embedding_norm = LayerNorm::load(device, &mut vb.pp("embd_norm"), config.norm_eps)?; + + // Load relative position embeddings with LayerNorm + // The output_norm in GGUF is the LayerNorm for relative position embeddings + // Load norm first to avoid borrow issues + let rel_pos_norm = LayerNorm::load(device, &mut vb.pp("output_norm"), config.norm_eps).ok(); + let rel_pos_embedding = RelativePositionEmbedding::load_with_norm( + device, + &mut vb.pp("rel_pos_embd"), + rel_pos_norm, + config.max_relative_positions, + )?; + + // Load transformer layers + let mut layers = Vec::with_capacity(config.num_layers); + for i in 0..config.num_layers { + let layer = MDebertaLayer::load( + device, + &mut vb.pp(format!("blk.{i}")), + config.num_heads, + config.head_dimension, + config.norm_eps, + )?; + layers.push(layer); + } + + Ok(Self { + token_embeddings, + embedding_norm, + rel_pos_embedding, + layers, + device: device.clone(), + config, + }) + } + + /// Forward pass through the model. + /// + /// # Arguments + /// * `input_ids` - Token IDs [batch, seq_len] + /// * `attention_mask` - Optional attention mask [batch, seq_len] + /// + /// # Returns + /// Hidden states [batch, seq_len, hidden_size] + pub fn forward( + &self, + input_ids: &Tensor<2, u32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + let [_batch_size, seq_len] = input_ids.shape(); + + // Get token embeddings + let mut hidden_states = self.token_embeddings.forward(input_ids); + + #[cfg(debug_assertions)] + { + let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); + let slice = data.as_slice(); + let mean: f32 = slice.iter().sum::() / slice.len() as f32; + let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + eprintln!("[DEBUG] After token_embeddings: mean={:.6}, std={:.6}", mean, std); + } + + // Apply embedding LayerNorm + hidden_states = self.embedding_norm.forward(&hidden_states); + + #[cfg(debug_assertions)] + { + let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); + let slice = data.as_slice(); + let hidden_size = self.config.hidden_size; + let mean: f32 = slice.iter().sum::() / slice.len() as f32; + let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + eprintln!("[DEBUG] After embedding_norm: mean={:.6}, std={:.6}", mean, std); + // Print raw embeddings at <> positions (1, 3, 5) and others + eprintln!("[DEBUG] Raw embeddings at positions (first 5 values):"); + for pos in [0, 1, 2, 3, 4, 5, 10, 17] { + if pos < seq_len { + let start = pos * hidden_size; + let vals: Vec = (0..5).map(|i| slice[start + i]).collect(); + eprintln!(" pos {}: {:?}", pos, vals); + } + } + } + + // Compute relative position indices and get embedding table + let rel_indices = self.rel_pos_embedding.compute_relative_indices(seq_len, &self.device); + let rel_pos_emb = self.rel_pos_embedding.get_embeddings(); + + // Pass through transformer layers with proper position attention + for (i, layer) in self.layers.iter().enumerate() { + hidden_states = layer.forward_with_rel(&hidden_states, &rel_pos_emb, &rel_indices, attention_mask); + + #[cfg(debug_assertions)] + if i == 0 || i == 11 { + let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); + let slice = data.as_slice(); + let mean: f32 = slice.iter().sum::() / slice.len() as f32; + let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + eprintln!("[DEBUG] After layer {}: mean={:.6}, std={:.6}", i, mean, std); + } + } + + #[cfg(debug_assertions)] + { + let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); + let slice = data.as_slice(); + let mean: f32 = slice.iter().sum::() / slice.len() as f32; + let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + eprintln!("[DEBUG] Encoder output: mean={:.6}, std={:.6}", mean, std); + } + + // Return last layer output directly (no final LayerNorm on hidden states) + hidden_states + } + + /// Get the embedding dimension. + pub fn embedding_dim(&self) -> usize { + self.config.hidden_size + } + + /// Get the maximum sequence length. + pub fn max_seq_len(&self) -> usize { + self.config.context_length + } + + /// Get the vocabulary size. + pub fn vocab_size(&self) -> usize { + self.config.vocab_size + } + + /// Get the device. + pub fn device(&self) -> &Device { + &self.device + } + + /// Get the configuration. + pub fn config(&self) -> &MDebertaConfig { + &self.config + } +} diff --git a/models/rgliner/src/raw/mod.rs b/models/rgliner/src/raw/mod.rs index abd4617e2..9804d6b5f 100644 --- a/models/rgliner/src/raw/mod.rs +++ b/models/rgliner/src/raw/mod.rs @@ -1,13 +1,22 @@ //! Raw model implementations for GLiNER. +pub mod mdeberta; pub mod modern_bert; +mod bilstm; +mod joint_scorer; mod label_encoder; +mod pair_projector; +mod relations_layer; mod scorer; mod span_layer; mod text_encoder; +pub use bilstm::BiLstm; +pub use joint_scorer::{JointScorer, PromptRepLayer}; pub use label_encoder::{CachedLabels, LabelEncoder}; +pub use pair_projector::PairProjector; +pub use relations_layer::RelationsRepLayer; pub use scorer::Scorer; pub use span_layer::SpanLayer; pub use text_encoder::TextEncoder; diff --git a/models/rgliner/src/raw/pair_projector.rs b/models/rgliner/src/raw/pair_projector.rs new file mode 100644 index 000000000..a43dee7dd --- /dev/null +++ b/models/rgliner/src/raw/pair_projector.rs @@ -0,0 +1,141 @@ +//! Entity pair projector for relation classification. +//! +//! Projects concatenated head and tail entity embeddings to a space +//! suitable for scoring against relation label embeddings. + +use fusor::layers::Linear; +use fusor::{Device, Result, Tensor, VarBuilder}; + +/// Entity pair projector. +/// +/// Architecture: Linear(hidden*2 -> hidden) -> ReLU -> Dropout -> Linear(hidden -> hidden) +/// +/// Takes concatenated head and tail entity embeddings and produces a pair representation +/// that can be scored against relation label embeddings. +pub struct PairProjector { + linear1: Linear, + linear2: Linear, +} + +impl PairProjector { + /// Load the pair projector from GGUF weights. + /// + /// The GGUF weights use numeric indices (0, 3) for the layers + /// corresponding to the PyTorch Sequential layer indices. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let linear1 = Linear::load(device, &mut vb.pp("0"))?; + let linear2 = Linear::load(device, &mut vb.pp("3"))?; + + Ok(Self { linear1, linear2 }) + } + + /// Project entity pairs to relation space. + /// + /// # Arguments + /// * `head_embeddings` - Head entity embeddings [num_pairs, hidden_size] + /// * `tail_embeddings` - Tail entity embeddings [num_pairs, hidden_size] + /// + /// # Returns + /// Pair representations [num_pairs, hidden_size] + pub fn forward( + &self, + head_embeddings: &Tensor<2, f32>, + tail_embeddings: &Tensor<2, f32>, + ) -> Tensor<2, f32> { + let [_num_pairs, _hidden_size] = head_embeddings.shape(); + + // Expand to 3D for cat operation, then squeeze back + let head_3d: Tensor<3, f32> = head_embeddings.unsqueeze(0).to_concrete(); + let tail_3d: Tensor<3, f32> = tail_embeddings.unsqueeze(0).to_concrete(); + + // Concatenate head and tail: [1, num_pairs, hidden_size * 2] + let concatenated = Tensor::cat([head_3d, tail_3d], 2); + + // First layer: Linear -> ReLU + let hidden = self.linear1.forward(&concatenated).relu(); + + // Second layer: Linear -> squeeze back to 2D + let result = self.linear2.forward(&hidden); + result.squeeze(0).to_concrete() + } + + /// Project entity pairs for batched processing. + /// + /// # Arguments + /// * `head_embeddings` - Head entity embeddings [batch, num_pairs, hidden_size] + /// * `tail_embeddings` - Tail entity embeddings [batch, num_pairs, hidden_size] + /// + /// # Returns + /// Pair representations [batch, num_pairs, hidden_size] + pub fn forward_batched( + &self, + head_embeddings: &Tensor<3, f32>, + tail_embeddings: &Tensor<3, f32>, + ) -> Tensor<3, f32> { + let [_batch_size, _num_pairs, _hidden_size] = head_embeddings.shape(); + + // Concatenate head and tail: [batch, num_pairs, hidden_size * 2] + let concatenated = Tensor::cat( + [head_embeddings.to_concrete(), tail_embeddings.to_concrete()], + 2, + ); + + // First layer: Linear -> ReLU + let hidden = self.linear1.forward(&concatenated).relu(); + + // Second layer: Linear + self.linear2.forward(&hidden) + } +} + +/// Scorer for relation classification. +/// +/// Computes scores between pair representations and relation label embeddings. +pub struct RelationScorer; + +impl RelationScorer { + /// Score pairs against relation labels. + /// + /// # Arguments + /// * `pair_embeddings` - Pair representations [num_pairs, hidden_size] + /// * `relation_embeddings` - Relation label embeddings [num_relations, hidden_size] + /// + /// # Returns + /// Scores [num_pairs, num_relations] (logits, apply sigmoid for probabilities) + pub fn forward( + pair_embeddings: &Tensor<2, f32>, + relation_embeddings: &Tensor<2, f32>, + ) -> Tensor<2, f32> { + // Dot product: pairs @ relations.T + let rel_t = relation_embeddings.transpose(0, 1); + pair_embeddings.mat_mul(&rel_t) + } + + /// Score pairs against relation labels (batched). + /// + /// # Arguments + /// * `pair_embeddings` - Pair representations [batch, num_pairs, hidden_size] + /// * `relation_embeddings` - Relation label embeddings [num_relations, hidden_size] + /// + /// # Returns + /// Scores [batch, num_pairs, num_relations] + pub fn forward_batched( + pair_embeddings: &Tensor<3, f32>, + relation_embeddings: &Tensor<2, f32>, + ) -> Tensor<3, f32> { + let [batch_size, num_pairs, hidden_size] = pair_embeddings.shape(); + let [num_relations, _] = relation_embeddings.shape(); + + // Flatten pairs: [batch * num_pairs, hidden_size] + let flat_pairs = pair_embeddings + .reshape([batch_size * num_pairs, hidden_size]) + .to_concrete(); + + // Dot product: [batch * num_pairs, hidden_size] @ [hidden_size, num_relations] + let rel_t = relation_embeddings.transpose(0, 1); + let scores = flat_pairs.mat_mul(&rel_t); + + // Reshape back: [batch, num_pairs, num_relations] + scores.reshape([batch_size, num_pairs, num_relations]).to_concrete() + } +} diff --git a/models/rgliner/src/raw/relations_layer.rs b/models/rgliner/src/raw/relations_layer.rs new file mode 100644 index 000000000..3589e9fe6 --- /dev/null +++ b/models/rgliner/src/raw/relations_layer.rs @@ -0,0 +1,120 @@ +//! Relation representation layer for adjacency matrix computation. +//! +//! Computes an adjacency matrix between entity spans to filter +//! candidate pairs for relation classification. + +use fusor::layers::Linear; +use fusor::{Device, Result, Tensor, VarBuilder}; + +/// Relation representation layer - can be learned (bilinear) or simple dot-product. +pub enum RelationsRepLayer { + /// Learned bilinear projection + Bilinear(BilinearRelationsLayer), + /// Simple dot-product similarity (no learned weights) + DotProduct, +} + +impl RelationsRepLayer { + /// Load the relations layer from GGUF weights. + /// Falls back to dot-product if weights don't exist. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + match BilinearRelationsLayer::load(device, vb) { + Ok(bilinear) => Ok(Self::Bilinear(bilinear)), + Err(_) => Ok(Self::DotProduct), + } + } + + /// Create a dot-product based relations layer (no learned weights). + pub fn identity(_device: &Device, _hidden_size: usize) -> Self { + Self::DotProduct + } + + /// Compute adjacency matrix for entity spans. + /// + /// # Arguments + /// * `entity_embeddings` - Entity span embeddings [batch, num_entities, hidden_size] + /// + /// # Returns + /// Adjacency logits [batch, num_entities, num_entities] (apply sigmoid externally) + pub fn forward(&self, entity_embeddings: &Tensor<3, f32>) -> Tensor<3, f32> { + match self { + Self::Bilinear(layer) => layer.forward(entity_embeddings), + Self::DotProduct => { + // Simple dot product: embeddings @ embeddings.T + let entity_t = entity_embeddings.transpose(1, 2); + entity_embeddings.mat_mul(&entity_t) + } + } + } + + /// Apply sigmoid to logits (for use after forward). + pub fn apply_sigmoid(logits: &[f32]) -> Vec { + logits.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect() + } + + /// Filter entity pairs based on adjacency threshold. + /// + /// # Arguments + /// * `adjacency_scores` - Adjacency matrix [batch, num_entities, num_entities] + /// * `threshold` - Minimum score for a pair to be considered + /// + /// # Returns + /// Vector of (batch_idx, head_idx, tail_idx, score) tuples for pairs above threshold + pub async fn filter_pairs( + &self, + adjacency_scores: &Tensor<3, f32>, + threshold: f32, + ) -> Result> { + let [batch_size, num_entities, _] = adjacency_scores.shape(); + + let scores_slice = adjacency_scores.clone().as_slice().await?; + let scores_data = scores_slice.as_slice(); + + let mut pairs = Vec::new(); + for b in 0..batch_size { + for i in 0..num_entities { + for j in 0..num_entities { + if i == j { + continue; // Skip self-relations + } + let idx = b * num_entities * num_entities + i * num_entities + j; + let score = scores_data[idx]; + if score >= threshold { + pairs.push((b, i, j, score)); + } + } + } + } + + Ok(pairs) + } +} + +/// Learned bilinear relation representation layer. +/// +/// Computes adjacency scores between entity pairs: +/// `adj_score[i,j] = sigmoid(entity_i @ W @ entity_j.T)` +pub struct BilinearRelationsLayer { + /// Bilinear projection weight [hidden_size, hidden_size] + projection: Linear, +} + +impl BilinearRelationsLayer { + /// Load from GGUF weights. + pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { + let projection = Linear::load(device, &mut vb.pp("projection"))?; + Ok(Self { projection }) + } + + /// Compute adjacency matrix for entity spans. + pub fn forward(&self, entity_embeddings: &Tensor<3, f32>) -> Tensor<3, f32> { + // Project entity embeddings: [batch, num_entities, hidden_size] + let projected = self.projection.forward(entity_embeddings); + + // Compute bilinear scores: projected @ entity_embeddings.T + // [batch, num_entities, hidden_size] @ [batch, hidden_size, num_entities] + // = [batch, num_entities, num_entities] + let entity_t = entity_embeddings.transpose(1, 2); + projected.mat_mul(&entity_t) + } +} diff --git a/models/rgliner/src/raw/span_layer.rs b/models/rgliner/src/raw/span_layer.rs index 93418904c..033493ca3 100644 --- a/models/rgliner/src/raw/span_layer.rs +++ b/models/rgliner/src/raw/span_layer.rs @@ -162,6 +162,67 @@ impl SpanLayer { (span_embeddings, span_indices) } + /// Compute span representations for specific (start, end) word positions. + /// + /// Matches Python TokenMarker forward: + /// 1. project_start(h) -> start_rep + /// 2. project_end(h) -> end_rep + /// 3. gather at span positions + /// 4. cat + relu + /// 5. out_project + /// + /// # Arguments + /// * `word_embeddings` - Word embeddings [batch=1, num_words, hidden] + /// * `spans` - List of (start_word, end_word) pairs + /// + /// # Returns + /// Span embeddings [num_spans, hidden] + pub fn forward_for_spans( + &self, + word_embeddings: &Tensor<3, f32>, + spans: &[(usize, usize)], + device: &Device, + ) -> Tensor<2, f32> { + let [batch_size, num_words, hidden_dim] = word_embeddings.shape(); + assert_eq!(batch_size, 1, "only batch_size=1 supported"); + let num_spans = spans.len(); + + // Apply project_start and project_end to the full word embeddings + let start_rep = self + .start_fc2 + .forward(&self.start_fc1.forward(word_embeddings).relu()); + let end_rep = self + .end_fc2 + .forward(&self.end_fc1.forward(word_embeddings).relu()); + + // Gather at span positions + let start_rep_2d = start_rep.squeeze(0).to_concrete(); + let end_rep_2d = end_rep.squeeze(0).to_concrete(); + let _ = num_words; + + let start_indices: Vec = spans.iter().map(|(s, _)| *s as u32).collect(); + let end_indices: Vec = spans.iter().map(|(_, e)| *e as u32).collect(); + let start_idx_tensor = Tensor::new(device, &start_indices); + let end_idx_tensor = Tensor::new(device, &end_indices); + + let start_gathered = start_rep_2d.index_select(0, &start_idx_tensor); + let end_gathered = end_rep_2d.index_select(0, &end_idx_tensor); + + // Concat along last dim: [num_spans, hidden*2] + let start_3d: Tensor<3, f32> = start_gathered.unsqueeze(0).to_concrete(); + let end_3d: Tensor<3, f32> = end_gathered.unsqueeze(0).to_concrete(); + let combined = Tensor::cat([start_3d, end_3d], 2).relu(); + + // Apply out_project: Linear -> ReLU -> Linear + let hidden = self.out_fc1.forward(&combined).relu(); + let out = self.out_fc2.forward(&hidden); + + // [1, num_spans, hidden] -> [num_spans, hidden] + let _ = num_spans; + let _ = hidden_dim; + out.squeeze(0).to_concrete() + } + fn gather_span_embeddings( &self, word_embeddings: &Tensor<3, f32>, diff --git a/models/rgliner/src/relation_decoding.rs b/models/rgliner/src/relation_decoding.rs new file mode 100644 index 000000000..b2c11f481 --- /dev/null +++ b/models/rgliner/src/relation_decoding.rs @@ -0,0 +1,292 @@ +//! Relation decoding with three-threshold approach. +//! +//! Three thresholds control the extraction pipeline: +//! 1. Entity threshold: Confidence cutoff for NER +//! 2. Adjacency threshold: Entity pair candidate filtering +//! 3. Relation threshold: Relation classification cutoff + +use crate::decoding::Entity; + +/// A recognized relation between two entities. +#[derive(Debug, Clone)] +pub struct Relation { + /// The head (source) entity. + pub head: Entity, + /// The tail (target) entity. + pub tail: Entity, + /// The relation type/label. + pub relation: String, + /// Confidence score (0.0 to 1.0). + pub score: f32, +} + +/// Configuration for relation decoding thresholds. +#[derive(Debug, Clone)] +pub struct RelationDecoderConfig { + /// Entity detection threshold (default: 0.4) + pub entity_threshold: f32, + /// Adjacency filtering threshold (default: 0.55) + pub adjacency_threshold: f32, + /// Relation classification threshold (default: 0.8) + pub relation_threshold: f32, +} + +impl Default for RelationDecoderConfig { + fn default() -> Self { + Self { + entity_threshold: 0.4, + adjacency_threshold: 0.55, + relation_threshold: 0.8, + } + } +} + +/// Relation decoder using three-threshold approach. +pub struct RelationDecoder { + config: RelationDecoderConfig, +} + +impl RelationDecoder { + /// Create a new relation decoder with default config. + pub fn new() -> Self { + Self { + config: RelationDecoderConfig::default(), + } + } + + /// Create with custom configuration. + pub fn with_config(config: RelationDecoderConfig) -> Self { + Self { config } + } + + /// Set entity threshold. + pub fn with_entity_threshold(mut self, threshold: f32) -> Self { + self.config.entity_threshold = threshold; + self + } + + /// Set adjacency threshold. + pub fn with_adjacency_threshold(mut self, threshold: f32) -> Self { + self.config.adjacency_threshold = threshold; + self + } + + /// Set relation threshold. + pub fn with_relation_threshold(mut self, threshold: f32) -> Self { + self.config.relation_threshold = threshold; + self + } + + /// Get the current configuration. + pub fn config(&self) -> &RelationDecoderConfig { + &self.config + } + + /// Decode relations from entity pairs and relation scores. + /// + /// # Arguments + /// * `entities` - Detected entities from NER stage + /// * `adjacency_scores` - Adjacency matrix scores [num_entities, num_entities] + /// * `relation_scores` - Relation scores for each pair [num_pairs, num_relations] + /// * `candidate_pairs` - (head_idx, tail_idx) pairs that passed adjacency threshold + /// * `relation_labels` - Relation type labels + /// + /// # Returns + /// Vector of decoded relations above the relation threshold + pub fn decode( + &self, + entities: &[Entity], + _adjacency_scores: &[f32], + relation_scores: &[f32], + candidate_pairs: &[(usize, usize)], + relation_labels: &[&str], + ) -> Vec { + let num_relations = relation_labels.len(); + let mut relations = Vec::new(); + + for (pair_idx, &(head_idx, tail_idx)) in candidate_pairs.iter().enumerate() { + if head_idx >= entities.len() || tail_idx >= entities.len() { + continue; + } + + // Get relation scores for this pair + let scores_start = pair_idx * num_relations; + let scores_end = scores_start + num_relations; + + if scores_end > relation_scores.len() { + continue; + } + + // Find relations above threshold + for (rel_idx, &score) in relation_scores[scores_start..scores_end].iter().enumerate() { + if score >= self.config.relation_threshold { + relations.push(Relation { + head: entities[head_idx].clone(), + tail: entities[tail_idx].clone(), + relation: relation_labels[rel_idx].to_string(), + score, + }); + } + } + } + + // Sort by score descending + relations.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal)); + + relations + } + + /// Filter entity pairs based on adjacency scores. + /// + /// # Arguments + /// * `adjacency_scores` - Flat adjacency matrix [num_entities * num_entities] + /// * `num_entities` - Number of entities + /// + /// # Returns + /// Vector of (head_idx, tail_idx) pairs above the adjacency threshold + pub fn filter_pairs(&self, adjacency_scores: &[f32], num_entities: usize) -> Vec<(usize, usize)> { + let mut pairs = Vec::new(); + + for i in 0..num_entities { + for j in 0..num_entities { + if i == j { + continue; // Skip self-relations + } + + let idx = i * num_entities + j; + if idx < adjacency_scores.len() { + let score = adjacency_scores[idx]; + if score >= self.config.adjacency_threshold { + pairs.push((i, j)); + } + } + } + } + + pairs + } + + /// Pool entity span embeddings by mean pooling. + /// + /// # Arguments + /// * `text_embeddings` - Text token embeddings [num_words, hidden_size] + /// * `entity` - Entity with word indices + /// + /// # Returns + /// Mean-pooled embedding for the entity span + pub fn pool_entity_embedding( + text_embeddings: &[f32], + hidden_size: usize, + entity: &Entity, + ) -> Vec { + let start = entity.start_word; + let end = entity.end_word; + let num_words = end - start + 1; + + if num_words == 0 { + return vec![0.0; hidden_size]; + } + + let mut pooled = vec![0.0f32; hidden_size]; + + for word_idx in start..=end { + let offset = word_idx * hidden_size; + for h in 0..hidden_size { + if offset + h < text_embeddings.len() { + pooled[h] += text_embeddings[offset + h]; + } + } + } + + // Divide by number of words for mean pooling + for h in 0..hidden_size { + pooled[h] /= num_words as f32; + } + + pooled + } +} + +impl Default for RelationDecoder { + fn default() -> Self { + Self::new() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn make_entity(start: usize, end: usize, label: &str) -> Entity { + Entity { + text: format!("entity_{start}_{end}"), + label: label.to_string(), + start_char: start * 5, + end_char: end * 5 + 5, + start_word: start, + end_word: end, + score: 0.9, + } + } + + #[test] + fn test_filter_pairs() { + let decoder = RelationDecoder::new().with_adjacency_threshold(0.5); + + let adjacency_scores = vec![ + 0.0, 0.8, 0.3, // entity 0 -> [0, 1, 2] + 0.7, 0.0, 0.6, // entity 1 -> [0, 1, 2] + 0.2, 0.4, 0.0, // entity 2 -> [0, 1, 2] + ]; + + let pairs = decoder.filter_pairs(&adjacency_scores, 3); + + // Should include (0,1), (1,0), (1,2) where score >= 0.5 + assert!(pairs.contains(&(0, 1))); // 0.8 + assert!(pairs.contains(&(1, 0))); // 0.7 + assert!(pairs.contains(&(1, 2))); // 0.6 + assert!(!pairs.contains(&(0, 2))); // 0.3 < 0.5 + } + + #[test] + fn test_decode_relations() { + let decoder = RelationDecoder::new() + .with_relation_threshold(0.7); + + let entities = vec![ + make_entity(0, 1, "organization"), + make_entity(3, 4, "person"), + make_entity(6, 6, "location"), + ]; + + let adjacency_scores = vec![0.9; 9]; // All pairs pass + let candidate_pairs = vec![(0, 1), (0, 2), (1, 2)]; + + // Relation scores: 3 pairs x 2 relations + let relation_scores = vec![ + 0.85, 0.3, // pair (0,1): "founded by" = 0.85, "located in" = 0.3 + 0.2, 0.9, // pair (0,2): "founded by" = 0.2, "located in" = 0.9 + 0.1, 0.4, // pair (1,2): both below threshold + ]; + + let relation_labels = &["founded by", "located in"]; + + let relations = decoder.decode( + &entities, + &adjacency_scores, + &relation_scores, + &candidate_pairs, + relation_labels, + ); + + assert_eq!(relations.len(), 2); + + // Check first relation (highest score should be located_in at 0.9) + assert_eq!(relations[0].relation, "located in"); + assert_eq!(relations[0].score, 0.9); + + // Check second relation + assert_eq!(relations[1].relation, "founded by"); + assert_eq!(relations[1].score, 0.85); + } +} diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs new file mode 100644 index 000000000..4b45778df --- /dev/null +++ b/models/rgliner/src/relex.rs @@ -0,0 +1,790 @@ +//! GLiNER-RelEx: Joint Named Entity Recognition and Relation Extraction. +//! +//! This module provides the `GlinerRelEx` struct for extracting entities and relations +//! from text using the GLiNER-RelEx model architecture. +//! +//! ## Architecture +//! +//! The model uses the following pipeline: +//! 1. mDeBERTa encoder for contextual embeddings +//! 2. BiLSTM for enhanced token representations +//! 3. Prompt representation layer for label embeddings +//! 4. Joint scorer for token-level BIO predictions +//! 5. Span layer for entity span representations +//! 6. Pair projector for relation classification +//! +//! ## Example +//! +//! ```rust,no_run +//! use rgliner::relex::*; +//! +//! # async fn example() -> anyhow::Result<()> { +//! let mut relex = GlinerRelEx::builder() +//! .with_source(GlinerRelExSource::relex_multi()) +//! .build() +//! .await?; +//! +//! let (entities, relations) = relex.extract( +//! "Apple was founded by Steve Jobs in California.", +//! &["person", "organization", "location"], +//! &["founded by", "located in"], +//! ).await?; +//! +//! for relation in relations { +//! println!("{} --[{}]--> {}", +//! relation.head.text, +//! relation.relation, +//! relation.tail.text +//! ); +//! } +//! # Ok(()) +//! # } +//! ``` + +use std::sync::Arc; + +use fusor::{Device, Tensor, VarBuilder}; +use kalosm_common::Cache; +use kalosm_model_types::{FileSource, ModelLoadingProgress}; +use tokenizers::Tokenizer; + +use crate::decoding::Entity; +use crate::error::{GlinerError, GlinerLoadingError}; +use crate::raw::mdeberta::MDebertaModel; +use crate::raw::{BiLstm, JointScorer, PairProjector, PromptRepLayer, RelationsRepLayer, SpanLayer}; +use crate::relation_decoding::{Relation, RelationDecoder, RelationDecoderConfig}; +use crate::relex_tokenization::{RelExTokenizer, SpecialTokenIds}; + +/// Source configuration for GLiNER-RelEx models. +/// +/// The GGUF file produced by `convert_relex_to_gguf.py` embeds the tokenizer +/// JSON and GLiNER config JSON as string metadata, so only the model file is +/// required. `tokenizer` and `config` can optionally override the embedded +/// copies (e.g., to swap in a custom tokenizer). +pub struct GlinerRelExSource { + /// Main model GGUF file (encoder + all layers + embedded tokenizer/config) + pub model: FileSource, + /// Optional tokenizer JSON override. If `None`, the tokenizer is read from + /// the `gliner.tokenizer_json` metadata embedded in the GGUF. + pub tokenizer: Option, + /// Optional GLiNER config JSON override. If `None`, the config is read from + /// the `gliner.config_json` metadata embedded in the GGUF. + pub config: Option, +} + +impl GlinerRelExSource { + /// GLiNER-RelEx Multi v1.0 source. + /// + /// Downloads the GGUF-converted weights from HuggingFace. Tokenizer and + /// config are embedded in the GGUF file. + pub fn relex_multi() -> Self { + Self { + model: FileSource::huggingface( + "knowledgator/gliner-relex-multi-v1.0-gguf".to_string(), + "main".to_string(), + "gliner-relex-multi-v1.0-Q8_0.gguf".to_string(), + ), + tokenizer: None, + config: None, + } + } + + /// Create a source from a local GGUF file. + /// + /// The tokenizer and config are expected to be embedded in the GGUF + /// metadata (produced by `convert_relex_to_gguf.py`). + pub fn local(model_path: impl Into) -> Self { + Self { + model: FileSource::local(model_path.into()), + tokenizer: None, + config: None, + } + } + + /// Override the tokenizer source (otherwise read from GGUF metadata). + pub fn with_tokenizer(mut self, tokenizer: FileSource) -> Self { + self.tokenizer = Some(tokenizer); + self + } + + /// Override the config source (otherwise read from GGUF metadata). + pub fn with_config(mut self, config: FileSource) -> Self { + self.config = Some(config); + self + } +} + +impl Default for GlinerRelExSource { + fn default() -> Self { + Self::relex_multi() + } +} + +/// Configuration for GLiNER-RelEx model. +#[derive(Debug, Clone)] +pub struct GlinerRelExConfig { + /// Maximum span width in words + pub max_width: usize, + /// Hidden dimension + pub hidden_size: usize, + /// Entity detection threshold + pub entity_threshold: f32, + /// Adjacency filtering threshold + pub adjacency_threshold: f32, + /// Relation classification threshold + pub relation_threshold: f32, + /// Special token IDs + pub special_tokens: SpecialTokenIds, +} + +impl Default for GlinerRelExConfig { + fn default() -> Self { + Self { + max_width: 12, + hidden_size: 768, + entity_threshold: 0.4, + adjacency_threshold: 0.55, + relation_threshold: 0.8, + special_tokens: SpecialTokenIds::default(), + } + } +} + +/// Builder for constructing a [`GlinerRelEx`] model. +#[derive(Default)] +pub struct GlinerRelExBuilder { + source: GlinerRelExSource, + cache: Cache, + device: Option, + config: GlinerRelExConfig, +} + +impl GlinerRelExBuilder { + /// Set the model source. + pub fn with_source(mut self, source: GlinerRelExSource) -> Self { + self.source = source; + self + } + + /// Set the entity threshold. + pub fn with_entity_threshold(mut self, threshold: f32) -> Self { + self.config.entity_threshold = threshold; + self + } + + /// Set the adjacency threshold. + pub fn with_adjacency_threshold(mut self, threshold: f32) -> Self { + self.config.adjacency_threshold = threshold; + self + } + + /// Set the relation threshold. + pub fn with_relation_threshold(mut self, threshold: f32) -> Self { + self.config.relation_threshold = threshold; + self + } + + /// Set the maximum span width. + pub fn with_max_width(mut self, max_width: usize) -> Self { + self.config.max_width = max_width; + self + } + + /// Set the device. + pub fn with_device(mut self, device: Device) -> Self { + self.device = Some(device); + self + } + + /// Set the cache location. + pub fn with_cache(mut self, cache: Cache) -> Self { + self.cache = cache; + self + } + + /// Build the model. + pub async fn build(self) -> Result { + self.build_with_loading_handler(ModelLoadingProgress::multi_bar_loading_indicator()) + .await + } + + /// Build the model with a loading handler. + pub async fn build_with_loading_handler( + self, + loading_handler: impl FnMut(ModelLoadingProgress) + Send + 'static, + ) -> Result { + GlinerRelEx::from_builder(self, loading_handler).await + } +} + +/// GLiNER-RelEx model for joint NER and relation extraction. +pub struct GlinerRelEx { + /// mDeBERTa encoder + encoder: MDebertaModel, + /// BiLSTM for enhanced token representations + bilstm: BiLstm, + /// Prompt representation layer for label projection + prompt_rep_layer: PromptRepLayer, + /// Joint scorer for token-level predictions + scorer: JointScorer, + /// Span representation layer + span_layer: SpanLayer, + /// Relations representation layer (adjacency scoring) + relations_layer: RelationsRepLayer, + /// Entity pair projector + pair_projector: PairProjector, + /// Tokenizer with special token handling + tokenizer: Arc, + /// Relation decoder + relation_decoder: RelationDecoder, + /// Device + device: Device, + /// Configuration + config: GlinerRelExConfig, +} + +fn default_device() -> Device { + std::panic::catch_unwind(Device::gpu_blocking) + .ok() + .and_then(Result::ok) + .unwrap_or_else(Device::cpu) +} + +impl GlinerRelEx { + /// Create a new builder. + pub fn builder() -> GlinerRelExBuilder { + GlinerRelExBuilder::default() + } + + /// Create with default settings. + pub async fn new() -> Result { + Self::builder().build().await + } + + async fn from_builder( + builder: GlinerRelExBuilder, + mut progress_handler: impl FnMut(ModelLoadingProgress) + Send + 'static, + ) -> Result { + let GlinerRelExBuilder { + source, + cache, + device, + config, + } = builder; + + // Download main model weights first - the GGUF may also contain the + // tokenizer and config as embedded metadata. + let model_source = format!("Model ({})", source.model); + let mut create_progress = ModelLoadingProgress::downloading_progress(model_source); + let model_bytes = cache + .get_bytes(&source.model, |progress| { + progress_handler(create_progress(progress)) + }) + .await?; + + // Initialize device + let device = device.unwrap_or_else(default_device); + + // Load model components from GGUF + let mut model_cursor = std::io::Cursor::new(&model_bytes); + let mut vb = VarBuilder::from_gguf(&mut model_cursor) + .map_err(|err| GlinerLoadingError::LoadModel(fusor::Error::from(err)))?; + + // Resolve tokenizer: explicit override > embedded metadata. + let tokenizer_bytes: Vec = if let Some(tokenizer_src) = source.tokenizer.as_ref() { + let tok_label = format!("Tokenizer ({})", tokenizer_src); + let mut create_progress = ModelLoadingProgress::downloading_progress(tok_label); + cache + .get_bytes(tokenizer_src, |progress| { + progress_handler(create_progress(progress)) + }) + .await? + .to_vec() + } else { + let meta = vb + .get_metadata("gliner.tokenizer_json") + .and_then(|v| v.to_string().ok()) + .ok_or_else(|| { + GlinerLoadingError::LoadModel(fusor::Error::msg( + "GGUF missing embedded tokenizer (metadata key `gliner.tokenizer_json`). \ + Re-run convert_relex_to_gguf.py or set a tokenizer source via \ + `GlinerRelExSource::with_tokenizer`.", + )) + })?; + meta.as_bytes().to_vec() + }; + + let tokenizer = + Tokenizer::from_bytes(&tokenizer_bytes).map_err(GlinerLoadingError::LoadTokenizer)?; + let relex_tokenizer = RelExTokenizer::with_special_tokens(tokenizer, config.special_tokens.clone()); + + // Load encoder (mDeBERTa) + let encoder = MDebertaModel::load(&device, &mut vb.pp("text"))?; + + // Load BiLSTM + let bilstm = BiLstm::load(&device, &mut vb.pp("rnn"))?; + + // Load prompt representation layer + let prompt_rep_layer = PromptRepLayer::load(&device, &mut vb.pp("prompt_rep_layer"))?; + + // Load joint scorer + let scorer = JointScorer::load(&device, &mut vb.pp("scorer"))?; + + // Load span layer + let span_layer = SpanLayer::load(&device, &mut vb, config.max_width)?; + + // Load relations layer (may not exist, use projection from pair_proj) + let relations_layer = RelationsRepLayer::load(&device, &mut vb.pp("relations")) + .unwrap_or_else(|_| RelationsRepLayer::identity(&device, config.hidden_size)); + + // Load pair projector + let pair_projector = PairProjector::load(&device, &mut vb.pp("pair_proj"))?; + + // Create relation decoder + let relation_decoder = RelationDecoder::with_config(RelationDecoderConfig { + entity_threshold: config.entity_threshold, + adjacency_threshold: config.adjacency_threshold, + relation_threshold: config.relation_threshold, + }); + + Ok(Self { + encoder, + bilstm, + prompt_rep_layer, + scorer, + span_layer, + relations_layer, + pair_projector, + tokenizer: Arc::new(relex_tokenizer), + relation_decoder, + device, + config, + }) + } + + /// Extract entities and relations from text. + /// + /// # Arguments + /// * `text` - Input text + /// * `entity_labels` - Entity type labels (e.g., ["person", "organization"]) + /// * `relation_labels` - Relation type labels (e.g., ["founded by", "works at"]) + /// + /// # Returns + /// Tuple of (entities, relations) + pub async fn extract( + &self, + text: &str, + entity_labels: &[&str], + relation_labels: &[&str], + ) -> Result<(Vec, Vec), GlinerError> { + // 1. Tokenize with special tokens + let tokenized = self.tokenizer.tokenize(text, entity_labels, relation_labels)?; + + #[cfg(debug_assertions)] + { + eprintln!("[DEBUG] Tokenized: {} tokens, {} words", + tokenized.token_ids.len(), tokenized.num_words); + eprintln!("[DEBUG] ent_positions: {:?}", tokenized.ent_positions); + eprintln!("[DEBUG] rel_positions: {:?}", tokenized.rel_positions); + eprintln!("[DEBUG] text_positions: {:?}", tokenized.text_positions); + eprintln!("[DEBUG] token_ids (first 20): {:?}", + &tokenized.token_ids[..20.min(tokenized.token_ids.len())]); + } + + if tokenized.num_words == 0 { + return Ok((Vec::new(), Vec::new())); + } + + // 2. Prepare input tensors + let token_ids = Tensor::new(&self.device, &tokenized.token_ids); + let token_ids: Tensor<2, u32> = token_ids.unsqueeze(0).to_concrete(); + + let attention_mask = Tensor::new(&self.device, &tokenized.attention_mask); + let attention_mask: Tensor<2, u32> = attention_mask.unsqueeze(0).to_concrete(); + + // 3. Forward pass through encoder + let encoder_output = self.encoder.forward(&token_ids, Some(&attention_mask)); + + #[cfg(debug_assertions)] + { + let enc_data = encoder_output.clone().as_slice().await.unwrap(); + let enc_slice = enc_data.as_slice(); + let mean: f32 = enc_slice.iter().sum::() / enc_slice.len() as f32; + let variance: f32 = enc_slice.iter().map(|x| (x - mean).powi(2)).sum::() / enc_slice.len() as f32; + eprintln!("[DEBUG] Encoder output stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", + mean, variance, + enc_slice.iter().cloned().fold(f32::INFINITY, f32::min), + enc_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max)); + + // Check encoder output at specific positions (for <> tokens) + let hidden_size = self.config.hidden_size; + eprintln!("[DEBUG] Encoder output at <> positions (first 5 values):"); + for &pos in &tokenized.ent_positions { + let start = pos * hidden_size; + let vals: Vec = (0..5).map(|i| enc_slice[start + i]).collect(); + eprintln!(" pos {}: {:?}", pos, vals); + } + // Also check a few other positions for comparison + eprintln!("[DEBUG] Encoder output at other positions:"); + for &pos in &[0, 2, 4, 10, 17] { + if pos < tokenized.token_ids.len() { + let start = pos * hidden_size; + let vals: Vec = (0..5).map(|i| enc_slice[start + i]).collect(); + eprintln!(" pos {}: {:?}", pos, vals); + } + } + } + + // 4. Extract word-level embeddings from encoder output, THEN apply BiLSTM + // (Python applies BiLSTM to word-level embeddings, not the full token sequence.) + let word_encoder_embs = self.gather_at_positions(&encoder_output, &tokenized.text_positions); + let lstm_output = self.bilstm.forward(&word_encoder_embs).await; + + #[cfg(debug_assertions)] + { + let lstm_data = lstm_output.clone().as_slice().await.unwrap(); + let lstm_slice = lstm_data.as_slice(); + let hidden_size = self.config.hidden_size; + eprintln!("[DEBUG] Word-level BiLSTM output (first 5 values per word):"); + for w in 0..tokenized.num_words { + let start = w * hidden_size; + let vals: Vec = (0..5).map(|i| lstm_slice[start + i]).collect(); + eprintln!(" word {}: {:?}", w, vals); + } + } + + // 5. Extract label embeddings at marker positions from ENCODER output and project them + // (Labels are extracted from encoder output, text tokens from BiLSTM output) + // Entity label embeddings: hidden states at <> positions + let ent_embs_raw = self.gather_at_positions(&encoder_output, &tokenized.ent_positions); + + #[cfg(debug_assertions)] + { + let raw_data = ent_embs_raw.clone().as_slice().await.unwrap(); + let raw_slice = raw_data.as_slice(); + let hidden_size = self.config.hidden_size; + eprintln!("[DEBUG] Raw entity embeddings check (first 5 values per label):"); + for l in 0..tokenized.ent_positions.len() { + let start = l * hidden_size; + let vals: Vec = (0..5).map(|i| raw_slice[start + i]).collect(); + eprintln!(" label {} (pos {}): {:?}", l, tokenized.ent_positions[l], vals); + } + } + + let ent_embs = self.prompt_rep_layer.forward_3d(&ent_embs_raw); + + #[cfg(debug_assertions)] + { + let ent_data = ent_embs.clone().as_slice().await.unwrap(); + let ent_slice = ent_data.as_slice(); + let hidden_size = self.config.hidden_size; + let mean: f32 = ent_slice.iter().sum::() / ent_slice.len() as f32; + let variance: f32 = ent_slice.iter().map(|x| (x - mean).powi(2)).sum::() / ent_slice.len() as f32; + eprintln!("[DEBUG] Entity label embs stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", + mean, variance, + ent_slice.iter().cloned().fold(f32::INFINITY, f32::min), + ent_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max)); + // Print projected values per label (compare with Python) + eprintln!("[DEBUG] Projected entity embeddings (first 5 values per label):"); + for l in 0..entity_labels.len() { + let start = l * hidden_size; + let vals: Vec = (0..5).map(|i| ent_slice[start + i]).collect(); + eprintln!(" {}: {:?}", entity_labels[l], vals); + } + } + + // Relation label embeddings: raw hidden states at <> positions + // (unlike entity labels, relation labels are NOT projected through prompt_rep_layer) + let rel_embs = self.gather_at_positions(&encoder_output, &tokenized.rel_positions); + + // 6. Text embeddings = BiLSTM output (already at word level) + let text_embs = lstm_output.clone(); + + #[cfg(debug_assertions)] + { + let text_data = text_embs.clone().as_slice().await.unwrap(); + let text_slice = text_data.as_slice(); + let mean: f32 = text_slice.iter().sum::() / text_slice.len() as f32; + let variance: f32 = text_slice.iter().map(|x| (x - mean).powi(2)).sum::() / text_slice.len() as f32; + eprintln!("[DEBUG] Text token embs stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", + mean, variance, + text_slice.iter().cloned().fold(f32::INFINITY, f32::min), + text_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max)); + } + + // 7. Score tokens against entity labels using joint scorer + let ent_embs_2d: Tensor<2, f32> = ent_embs.squeeze(0).to_concrete(); + let token_scores = self.scorer.forward_entity_scores(&text_embs, &ent_embs_2d).await; + + // 8. Decode entities from token-level scores + let entities = self.decode_entities_from_tokens( + &token_scores, + entity_labels, + &tokenized.word_offsets, + text, + ).await?; + + // If no entities or no relation labels, return early + if entities.len() < 2 || relation_labels.is_empty() { + return Ok((entities, Vec::new())); + } + + // 9. Compute span representations for each entity using span_layer + // (matches Python's TokenMarker: project_start/project_end MLPs + out_project MLP) + let entity_spans: Vec<(usize, usize)> = entities + .iter() + .map(|e| (e.start_word, e.end_word)) + .collect(); + let span_reps = self.span_layer.forward_for_spans(&text_embs, &entity_spans, &self.device); + // span_reps shape: [num_entities, hidden] + + #[cfg(debug_assertions)] + { + let sr_data = span_reps.clone().as_slice().await?; + let sr = sr_data.as_slice(); + let hidden = self.config.hidden_size; + eprintln!("[DEBUG] Span reps (first 5 values per entity):"); + for (i, e) in entities.iter().enumerate() { + let start = i * hidden; + let vals: Vec = (0..5).map(|k| sr[start + k]).collect(); + eprintln!(" {} ({}, {}): {:?}", e.text, e.start_word, e.end_word, vals); + } + // Print rel_embs + let re_data = rel_embs.clone().as_slice().await?; + let re = re_data.as_slice(); + eprintln!("[DEBUG] Rel embs (first 5 values per label):"); + for (i, l) in relation_labels.iter().enumerate() { + let start = i * hidden; + let vals: Vec = (0..5).map(|k| re[start + k]).collect(); + eprintln!(" {}: {:?}", l, vals); + } + } + + let num_entities = entities.len(); + let hidden_size = self.config.hidden_size; + + // 10. Build all entity pairs (head, tail) with head != tail + let mut candidate_pairs: Vec<(usize, usize)> = Vec::new(); + for head in 0..num_entities { + for tail in 0..num_entities { + if head != tail { + candidate_pairs.push((head, tail)); + } + } + } + + // 11. Gather head and tail span reps using index_select + let span_reps_data = span_reps.clone().as_slice().await?; + let span_reps_slice = span_reps_data.as_slice(); + let mut head_embs = Vec::with_capacity(candidate_pairs.len() * hidden_size); + let mut tail_embs = Vec::with_capacity(candidate_pairs.len() * hidden_size); + for &(head_idx, tail_idx) in &candidate_pairs { + let h_start = head_idx * hidden_size; + let t_start = tail_idx * hidden_size; + head_embs.extend_from_slice(&span_reps_slice[h_start..h_start + hidden_size]); + tail_embs.extend_from_slice(&span_reps_slice[t_start..t_start + hidden_size]); + } + + let head_tensor = Tensor::new(&self.device, &head_embs) + .reshape([candidate_pairs.len(), hidden_size]) + .to_concrete(); + let tail_tensor = Tensor::new(&self.device, &tail_embs) + .reshape([candidate_pairs.len(), hidden_size]) + .to_concrete(); + + // 12. Apply pair_projector: concat(head, tail) -> MLP -> pair_rep + let pair_embs = self.pair_projector.forward(&head_tensor, &tail_tensor); + + // 13. Score pairs against relation labels via dot product (no sigmoid yet) + let rel_embs_squeezed: Tensor<2, f32> = rel_embs.squeeze(0).to_concrete(); + let rel_scores = pair_embs.mat_mul(&rel_embs_squeezed.transpose(0, 1)); + + // 14. Apply sigmoid and filter by relation_threshold + let rel_scores_slice = rel_scores.clone().as_slice().await?; + let n_rels = relation_labels.len(); + let mut relations = Vec::new(); + let threshold = self.config.relation_threshold; + + #[cfg(debug_assertions)] + { + eprintln!( + "[DEBUG] Relation scoring: {} pairs, {} relations, threshold={}", + candidate_pairs.len(), n_rels, threshold); + for (pair_idx, &(h, t)) in candidate_pairs.iter().enumerate().take(6) { + let base = pair_idx * n_rels; + let raw: Vec = (0..n_rels).map(|c| rel_scores_slice.as_slice()[base + c]).collect(); + let sig: Vec = raw.iter().map(|x| 1.0 / (1.0 + (-x).exp())).collect(); + eprintln!( + " pair ({}->{}) [{} -> {}]: raw={:?}, sig={:?}", + h, t, entities[h].text, entities[t].text, raw, sig); + } + } + + for (pair_idx, &(head_idx, tail_idx)) in candidate_pairs.iter().enumerate() { + let base = pair_idx * n_rels; + for rel_idx in 0..n_rels { + let raw = rel_scores_slice.as_slice()[base + rel_idx]; + let prob = 1.0 / (1.0 + (-raw).exp()); + if prob > threshold { + relations.push(Relation { + head: entities[head_idx].clone(), + tail: entities[tail_idx].clone(), + relation: relation_labels[rel_idx].to_string(), + score: prob, + }); + } + } + } + + // Sort by score descending + relations.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal)); + + Ok((entities, relations)) + } + + /// Decode entities using span-boundary detection with start/end/inside scores. + /// + /// `token_scores` has shape [batch, seq_len, n_labels, 3] where the last dim + /// is [start, end, inside] sigmoid probabilities. + async fn decode_entities_from_tokens( + &self, + token_scores: &Tensor<4, f32>, + entity_labels: &[&str], + word_offsets: &[(usize, usize)], + text: &str, + ) -> Result, GlinerError> { + let [_batch_size, num_tokens, num_labels, num_channels] = token_scores.shape(); + assert_eq!(num_channels, 3, "expected [start, end, inside]"); + let scores_data = token_scores.clone().as_slice().await?; + let scores = scores_data.as_slice(); + + let threshold = self.config.entity_threshold; + + #[cfg(debug_assertions)] + { + eprintln!("[DEBUG] Entity decoding: num_tokens={}, num_labels={}, threshold={}", + num_tokens, num_labels, threshold); + eprintln!("[DEBUG] Sigmoid scores [start, end, inside] (first 5 tokens):"); + for t in 0..5.min(num_tokens) { + for l in 0..num_labels { + let base = t * num_labels * 3 + l * 3; + eprintln!( + " token {} label {} ({}): start={:.4}, end={:.4}, inside={:.4}", + t, l, entity_labels[l], + scores[base], scores[base + 1], scores[base + 2]); + } + } + } + + // Candidate spans: (start, end, label, score) + let mut candidates: Vec<(usize, usize, usize, f32)> = Vec::new(); + + let score_at = |tok: usize, lab: usize, ch: usize| -> f32 { + scores[tok * num_labels * 3 + lab * 3 + ch] + }; + + for label_idx in 0..num_labels { + for start_tok in 0..num_tokens { + let start_score = score_at(start_tok, label_idx, 0); + if start_score < threshold { continue; } + + for end_tok in start_tok..num_tokens { + let end_score = score_at(end_tok, label_idx, 1); + if end_score < threshold { continue; } + + // Check all inside scores from start_tok to end_tok + let mut min_score = start_score.min(end_score); + let mut valid = true; + for t in start_tok..=end_tok { + let inside = score_at(t, label_idx, 2); + if inside < threshold { valid = false; break; } + if inside < min_score { min_score = inside; } + } + if !valid { continue; } + + candidates.push((start_tok, end_tok, label_idx, min_score)); + } + } + } + + // Sort candidates by score descending + candidates.sort_by(|a, b| b.3.partial_cmp(&a.3).unwrap_or(std::cmp::Ordering::Equal)); + + // Greedy filter non-overlapping spans (flat_ner equivalent) + let mut taken: Vec<(usize, usize)> = Vec::new(); + let mut entities = Vec::new(); + for (start_tok, end_tok, label_idx, score) in candidates { + let overlap = taken.iter().any(|&(a, b)| !(end_tok < a || start_tok > b)); + if overlap { continue; } + taken.push((start_tok, end_tok)); + + if start_tok < word_offsets.len() && end_tok < word_offsets.len() { + let (start_char, _) = word_offsets[start_tok]; + let (_, end_char) = word_offsets[end_tok]; + entities.push(Entity { + text: text[start_char..end_char].to_string(), + label: entity_labels[label_idx].to_string(), + start_char, + end_char, + start_word: start_tok, + end_word: end_tok, + score, + }); + } + } + + // Sort by score descending + entities.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal)); + + Ok(entities) + } + + /// Gather hidden states at specific positions. + fn gather_at_positions( + &self, + hidden_states: &Tensor<3, f32>, + positions: &[usize], + ) -> Tensor<3, f32> { + let [batch_size, _seq_len, hidden_size] = hidden_states.shape(); + let num_positions = positions.len(); + + if num_positions == 0 { + return Tensor::zeros(&self.device, [batch_size, 1, hidden_size]); + } + + // Build index tensor + let indices: Vec = positions.iter().map(|&p| p as u32).collect(); + let index_tensor = Tensor::new(&self.device, &indices); + + // For batch size 1, we can use index_select + let hidden_2d = hidden_states.squeeze(0).to_concrete(); + let gathered = hidden_2d.index_select(0, &index_tensor); + + gathered.unsqueeze(0).to_concrete() + } + + /// Build a tensor from entity embeddings. + fn build_entity_tensor(&self, embeddings: &[Vec], device: &Device) -> Tensor<3, f32> { + let num_entities = embeddings.len(); + if num_entities == 0 || embeddings[0].is_empty() { + return Tensor::zeros(device, [1, 1, self.config.hidden_size]); + } + + let hidden_size = embeddings[0].len(); + let flat: Vec = embeddings.iter().flatten().copied().collect(); + + Tensor::new(device, &flat) + .reshape([1, num_entities, hidden_size]) + .to_concrete() + } + + /// Get the device. + pub fn device(&self) -> &Device { + &self.device + } + + /// Get the configuration. + pub fn config(&self) -> &GlinerRelExConfig { + &self.config + } +} diff --git a/models/rgliner/src/relex_tokenization.rs b/models/rgliner/src/relex_tokenization.rs new file mode 100644 index 000000000..a82cc59d3 --- /dev/null +++ b/models/rgliner/src/relex_tokenization.rs @@ -0,0 +1,278 @@ +//! RelEx tokenization with special token handling. +//! +//! Builds joint input sequences in the format: +//! `[CLS] <> label1 <> label2 <> <> rel1 <> rel2 <> text... [SEP]` + +use crate::error::GlinerError; +use std::sync::Arc; +use tokenizers::Tokenizer; + +/// Special token IDs for GLiNER-RelEx. +#[derive(Debug, Clone)] +pub struct SpecialTokenIds { + /// [CLS] token ID. + pub cls_id: u32, + /// [SEP] token ID (end-of-sequence separator). + pub sep_id: u32, + /// [PAD] token ID. + pub pad_id: u32, + /// <> token ID (entity marker). + pub ent_id: u32, + /// <> token ID (relation marker). + pub rel_id: u32, + /// <> token ID (internal separator between entity/relation/text blocks). + pub inner_sep_id: u32, +} + +impl Default for SpecialTokenIds { + fn default() -> Self { + Self { + cls_id: 1, // [CLS] token + sep_id: 2, // [SEP] token (end-of-sequence) + pad_id: 0, // [PAD] token + ent_id: 250102, // <> token + rel_id: 250104, // <> token + inner_sep_id: 250103, // <> token (internal separator) + } + } +} + +/// Tokenized RelEx input with position tracking. +#[derive(Debug, Clone)] +pub struct RelExTokenizedInput { + /// Token IDs for the full sequence + pub token_ids: Vec, + /// Attention mask (1 for real tokens, 0 for padding) + pub attention_mask: Vec, + /// Positions of <> tokens (indices into token_ids) + pub ent_positions: Vec, + /// Positions of <> tokens (indices into token_ids) + pub rel_positions: Vec, + /// Positions of first subtoken for each text word (indices into token_ids) + pub text_positions: Vec, + /// Word offsets in the original text (start_char, end_char) + pub word_offsets: Vec<(usize, usize)>, + /// Number of text words + pub num_words: usize, + /// Number of entity labels + pub num_entity_labels: usize, + /// Number of relation labels + pub num_relation_labels: usize, +} + +/// RelEx tokenizer for building joint input sequences. +pub struct RelExTokenizer { + tokenizer: Arc, + special_tokens: SpecialTokenIds, +} + +impl RelExTokenizer { + /// Create a new RelEx tokenizer. + pub fn new(tokenizer: Tokenizer) -> Self { + Self { + tokenizer: Arc::new(tokenizer), + special_tokens: SpecialTokenIds::default(), + } + } + + /// Create with custom special token IDs. + pub fn with_special_tokens(tokenizer: Tokenizer, special_tokens: SpecialTokenIds) -> Self { + Self { + tokenizer: Arc::new(tokenizer), + special_tokens, + } + } + + /// Tokenize text, entity labels, and relation labels into a joint sequence. + /// + /// Output format (matches Python GLiNER): + /// `[CLS] <> label1 <> label2 <> <> rel1 <> rel2 <> word1 word2 ... [SEP]` + /// + /// Key details: + /// - `<>` (inner_sep_id) is used between entity/relation/text blocks (not `[SEP]`) + /// - `[SEP]` is used only at the end of the sequence + /// - Each word in the text is tokenized independently so that SentencePiece adds the + /// leading `▁` marker for each word. + pub fn tokenize( + &self, + text: &str, + entity_labels: &[&str], + relation_labels: &[&str], + ) -> Result { + let mut token_ids = Vec::new(); + let mut ent_positions = Vec::new(); + let mut rel_positions = Vec::new(); + let mut text_positions = Vec::new(); + let mut word_offsets = Vec::new(); + + // Start with [CLS] + token_ids.push(self.special_tokens.cls_id); + + // Encode entity labels block: <> label1 <> label2 ... + for label in entity_labels { + ent_positions.push(token_ids.len()); + token_ids.push(self.special_tokens.ent_id); + + let label_encoding = self + .tokenizer + .encode(label.to_string(), false) + .map_err(|e| GlinerError::TokenizationError(e.to_string()))?; + token_ids.extend(label_encoding.get_ids().iter().copied()); + } + // Internal separator + token_ids.push(self.special_tokens.inner_sep_id); + + // Encode relation labels block: <> rel1 <> rel2 ... + for label in relation_labels { + rel_positions.push(token_ids.len()); + token_ids.push(self.special_tokens.rel_id); + + let label_encoding = self + .tokenizer + .encode(label.to_string(), false) + .map_err(|e| GlinerError::TokenizationError(e.to_string()))?; + token_ids.extend(label_encoding.get_ids().iter().copied()); + } + // Internal separator between relations and text + token_ids.push(self.special_tokens.inner_sep_id); + + // Encode text with word-level tracking: each word is tokenized separately + let words = self.split_words(text); + for (word, (start_char, end_char)) in words { + text_positions.push(token_ids.len()); + word_offsets.push((start_char, end_char)); + + let word_encoding = self + .tokenizer + .encode(word.to_string(), false) + .map_err(|e| GlinerError::TokenizationError(e.to_string()))?; + token_ids.extend(word_encoding.get_ids().iter().copied()); + } + + // Final [SEP] + token_ids.push(self.special_tokens.sep_id); + + let num_words = text_positions.len(); + let attention_mask = vec![1u32; token_ids.len()]; + + Ok(RelExTokenizedInput { + token_ids, + attention_mask, + ent_positions, + rel_positions, + text_positions, + word_offsets, + num_words, + num_entity_labels: entity_labels.len(), + num_relation_labels: relation_labels.len(), + }) + } + + /// Split text into words with character offsets. + /// + /// Matches Python GLiNER's `WhitespaceTokenSplitter` regex: + /// `\w+(?:[-_]\w+)*|\S` + /// + /// This yields: + /// - Runs of word characters (alphanumeric/underscore), with hyphens/underscores + /// joining word-like groups (e.g. "foo-bar", "x_1") + /// - OR any single non-whitespace character (punctuation as its own token) + fn split_words<'a>(&self, text: &'a str) -> Vec<(&'a str, (usize, usize))> { + let mut words = Vec::new(); + let bytes = text.as_bytes(); + let n = bytes.len(); + let mut i = 0; + + // Helpers operating on byte indices. Text is ASCII-safe for typical inputs; + // for non-ASCII, is_word_char treats each UTF-8 byte - safe approximation that + // matches what Python's \w would do for ASCII-only text. + let is_word_char = |c: u8| c.is_ascii_alphanumeric() || c == b'_'; + let is_whitespace = |c: u8| matches!(c, b' ' | b'\t' | b'\n' | b'\r'); + + while i < n { + let c = bytes[i]; + if is_whitespace(c) { + i += 1; + continue; + } + if is_word_char(c) { + // Match \w+(?:[-_]\w+)* + let start = i; + while i < n && is_word_char(bytes[i]) { + i += 1; + } + // Try to extend with (-|_)\w+ groups (greedy) + loop { + if i + 1 < n && (bytes[i] == b'-' || bytes[i] == b'_') && is_word_char(bytes[i + 1]) { + i += 1; + while i < n && is_word_char(bytes[i]) { + i += 1; + } + } else { + break; + } + } + let end = i; + words.push((&text[start..end], (start, end))); + } else { + // \S - single non-whitespace character (byte here; for ASCII this is fine) + // Advance by one UTF-8 codepoint + let char_len = std::str::from_utf8(&bytes[i..i.saturating_add(4).min(n)]) + .ok() + .and_then(|s| s.chars().next()) + .map(|c| c.len_utf8()) + .unwrap_or(1); + let start = i; + let end = i + char_len; + words.push((&text[start..end], (start, end))); + i = end; + } + } + + words + } + + /// Get the underlying tokenizer. + pub fn tokenizer(&self) -> &Tokenizer { + &self.tokenizer + } + + /// Get special token IDs. + pub fn special_tokens(&self) -> &SpecialTokenIds { + &self.special_tokens + } +} + +#[cfg(test)] +mod tests { + /// Helper function to test word splitting logic without needing a tokenizer. + fn split_words(text: &str) -> Vec<(&str, (usize, usize))> { + let mut words = Vec::new(); + let mut char_idx = 0; + + for word in text.split_whitespace() { + if let Some(pos) = text[char_idx..].find(word) { + let start = char_idx + pos; + let end = start + word.len(); + words.push((word, (start, end))); + char_idx = end; + } + } + + words + } + + #[test] + fn test_split_words() { + let text = "Apple Inc. was founded by Steve Jobs."; + let words = split_words(text); + + assert_eq!(words.len(), 7); + assert_eq!(words[0].0, "Apple"); + assert_eq!(words[0].1, (0, 5)); + assert_eq!(words[1].0, "Inc."); + assert_eq!(words[1].1, (6, 10)); + assert_eq!(words[5].0, "Steve"); + assert_eq!(words[6].0, "Jobs."); + } +} diff --git a/models/rgliner/src/source.rs b/models/rgliner/src/source.rs index 8cedaaf4b..2a3b6560d 100644 --- a/models/rgliner/src/source.rs +++ b/models/rgliner/src/source.rs @@ -143,22 +143,139 @@ impl GlinerSource { label_encoder_config: Self::huggingface_or_cached( "Demonthos/gliner-gguf", "main", - "label-encoder-config.json", + "edge-label-encoder-config.json", ), label_encoder_tokenizer: Self::huggingface_or_cached( "Demonthos/gliner-gguf", "main", - "label-encoder-tokenizer.json", + "edge-label-encoder-tokenizer.json", ), tokenizer: Self::huggingface_or_cached( "Demonthos/gliner-gguf", "main", - "text-tokenizer.json", + "edge-text-tokenizer.json", ), config: Self::huggingface_or_cached( "Demonthos/gliner-gguf", "main", - "text-gliner-config.json", + "edge-text-gliner-config.json", + ), + } + } + + /// Demonthos GLiNER GGUF small upload. + /// + /// Uses the GGUF weights and sidecar tokenizer/config files from + /// `Demonthos/gliner-gguf`. + pub fn demonthos_small() -> Self { + Self { + model: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "gliner-small.gguf", + ), + label_encoder: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "gliner-small-label-encoder.gguf", + ), + label_encoder_config: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "small-label-encoder-config.json", + ), + label_encoder_tokenizer: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "small-label-encoder-tokenizer.json", + ), + tokenizer: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "small-text-tokenizer.json", + ), + config: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "small-text-gliner-config.json", + ), + } + } + + /// Demonthos GLiNER GGUF base upload. + /// + /// Uses the GGUF weights and sidecar tokenizer/config files from + /// `Demonthos/gliner-gguf`. + pub fn demonthos_base() -> Self { + Self { + model: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "gliner-base.gguf", + ), + label_encoder: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "gliner-base-label-encoder.gguf", + ), + label_encoder_config: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "base-label-encoder-config.json", + ), + label_encoder_tokenizer: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "base-label-encoder-tokenizer.json", + ), + tokenizer: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "base-text-tokenizer.json", + ), + config: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "base-text-gliner-config.json", + ), + } + } + + /// Demonthos GLiNER GGUF large upload. + /// + /// Uses the GGUF weights and sidecar tokenizer/config files from + /// `Demonthos/gliner-gguf`. + pub fn demonthos_large() -> Self { + Self { + model: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "gliner-large.gguf", + ), + label_encoder: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "gliner-large-label-encoder.gguf", + ), + label_encoder_config: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "large-label-encoder-config.json", + ), + label_encoder_tokenizer: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "large-label-encoder-tokenizer.json", + ), + tokenizer: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "large-text-tokenizer.json", + ), + config: Self::huggingface_or_cached( + "Demonthos/gliner-gguf", + "main", + "large-text-gliner-config.json", ), } } From e275069f1e7a1625e9cd957be0f41363b7fe6cf3 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 19:23:02 -0500 Subject: [PATCH 04/34] cli version --- Cargo.lock | 1 + models/rgliner/Cargo.toml | 1 + models/rgliner/examples/relex.rs | 115 +++++++++++++++++++++++-------- 3 files changed, 88 insertions(+), 29 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 4fc1787ba..72d84b8e1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7675,6 +7675,7 @@ name = "rgliner" version = "0.4.0" dependencies = [ "anyhow", + "clap 4.6.0", "fusor", "fusor-core", "fusor-gguf", diff --git a/models/rgliner/Cargo.toml b/models/rgliner/Cargo.toml index c4a8a63d0..6a76abeca 100644 --- a/models/rgliner/Cargo.toml +++ b/models/rgliner/Cargo.toml @@ -28,6 +28,7 @@ pollster = "0.4.0" [dev-dependencies] anyhow.workspace = true tokio = { version = "1", features = ["full"] } +clap = { version = "4", features = ["derive"] } [features] default = ["tokio"] diff --git a/models/rgliner/examples/relex.rs b/models/rgliner/examples/relex.rs index 799e41899..e9e1d09cb 100644 --- a/models/rgliner/examples/relex.rs +++ b/models/rgliner/examples/relex.rs @@ -1,60 +1,117 @@ //! Example of using GlinerRelEx for joint NER and relation extraction. //! -//! Run with: +//! Examples: //! ``` -//! cargo run --example relex -p rgliner +//! cargo run --example relex -p rgliner --release -- \ +//! --text "Apple was founded by Steve Jobs in California." \ +//! --entity-labels person,organization,location \ +//! --relation-labels "founded by,located in" //! ``` //! //! The GGUF file has the tokenizer and GLiNER config baked in as metadata, //! so only the model file path is needed. +use clap::Parser; use rgliner::relex::{GlinerRelEx, GlinerRelExSource}; -use std::env; use std::path::PathBuf; -fn get_model_path() -> Option { - let weights_dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("weights"); - let default_path = weights_dir.join("gliner-relex-multi-v1.0.gguf"); - let model_path = env::var("GLINER_MODEL") - .map(PathBuf::from) - .unwrap_or(default_path); +#[derive(Parser, Debug)] +#[command( + about = "GLiNER-RelEx joint NER and relation extraction", + long_about = None, +)] +struct Args { + /// Input text to analyze. + #[arg(short, long)] + text: String, - if !model_path.exists() { - eprintln!("Model file not found: {:?}", model_path); - return None; + /// Entity labels to detect (comma-separated). + #[arg(short = 'e', long, value_delimiter = ',', required = true)] + entity_labels: Vec, + + /// Relation labels to detect (comma-separated). If empty, only entities are returned. + #[arg(short = 'r', long, value_delimiter = ',', default_value = "")] + relation_labels: Vec, + + /// Path to the GGUF model file. If omitted, uses + /// `/weights/gliner-relex-multi-v1.0.gguf` or `$GLINER_MODEL`. + #[arg(short = 'm', long)] + model: Option, + + /// Minimum confidence for entity detection. + #[arg(long, default_value_t = 0.5)] + entity_threshold: f32, + + /// Minimum confidence for relation classification. + #[arg(long, default_value_t = 0.5)] + relation_threshold: f32, + + /// Maximum adjacency score to keep an entity pair. + #[arg(long, default_value_t = 0.5)] + adjacency_threshold: f32, +} + +fn resolve_model_path(arg: Option) -> anyhow::Result { + if let Some(path) = arg { + return Ok(path); + } + if let Ok(env_path) = std::env::var("GLINER_MODEL") { + return Ok(PathBuf::from(env_path)); } - Some(model_path) + let default = PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("weights") + .join("gliner-relex-multi-v1.0.gguf"); + if default.exists() { + return Ok(default); + } + anyhow::bail!( + "No model path provided. Use --model, set $GLINER_MODEL, or run:\n \ + python scripts/convert_relex_to_gguf.py -m knowledgator/gliner-relex-multi-v1.0 \ + -o weights/gliner-relex-multi-v1.0.gguf" + ) } #[tokio::main] async fn main() -> anyhow::Result<()> { - let model_path = get_model_path().ok_or_else(|| { - anyhow::anyhow!( - "Model file not found. Convert with:\n \ - python scripts/convert_relex_to_gguf.py -m knowledgator/gliner-relex-multi-v1.0 \ - -o weights/gliner-relex-multi-v1.0.gguf" - ) - })?; - - println!("Loading model from: {:?}", model_path); + let args = Args::parse(); + let model_path = resolve_model_path(args.model)?; - let source = GlinerRelExSource::local(model_path); + // `value_delimiter = ','` with `default_value = ""` produces a single empty + // string when the user passes nothing - filter it out. + let entity_labels: Vec<&str> = args + .entity_labels + .iter() + .map(|s| s.trim()) + .filter(|s| !s.is_empty()) + .collect(); + let relation_labels: Vec<&str> = args + .relation_labels + .iter() + .map(|s| s.trim()) + .filter(|s| !s.is_empty()) + .collect(); - let text = "Perfect! Now I have all the information I need to create a comprehensive GGUF inference framework from scratch. Let me create a single-file implementation with full SIMD support targeting nightly Rust"; - let entity_labels = ["technology", "language", "file format"]; - let relation_labels = ["supported by", "implemented with"]; + if entity_labels.is_empty() { + anyhow::bail!("--entity-labels must contain at least one non-empty label"); + } - println!("\nText: {}", text); + println!("Loading model from: {:?}", model_path); + println!("Text: {}", args.text); println!("Entity labels: {:?}", entity_labels); println!("Relation labels: {:?}", relation_labels); + let source = GlinerRelExSource::local(model_path); let relex = GlinerRelEx::builder() .with_source(source) - .with_entity_threshold(0.1) + .with_entity_threshold(args.entity_threshold) + .with_relation_threshold(args.relation_threshold) + .with_adjacency_threshold(args.adjacency_threshold) .build() .await?; - let (entities, relations) = relex.extract(text, &entity_labels, &relation_labels).await?; + let (entities, relations) = relex + .extract(&args.text, &entity_labels, &relation_labels) + .await?; println!("\nEntities found ({}):", entities.len()); for entity in &entities { From 40cd435a5ab6f0adbc7f38ddb4baff0a6b3ea3e7 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 19:57:09 -0500 Subject: [PATCH 05/34] larger models --- .../rgliner/scripts/convert_relex_to_gguf.py | 4 + models/rgliner/src/raw/mdeberta/model.rs | 40 +++- models/rgliner/src/relex.rs | 211 ++++++++++++++++-- models/rgliner/src/relex_tokenization.rs | 36 ++- 4 files changed, 266 insertions(+), 25 deletions(-) diff --git a/models/rgliner/scripts/convert_relex_to_gguf.py b/models/rgliner/scripts/convert_relex_to_gguf.py index e60535bf0..378b76c43 100644 --- a/models/rgliner/scripts/convert_relex_to_gguf.py +++ b/models/rgliner/scripts/convert_relex_to_gguf.py @@ -197,6 +197,10 @@ def map_relex_weight_name(pytorch_name: str) -> str: """Map PyTorch weight names to GGUF conventions for GLiNER-RelEx.""" name = pytorch_name + # ===== Encoder output projection (large variants: 1024 -> 768) ===== + name = name.replace("token_rep_layer.projection.weight", "text.output_proj.weight") + name = name.replace("token_rep_layer.projection.bias", "text.output_proj.bias") + # ===== mDeBERTa Encoder (token_rep_layer.bert_layer.model.*) ===== name = name.replace("token_rep_layer.bert_layer.model.embeddings.word_embeddings.weight", "text.token_embd.weight") name = name.replace("token_rep_layer.bert_layer.model.embeddings.LayerNorm.weight", "text.embd_norm.weight") diff --git a/models/rgliner/src/raw/mdeberta/model.rs b/models/rgliner/src/raw/mdeberta/model.rs index d9a3dd867..0902fcf3c 100644 --- a/models/rgliner/src/raw/mdeberta/model.rs +++ b/models/rgliner/src/raw/mdeberta/model.rs @@ -3,7 +3,7 @@ #[cfg(debug_assertions)] use pollster; -use fusor::layers::{Embedding, LayerNorm}; +use fusor::layers::{Embedding, LayerNorm, Linear}; use fusor::{Device, Result, Tensor, VarBuilder}; use super::attention::RelativePositionEmbedding; @@ -23,6 +23,10 @@ pub struct MDebertaModel { rel_pos_embedding: RelativePositionEmbedding, /// Transformer layers layers: Vec, + /// Optional encoder output projection (used by the `large` variants to + /// map the encoder's 1024-dim hidden state down to the 768-dim space + /// expected by the downstream heads). Absent on base/multi variants. + output_proj: Option>, /// Device device: Device, /// Configuration @@ -64,11 +68,25 @@ impl MDebertaModel { layers.push(layer); } + // Optional post-encoder projection (only present on the large variants, + // which use DeBERTa-v3-large at 1024-dim and project down to 768). + let output_proj = Linear::load(device, &mut vb.pp("output_proj")).ok(); + + #[cfg(debug_assertions)] + if let Some(ref p) = output_proj { + eprintln!( + "[DEBUG] Encoder output projection loaded: {} -> {}", + p.in_features(), + p.out_features() + ); + } + Ok(Self { token_embeddings, embedding_norm, rel_pos_embedding, layers, + output_proj, device: device.clone(), config, }) @@ -147,10 +165,26 @@ impl MDebertaModel { let slice = data.as_slice(); let mean: f32 = slice.iter().sum::() / slice.len() as f32; let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); - eprintln!("[DEBUG] Encoder output: mean={:.6}, std={:.6}", mean, std); + eprintln!("[DEBUG] Encoder output (pre-projection): mean={:.6}, std={:.6}", mean, std); + } + + // Apply optional post-encoder projection (large variants). + if let Some(ref proj) = self.output_proj { + hidden_states = proj.forward(&hidden_states); + + #[cfg(debug_assertions)] + { + let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); + let slice = data.as_slice(); + let mean: f32 = slice.iter().sum::() / slice.len() as f32; + let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + eprintln!( + "[DEBUG] Encoder output (post-projection): mean={:.6}, std={:.6}", + mean, std + ); + } } - // Return last layer output directly (no final LayerNorm on hidden states) hidden_states } diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index 4b45778df..d74474319 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -75,8 +75,10 @@ pub struct GlinerRelExSource { impl GlinerRelExSource { /// GLiNER-RelEx Multi v1.0 source. /// - /// Downloads the GGUF-converted weights from HuggingFace. Tokenizer and - /// config are embedded in the GGUF file. + /// Multilingual variant built on `mdeberta-v3-base` with `span_mode = token_level`. + /// Downloads the GGUF-converted weights from HuggingFace. + /// + /// Tokenizer and config are embedded in the GGUF file. pub fn relex_multi() -> Self { Self { model: FileSource::huggingface( @@ -89,6 +91,24 @@ impl GlinerRelExSource { } } + /// GLiNER-RelEx Base v1.0 source. + /// + /// English-only variant built on `deberta-v3-base` with `span_mode = token_level`. + /// Smaller than the multilingual variant but limited to English text. + /// + /// Tokenizer and config are embedded in the GGUF file. + pub fn relex_base() -> Self { + Self { + model: FileSource::huggingface( + "knowledgator/gliner-relex-base-v1.0-gguf".to_string(), + "main".to_string(), + "gliner-relex-base-v1.0-Q8_0.gguf".to_string(), + ), + tokenizer: None, + config: None, + } + } + /// Create a source from a local GGUF file. /// /// The tokenizer and config are expected to be embedded in the GGUF @@ -217,6 +237,17 @@ impl GlinerRelExBuilder { } } +/// Span-scoring modes supported by the Rust inference path. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SpanMode { + /// Per-token BIO-style output (`[start, end, inside]` sigmoids per (token, label)). + /// Used by the `multi` and `base` variants. + TokenLevel, + /// Per-span scoring: enumerate all spans up to `max_width`, score each against + /// the projected entity prompts. Used by the `large` variants. + MarkerV0, +} + /// GLiNER-RelEx model for joint NER and relation extraction. pub struct GlinerRelEx { /// mDeBERTa encoder @@ -225,8 +256,8 @@ pub struct GlinerRelEx { bilstm: BiLstm, /// Prompt representation layer for label projection prompt_rep_layer: PromptRepLayer, - /// Joint scorer for token-level predictions - scorer: JointScorer, + /// Joint scorer for token-level predictions (None for markerV0 variants). + scorer: Option, /// Span representation layer span_layer: SpanLayer, /// Relations representation layer (adjacency scoring) @@ -237,6 +268,8 @@ pub struct GlinerRelEx { tokenizer: Arc, /// Relation decoder relation_decoder: RelationDecoder, + /// How entities are scored (derived from `gliner.span_mode` metadata). + span_mode: SpanMode, /// Device device: Device, /// Configuration @@ -290,6 +323,24 @@ impl GlinerRelEx { let mut vb = VarBuilder::from_gguf(&mut model_cursor) .map_err(|err| GlinerLoadingError::LoadModel(fusor::Error::from(err)))?; + // Determine span mode to pick the right decoder path. Supported modes + // are `token_level` (base/multi) and `markerV0` (large). + let span_mode_str = vb + .get_metadata("gliner.span_mode") + .and_then(|v| v.to_string().ok()) + .map(|s| s.to_string()) + .unwrap_or_else(|| "token_level".to_string()); + let span_mode = match span_mode_str.as_str() { + "token_level" => SpanMode::TokenLevel, + "markerV0" => SpanMode::MarkerV0, + other => { + return Err(GlinerLoadingError::LoadModel(fusor::Error::msg(format!( + "Unsupported gliner.span_mode '{other}'. \ + Supported values: 'token_level', 'markerV0'." + )))); + } + }; + // Resolve tokenizer: explicit override > embedded metadata. let tokenizer_bytes: Vec = if let Some(tokenizer_src) = source.tokenizer.as_ref() { let tok_label = format!("Tokenizer ({})", tokenizer_src); @@ -316,7 +367,17 @@ impl GlinerRelEx { let tokenizer = Tokenizer::from_bytes(&tokenizer_bytes).map_err(GlinerLoadingError::LoadTokenizer)?; - let relex_tokenizer = RelExTokenizer::with_special_tokens(tokenizer, config.special_tokens.clone()); + // Resolve special tokens from the tokenizer so we pick up the right IDs + // regardless of variant (multi uses 250102/250103/250104, base/large + // use 128001/128002/128003). Falls back to the user-supplied IDs. + let mut effective_config = config; + effective_config.special_tokens = + SpecialTokenIds::from_tokenizer(&tokenizer, effective_config.special_tokens); + let relex_tokenizer = RelExTokenizer::with_special_tokens( + tokenizer, + effective_config.special_tokens.clone(), + ); + let config = effective_config; // Load encoder (mDeBERTa) let encoder = MDebertaModel::load(&device, &mut vb.pp("text"))?; @@ -328,7 +389,10 @@ impl GlinerRelEx { let prompt_rep_layer = PromptRepLayer::load(&device, &mut vb.pp("prompt_rep_layer"))?; // Load joint scorer - let scorer = JointScorer::load(&device, &mut vb.pp("scorer"))?; + let scorer = match span_mode { + SpanMode::TokenLevel => Some(JointScorer::load(&device, &mut vb.pp("scorer"))?), + SpanMode::MarkerV0 => None, + }; // Load span layer let span_layer = SpanLayer::load(&device, &mut vb, config.max_width)?; @@ -357,6 +421,7 @@ impl GlinerRelEx { pair_projector, tokenizer: Arc::new(relex_tokenizer), relation_decoder, + span_mode, device, config, }) @@ -512,17 +577,34 @@ impl GlinerRelEx { text_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max)); } - // 7. Score tokens against entity labels using joint scorer + // 7–8. Decode entities using the mode matching the trained head. let ent_embs_2d: Tensor<2, f32> = ent_embs.squeeze(0).to_concrete(); - let token_scores = self.scorer.forward_entity_scores(&text_embs, &ent_embs_2d).await; - - // 8. Decode entities from token-level scores - let entities = self.decode_entities_from_tokens( - &token_scores, - entity_labels, - &tokenized.word_offsets, - text, - ).await?; + let entities = match self.span_mode { + SpanMode::TokenLevel => { + let scorer = self.scorer.as_ref().expect("token_level requires scorer"); + let token_scores = scorer + .forward_entity_scores(&text_embs, &ent_embs_2d) + .await; + self.decode_entities_from_tokens( + &token_scores, + entity_labels, + &tokenized.word_offsets, + text, + ) + .await? + } + SpanMode::MarkerV0 => { + self.decode_entities_marker_v0( + &text_embs, + &ent_embs_2d, + entity_labels, + &tokenized.word_offsets, + tokenized.num_words, + text, + ) + .await? + } + }; // If no entities or no relation labels, return early if entities.len() < 2 || relation_labels.is_empty() { @@ -642,6 +724,103 @@ impl GlinerRelEx { Ok((entities, relations)) } + /// Decode entities for `span_mode = markerV0` (used by the `large` variants). + /// + /// Enumerates every `(start, end)` pair up to `config.max_width` words, + /// computes the span representation via `SpanLayer::forward_for_spans`, + /// scores each span against every projected entity prompt via a dot + /// product, applies sigmoid + `entity_threshold`, and greedy-filters + /// overlapping spans (keeping the highest-scoring one). + async fn decode_entities_marker_v0( + &self, + text_embs: &Tensor<3, f32>, + ent_embs_2d: &Tensor<2, f32>, + entity_labels: &[&str], + word_offsets: &[(usize, usize)], + num_words: usize, + text: &str, + ) -> Result, GlinerError> { + let threshold = self.config.entity_threshold; + let max_width = self.config.max_width; + let hidden = self.config.hidden_size; + let n_labels = entity_labels.len(); + + if num_words == 0 || n_labels == 0 { + return Ok(Vec::new()); + } + + // Enumerate spans: (start, end) with end-start+1 <= max_width. + let mut spans: Vec<(usize, usize)> = Vec::new(); + for start in 0..num_words { + for width in 1..=max_width.min(num_words - start) { + spans.push((start, start + width - 1)); + } + } + + // Compute span reps [num_spans, hidden]. + let span_reps = self + .span_layer + .forward_for_spans(text_embs, &spans, &self.device); + + // Score: [num_spans, hidden] @ [hidden, n_labels] -> [num_spans, n_labels]. + let label_rep_t: Tensor<2, f32> = ent_embs_2d.transpose(0, 1).to_concrete(); + let logits = span_reps.mat_mul(&label_rep_t); + + let logits_data = logits.clone().as_slice().await?; + let logits_slice = logits_data.as_slice(); + + #[cfg(debug_assertions)] + eprintln!( + "[DEBUG] markerV0 decoding: {} spans x {} labels, threshold={}", + spans.len(), + n_labels, + threshold, + ); + + // Candidate (start, end, label, score) above threshold. + let mut candidates: Vec<(usize, usize, usize, f32)> = Vec::new(); + for (span_idx, &(s, e)) in spans.iter().enumerate() { + for l in 0..n_labels { + let raw = logits_slice[span_idx * n_labels + l]; + let prob = 1.0 / (1.0 + (-raw).exp()); + if prob >= threshold { + candidates.push((s, e, l, prob)); + } + } + } + + // Sort by score descending and greedy non-overlapping filter. + candidates.sort_by(|a, b| b.3.partial_cmp(&a.3).unwrap_or(std::cmp::Ordering::Equal)); + + let mut taken: Vec<(usize, usize)> = Vec::new(); + let mut entities = Vec::new(); + for (s, e, l, score) in candidates { + let overlap = taken.iter().any(|&(a, b)| !(e < a || s > b)); + if overlap { + continue; + } + taken.push((s, e)); + if s < word_offsets.len() && e < word_offsets.len() { + let (start_char, _) = word_offsets[s]; + let (_, end_char) = word_offsets[e]; + entities.push(Entity { + text: text[start_char..end_char].to_string(), + label: entity_labels[l].to_string(), + start_char, + end_char, + start_word: s, + end_word: e, + score, + }); + } + } + + // Ensure output is sorted by score descending for presentation. + entities.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal)); + let _ = hidden; + Ok(entities) + } + /// Decode entities using span-boundary detection with start/end/inside scores. /// /// `token_scores` has shape [batch, seq_len, n_labels, 3] where the last dim diff --git a/models/rgliner/src/relex_tokenization.rs b/models/rgliner/src/relex_tokenization.rs index a82cc59d3..d6bbdf9b8 100644 --- a/models/rgliner/src/relex_tokenization.rs +++ b/models/rgliner/src/relex_tokenization.rs @@ -26,13 +26,37 @@ pub struct SpecialTokenIds { impl Default for SpecialTokenIds { fn default() -> Self { + // Defaults match mdeberta-v3-base (used by gliner-relex-multi-v1.0). + // For other variants (deberta-v3-base/large), use + // `SpecialTokenIds::from_tokenizer` to resolve the IDs dynamically. Self { - cls_id: 1, // [CLS] token - sep_id: 2, // [SEP] token (end-of-sequence) - pad_id: 0, // [PAD] token - ent_id: 250102, // <> token - rel_id: 250104, // <> token - inner_sep_id: 250103, // <> token (internal separator) + cls_id: 1, + sep_id: 2, + pad_id: 0, + ent_id: 250102, + rel_id: 250104, + inner_sep_id: 250103, + } + } +} + +impl SpecialTokenIds { + /// Resolve IDs by querying the tokenizer for each special token. + /// + /// This handles vocab differences between variants (e.g. multi vs base/large + /// where `<>` is id 250102 vs 128001). Falls back to the corresponding + /// field in `fallback` if the tokenizer doesn't contain a particular token. + pub fn from_tokenizer(tokenizer: &tokenizers::Tokenizer, fallback: Self) -> Self { + let lookup = |tok: &str, default: u32| -> u32 { + tokenizer.token_to_id(tok).unwrap_or(default) + }; + Self { + cls_id: lookup("[CLS]", fallback.cls_id), + sep_id: lookup("[SEP]", fallback.sep_id), + pad_id: lookup("[PAD]", fallback.pad_id), + ent_id: lookup("<>", fallback.ent_id), + rel_id: lookup("<>", fallback.rel_id), + inner_sep_id: lookup("<>", fallback.inner_sep_id), } } } From 7dec8b02bf2d04411e446a862cc6c83602654ea2 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 20:08:45 -0500 Subject: [PATCH 06/34] ignore gguf files --- .claude/settings.local.json | 14 +++++++++++++- .gitignore | 1 + models/rgliner/src/relex.rs | 19 +++++++++++++++++++ 3 files changed, 33 insertions(+), 1 deletion(-) diff --git a/.claude/settings.local.json b/.claude/settings.local.json index 5eaa6721c..53905fd06 100644 --- a/.claude/settings.local.json +++ b/.claude/settings.local.json @@ -18,7 +18,19 @@ "Bash(GLINER_MODEL=/Users/evanalmloff/Desktop/Github/ner/models/rgliner/weights/gliner-relex-multi-v1.0.gguf cargo run:*)", "Bash(GLINER_MODEL=./weights/gliner-relex-multi-v1.0.gguf cargo run:*)", "Bash(cargo add:*)", - "Bash(pip install:*)" + "Bash(pip install:*)", + "Bash(mkdir -p /tmp/gguf-backup)", + "Bash(mv models/rgliner/weights/gliner-relex-base-v1.0.gguf /tmp/gguf-backup/)", + "Bash(mv models/rgliner/weights/gliner-relex-large-v0.5.gguf /tmp/gguf-backup/)", + "Bash(mv models/rgliner/weights/gliner-relex-large-v1.0.gguf /tmp/gguf-backup/)", + "Bash(mv models/rgliner/weights/gliner-relex-multi-v1.0.gguf /tmp/gguf-backup/)", + "Bash(mv models/rgliner/weights/gliner-relex-multi-v1.0.gguf /tmp/gguf-backup/gliner-relex-multi-v1.0.gguf.restored)", + "Bash(git filter-repo:*)", + "Bash(git filter-branch:*)", + "Bash(git stash:*)", + "Bash(FILTER_BRANCH_SQUELCH_WARNING=1 git filter-branch -f --index-filter 'git rm --cached --ignore-unmatch models/rgliner/weights/gliner-relex-multi-v1.0.gguf' d6d5c674..HEAD)", + "Bash(git ls-tree:*)", + "Bash(git push:*)" ] } } diff --git a/.gitignore b/.gitignore index 786962a8b..912bb3c21 100644 --- a/.gitignore +++ b/.gitignore @@ -17,3 +17,4 @@ tokenizer.json out.txt todo.md rust-analyzer +*.gguf diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index d74474319..780f3f729 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -109,6 +109,25 @@ impl GlinerRelExSource { } } + /// GLiNER-RelEx Large v1.0 source. + /// + /// English-only variant built on `deberta-v3-large` with `span_mode = markerV0` + /// and a 1024→768 projection between the encoder and downstream heads. + /// The most accurate variant but also the largest. + /// + /// Tokenizer and config are embedded in the GGUF file. + pub fn relex_large() -> Self { + Self { + model: FileSource::huggingface( + "knowledgator/gliner-relex-large-v1.0-gguf".to_string(), + "main".to_string(), + "gliner-relex-large-v1.0-Q8_0.gguf".to_string(), + ), + tokenizer: None, + config: None, + } + } + /// Create a source from a local GGUF file. /// /// The tokenizer and config are expected to be embedded in the GGUF From 1dfb3ec0a5948fbbb1ea29d28ef761271bbf3fde Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 20:45:51 -0500 Subject: [PATCH 07/34] fix formatting --- .../core/src/quantized/matmul/sgemv/mod.rs | 7 +- models/rbert/src/lib.rs | 14 +- models/rbert/src/raw/attention.rs | 5 +- models/rbert/src/raw/self_attention.rs | 7 +- models/rgliner/convert_to_gguf.py | 89 ++++++- .../rgliner/scripts/convert_relex_to_gguf.py | 250 ++++++++++++++++-- models/rgliner/src/lib.rs | 184 ++++++++++--- models/rgliner/src/raw/bilstm.rs | 14 +- models/rgliner/src/raw/joint_scorer.rs | 82 ++++-- models/rgliner/src/raw/mdeberta/attention.rs | 52 +++- models/rgliner/src/raw/mdeberta/config.rs | 10 +- models/rgliner/src/raw/mdeberta/layer.rs | 27 +- models/rgliner/src/raw/mdeberta/model.rs | 44 ++- models/rgliner/src/raw/pair_projector.rs | 4 +- models/rgliner/src/relation_decoding.rs | 21 +- models/rgliner/src/relex.rs | 160 +++++++---- models/rgliner/src/relex_tokenization.rs | 10 +- models/rgliner/src/source.rs | 12 +- 18 files changed, 786 insertions(+), 206 deletions(-) diff --git a/fusor-ml/core/src/quantized/matmul/sgemv/mod.rs b/fusor-ml/core/src/quantized/matmul/sgemv/mod.rs index ddc84c7fd..0300d1dc0 100644 --- a/fusor-ml/core/src/quantized/matmul/sgemv/mod.rs +++ b/fusor-ml/core/src/quantized/matmul/sgemv/mod.rs @@ -179,7 +179,12 @@ pub(crate) fn sgemv( } /// Calculate the number of N-dimension workgroups based on matrix type -pub(crate) fn n_workgroups(op: &QMatMulOperation, matrix: &QMatrix, n: u32, device: &Device) -> u32 { +pub(crate) fn n_workgroups( + op: &QMatMulOperation, + matrix: &QMatrix, + n: u32, + device: &Device, +) -> u32 { // Only use specialized dispatch sizes if we can use specialized SGEMV if use_specialized_sgemv(op, device) { if matrix.datatype == GgmlType::Q6K { diff --git a/models/rbert/src/lib.rs b/models/rbert/src/lib.rs index b68e65630..d6665d733 100644 --- a/models/rbert/src/lib.rs +++ b/models/rbert/src/lib.rs @@ -698,8 +698,11 @@ impl Bert { match &*self.model { EmbeddingModel::Bert(model) => { let token_type_ids = token_ids.zeros_like(); - model.debug_first_layer(&token_ids, &token_type_ids, Some(&attention_mask)) - .ok_or_else(|| BertError::Fusor(fusor::Error::msg("BERT encoder has no layers"))) + model + .debug_first_layer(&token_ids, &token_type_ids, Some(&attention_mask)) + .ok_or_else(|| { + BertError::Fusor(fusor::Error::msg("BERT encoder has no layers")) + }) } EmbeddingModel::Qwen(_) => Err(BertError::Fusor(fusor::Error::msg( "debug_batch_first_layer is only implemented for BERT models", @@ -757,8 +760,11 @@ impl Bert { match &*self.model { EmbeddingModel::Bert(model) => { let token_type_ids = token_ids.zeros_like(); - model.debug_first_layer_attention(&token_ids, &token_type_ids, Some(&attention_mask)) - .ok_or_else(|| BertError::Fusor(fusor::Error::msg("BERT encoder has no layers"))) + model + .debug_first_layer_attention(&token_ids, &token_type_ids, Some(&attention_mask)) + .ok_or_else(|| { + BertError::Fusor(fusor::Error::msg("BERT encoder has no layers")) + }) } EmbeddingModel::Qwen(_) => Err(BertError::Fusor(fusor::Error::msg( "debug_batch_first_layer_attention is only implemented for BERT models", diff --git a/models/rbert/src/raw/attention.rs b/models/rbert/src/raw/attention.rs index 2aeb7ec31..15e0b890f 100644 --- a/models/rbert/src/raw/attention.rs +++ b/models/rbert/src/raw/attention.rs @@ -47,8 +47,9 @@ impl BertAttention { Tensor<3, f32>, ) { let _enter = self.span.enter(); - let (query_layer, key_layer, value_layer, self_outputs) = - self.self_attention.debug_forward(hidden_states, attention_mask); + let (query_layer, key_layer, value_layer, self_outputs) = self + .self_attention + .debug_forward(hidden_states, attention_mask); let attention_output = self.self_output.forward(&self_outputs, hidden_states); ( query_layer, diff --git a/models/rbert/src/raw/self_attention.rs b/models/rbert/src/raw/self_attention.rs index b914da45a..bb6a0faa7 100644 --- a/models/rbert/src/raw/self_attention.rs +++ b/models/rbert/src/raw/self_attention.rs @@ -85,7 +85,12 @@ impl BertSelfAttention { &self, hidden_states: &Tensor<3, f32>, attention_mask: Option<&Tensor<2, u32>>, - ) -> (Tensor<4, f32>, Tensor<4, f32>, Tensor<4, f32>, Tensor<3, f32>) { + ) -> ( + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<4, f32>, + Tensor<3, f32>, + ) { let _enter = self.span.enter(); let query_layer = self.query.forward(hidden_states); let key_layer = self.key.forward(hidden_states); diff --git a/models/rgliner/convert_to_gguf.py b/models/rgliner/convert_to_gguf.py index ba0609fe0..039156189 100644 --- a/models/rgliner/convert_to_gguf.py +++ b/models/rgliner/convert_to_gguf.py @@ -13,8 +13,11 @@ import argparse import json import os +import shutil import struct +import subprocess import sys +import tempfile from pathlib import Path from typing import Any, Dict, List, Tuple @@ -23,6 +26,21 @@ from huggingface_hub import hf_hub_download, snapshot_download +# Quantization targets. In-process ones are packed directly by `gguf.quants.quantize`; +# k-quants are produced by shelling out to `llama-quantize` post-hoc. +IN_PROCESS_QUANTS = { + "f32", "f16", "bf16", + "q4_0", "q4_1", "q5_0", "q5_1", "q8_0", +} +LLAMA_QUANT_TYPES = { + "q2_k", "q3_k", "q3_k_s", "q3_k_m", "q3_k_l", + "q4_k", "q4_k_s", "q4_k_m", + "q5_k", "q5_k_s", "q5_k_m", + "q6_k", +} +QUANT_TYPES = IN_PROCESS_QUANTS | LLAMA_QUANT_TYPES + + # GGUF constants GGUF_MAGIC = 0x46554747 # "GGUF" in little-endian GGUF_VERSION = 3 @@ -54,21 +72,73 @@ GGML_TYPE_BF16 = 30 +def _ggml_type_for(quant: str) -> int: + return { + "f32": GGML_TYPE_F32, + "f16": GGML_TYPE_F16, + "bf16": GGML_TYPE_BF16, + "q4_0": GGML_TYPE_Q4_0, + "q4_1": GGML_TYPE_Q4_1, + "q5_0": GGML_TYPE_Q5_0, + "q5_1": GGML_TYPE_Q5_1, + "q8_0": GGML_TYPE_Q8_0, + }[quant] + + +def _name_for_type(ggml_type: int) -> str: + return { + GGML_TYPE_F32: "F32", + GGML_TYPE_F16: "F16", + GGML_TYPE_BF16: "BF16", + GGML_TYPE_Q4_0: "Q4_0", + GGML_TYPE_Q4_1: "Q4_1", + GGML_TYPE_Q5_0: "Q5_0", + GGML_TYPE_Q5_1: "Q5_1", + GGML_TYPE_Q8_0: "Q8_0", + }.get(ggml_type, f"type={ggml_type}") + + +def _gguf_block_quant(data: np.ndarray, quant: str) -> np.ndarray: + """Return the packed byte array for a block-quantised tensor.""" + import gguf + qtype = { + "q4_0": gguf.GGMLQuantizationType.Q4_0, + "q4_1": gguf.GGMLQuantizationType.Q4_1, + "q5_0": gguf.GGMLQuantizationType.Q5_0, + "q5_1": gguf.GGMLQuantizationType.Q5_1, + "q8_0": gguf.GGMLQuantizationType.Q8_0, + }[quant] + return gguf.quants.quantize(data, qtype) + + +class _U32(int): + """Force a metadata int to be written as GGUF u32 (required by llama-quantize's arch loader).""" + + class GGUFWriter: """Simple GGUF file writer.""" def __init__(self, path: str): self.path = path self.metadata: Dict[str, Any] = {} - self.tensors: List[Tuple[str, np.ndarray, int]] = [] # (name, data, ggml_type) + # (name, data, ggml_type, logical_shape) + self.tensors: List[Tuple[str, np.ndarray, int, Tuple[int, ...]]] = [] def add_metadata(self, key: str, value: Any): """Add metadata key-value pair.""" self.metadata[key] = value - def add_tensor(self, name: str, data: np.ndarray, ggml_type: int = GGML_TYPE_F32): - """Add a tensor.""" - self.tensors.append((name, data, ggml_type)) + def add_tensor( + self, + name: str, + data: np.ndarray, + ggml_type: int = GGML_TYPE_F32, + shape: Tuple[int, ...] = None, + ): + """Add a tensor. `shape` is the logical (un-packed) shape; defaults to data.shape.""" + if shape is None: + shape = tuple(data.shape) + self.tensors.append((name, data, ggml_type, tuple(shape))) def _write_string(self, f, s: str): """Write a GGUF string (length-prefixed UTF-8).""" @@ -81,6 +151,9 @@ def _write_metadata_value(self, f, value: Any): if isinstance(value, bool): f.write(struct.pack(' int: + return { + "f32": GGML_TYPE_F32, + "f16": GGML_TYPE_F16, + "bf16": GGML_TYPE_BF16, + "q4_0": GGML_TYPE_Q4_0, + "q4_1": GGML_TYPE_Q4_1, + "q5_0": GGML_TYPE_Q5_0, + "q5_1": GGML_TYPE_Q5_1, + "q8_0": GGML_TYPE_Q8_0, + }[quant] + + +def _gguf_block_quant(data: np.ndarray, quant: str) -> np.ndarray: + """Return the packed byte array for a block-quantised tensor. + + Uses `gguf.quants.quantize` (Python package from llama.cpp). The inner row + dimension must be a multiple of 32 for q4_0/q5_0/q8_0, else we fall back to + f16 for that tensor (caller decides). + """ + import gguf # Deferred import - only needed for block quants. + qtype = { + "q4_0": gguf.GGMLQuantizationType.Q4_0, + "q4_1": gguf.GGMLQuantizationType.Q4_1, + "q5_0": gguf.GGMLQuantizationType.Q5_0, + "q5_1": gguf.GGMLQuantizationType.Q5_1, + "q8_0": gguf.GGMLQuantizationType.Q8_0, + }[quant] + return gguf.quants.quantize(data, qtype) + + +class _U32(int): + """Force a metadata int to be written as GGUF u32 (needed by llama-quantize for some keys).""" + + class GGUFWriter: """Simple GGUF file writer.""" @@ -60,8 +134,17 @@ def __init__(self, path: str): def add_metadata(self, key: str, value: Any): self.metadata[key] = value - def add_tensor(self, name: str, data: np.ndarray, ggml_type: int = GGML_TYPE_F32): - self.tensors.append((name, data, ggml_type)) + def add_tensor( + self, + name: str, + data: np.ndarray, + ggml_type: int = GGML_TYPE_F32, + shape: Tuple[int, ...] = None, + ): + """Add a tensor. `shape` is the logical (un-packed) shape; defaults to data.shape.""" + if shape is None: + shape = tuple(data.shape) + self.tensors.append((name, data, ggml_type, tuple(shape))) def _write_string(self, f, s: str): encoded = s.encode('utf-8') @@ -72,6 +155,9 @@ def _write_metadata_value(self, f, value: Any): if isinstance(value, bool): f.write(struct.pack('> 16) & 0xFFFF).astype(np.uint16) self._write_string(f, name) - f.write(struct.pack(' str: return name +def _llama_quantize(f32_path: str, output_path: str, quant_type: str) -> None: + """Run `llama-quantize` to convert an f32 GGUF to a k-quant type.""" + binary = shutil.which("llama-quantize") + if binary is None: + raise RuntimeError( + "`llama-quantize` not found in PATH. Install llama.cpp to use k-quant types:\n" + " brew install llama.cpp # macOS\n" + " or build from https://github.com/ggml-org/llama.cpp" + ) + cmd = [binary] + # Keep `scorer.out_mlp.3.weight` at F32 - fusor's quantised matmul produces + # NaN on all-but-first rows when the output dim is tiny (shape is [3, 3072]). + cmd += ["--tensor-type", "scorer.out_mlp.3.weight=f32"] + cmd += [f32_path, output_path, quant_type.upper()] + print(f"\n$ {' '.join(cmd)}") + result = subprocess.run(cmd) + if result.returncode != 0: + raise RuntimeError( + f"llama-quantize failed (exit {result.returncode}) for {quant_type!r}. " + f"See output above. The f32 GGUF was kept at {f32_path} for inspection." + ) + + def convert_relex_to_gguf( model_id: str, output_path: str, @@ -266,19 +375,35 @@ def convert_relex_to_gguf( for name, tensor in state_dict.items(): print(f" {name}: {tensor.shape} {tensor.dtype}") - if quantize == "f32": - ggml_type = GGML_TYPE_F32 - elif quantize == "f16": - ggml_type = GGML_TYPE_F16 - elif quantize == "bf16": - ggml_type = GGML_TYPE_BF16 - else: - raise ValueError(f"Unsupported quantization: {quantize}") + quantize = quantize.lower() + if quantize not in QUANT_TYPES: + raise ValueError( + f"Unsupported quantization: {quantize!r}. Supported: {', '.join(sorted(QUANT_TYPES))}" + ) + post_quant = None + if quantize in LLAMA_QUANT_TYPES: + post_quant = quantize + quantize = "f32" + default_ggml_type = _ggml_type_for(quantize) + + # For k-quant targets, write f32 first to a temp path then shell out. + final_output_path = output_path + if post_quant is not None: + tmp = tempfile.NamedTemporaryFile( + prefix=os.path.basename(output_path).rsplit(".", 1)[0] + ".f32.", + suffix=".gguf", + delete=False, + dir=os.path.dirname(os.path.abspath(output_path)) or None, + ) + output_path = tmp.name + tmp.close() writer = GGUFWriter(output_path) - # Add metadata - writer.add_metadata("general.architecture", "gliner-relex") + # Masquerade as `bert` when we'll feed this to llama-quantize so its loader + # accepts the custom arch. rgliner doesn't read general.architecture. + arch = "bert" if post_quant is not None else "gliner-relex" + writer.add_metadata("general.architecture", arch) writer.add_metadata("general.name", model_id.split("/")[-1]) writer.add_metadata("general.quantization_version", 2) @@ -327,8 +452,26 @@ def convert_relex_to_gguf( writer.add_metadata("gliner.tokenizer_json", tokenizer_json) writer.add_metadata("gliner.config_json", gliner_config_json) + # When we're masquerading as `bert` for llama-quantize, mirror the required + # arch-scoped keys so its loader validates. llama.cpp's BERT loader expects + # u32 (not u64) for these fields; rgliner never reads them. + if post_quant is not None: + writer.add_metadata("bert.context_length", _U32(context_length)) + writer.add_metadata("bert.embedding_length", _U32(hidden_size)) + writer.add_metadata("bert.feed_forward_length", _U32(intermediate_size)) + writer.add_metadata("bert.block_count", _U32(num_layers)) + writer.add_metadata("bert.attention.head_count", _U32(num_heads)) + writer.add_metadata("bert.attention.layer_norm_epsilon", 1e-7) + # Convert tensors - print(f"\nConverting {len(state_dict)} tensors to GGUF...") + print(f"\nConverting {len(state_dict)} tensors to GGUF ({quantize.upper()})...") + + block_quant = quantize.startswith("q") + # Allow debugging by overriding via env var: + # GLINER_QUANT_INCLUDE="text.blk.0.ffn" -> only quantise tensors containing this substring + # GLINER_QUANT_EXCLUDE="token_embd" -> never quantise tensors containing this substring + quant_include = os.environ.get("GLINER_QUANT_INCLUDE") + quant_exclude = os.environ.get("GLINER_QUANT_EXCLUDE") for pytorch_name, tensor in state_dict.items(): gguf_name = map_relex_weight_name(pytorch_name) @@ -343,15 +486,75 @@ def convert_relex_to_gguf( dtype=np.float32 ).reshape(t.shape) - print(f" {pytorch_name} -> {gguf_name} {data.shape}") - writer.add_tensor(gguf_name, data, ggml_type) + logical_shape = tuple(data.shape) + tensor_type = default_ggml_type + + should_quant = block_quant + if should_quant and quant_include is not None and quant_include not in gguf_name: + should_quant = False + if should_quant and quant_exclude is not None and quant_exclude in gguf_name: + should_quant = False + + if should_quant: + # Row-wise block quantisation requires the inner dim to be a + # multiple of 32. 1-D tensors (biases, norms) and non-conforming + # tensors (odd vocab sizes etc.) stay at F32. + # + # Additional constraint: fusor's quantised matmul kernel produces + # NaN for all but the first output row when the output dimension + # is very small (the classifier head `scorer.out_mlp.3.weight` is + # [3, 3072]). Keep these tiny-output tensors at F32 - they're + # negligible bytes anyway. + MIN_OUT_DIM_FOR_QUANT = 32 + if ( + data.ndim >= 2 + and data.shape[-1] % 32 == 0 + and data.shape[0] >= MIN_OUT_DIM_FOR_QUANT + ): + data = _gguf_block_quant(data, quantize) + tensor_type = default_ggml_type + else: + data = np.ascontiguousarray(data, dtype=np.float32) + tensor_type = GGML_TYPE_F32 + else: + data = np.ascontiguousarray(data, dtype=np.float32) + tensor_type = GGML_TYPE_F32 + + print(f" {pytorch_name} -> {gguf_name} {logical_shape} [{_name_for_type(tensor_type)}]") + writer.add_tensor(gguf_name, data, tensor_type, shape=logical_shape) writer.write() print(f"\nOutput: {output_path}") print(f"Size: {os.path.getsize(output_path) / 1024 / 1024:.2f} MB") + + if post_quant is not None: + print(f"\nQuantizing to {post_quant.upper()} via llama-quantize...") + try: + _llama_quantize(output_path, final_output_path, post_quant) + finally: + try: + os.remove(output_path) + except OSError: + pass + print(f"\nFinal output: {final_output_path}") + print(f"Final size: {os.path.getsize(final_output_path) / 1024 / 1024:.2f} MB") + print("\nConversion complete!") +def _name_for_type(ggml_type: int) -> str: + return { + GGML_TYPE_F32: "F32", + GGML_TYPE_F16: "F16", + GGML_TYPE_BF16: "BF16", + GGML_TYPE_Q4_0: "Q4_0", + GGML_TYPE_Q4_1: "Q4_1", + GGML_TYPE_Q5_0: "Q5_0", + GGML_TYPE_Q5_1: "Q5_1", + GGML_TYPE_Q8_0: "Q8_0", + }.get(ggml_type, f"type={ggml_type}") + + def main(): parser = argparse.ArgumentParser(description="Convert GLiNER-RelEx PyTorch models to GGUF") parser.add_argument( @@ -370,8 +573,13 @@ def main(): "--quantize", "-q", type=str, default="f32", - choices=["f32", "f16", "bf16"], - help="Quantization type (default: f32)" + choices=sorted(QUANT_TYPES), + help=( + "Quantization type (default: f32). " + "All quantisation is done in-process via the `gguf` Python package. " + "k-quants (q4_k, q5_k, q6_k) aren't in this list - they are not " + "round-trip compatible with fusor's current quantised tensor reader." + ), ) parser.add_argument( "--cache-dir", diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs index 692dd7cce..57a378387 100644 --- a/models/rgliner/src/lib.rs +++ b/models/rgliner/src/lib.rs @@ -716,8 +716,16 @@ mod gpu_parity_tests { .unwrap(); let single_label = ["organization"]; - let cpu_single_label_embeddings = cpu.label_encoder.encode_labels(&single_label).await.unwrap(); - let gpu_single_label_embeddings = gpu.label_encoder.encode_labels(&single_label).await.unwrap(); + let cpu_single_label_embeddings = cpu + .label_encoder + .encode_labels(&single_label) + .await + .unwrap(); + let gpu_single_label_embeddings = gpu + .label_encoder + .encode_labels(&single_label) + .await + .unwrap(); let _ = print_diff( "single_label_embeddings", &cpu_single_label_embeddings, @@ -749,9 +757,16 @@ mod gpu_parity_tests { .await .unwrap(); - let (cpu_label_states, _) = cpu.label_encoder.debug_sentence_hidden_states(&labels).unwrap(); - let (gpu_label_states, _) = gpu.label_encoder.debug_sentence_hidden_states(&labels).unwrap(); - for (idx, (cpu_state, gpu_state)) in cpu_label_states.iter().zip(&gpu_label_states).enumerate() + let (cpu_label_states, _) = cpu + .label_encoder + .debug_sentence_hidden_states(&labels) + .unwrap(); + let (gpu_label_states, _) = gpu + .label_encoder + .debug_sentence_hidden_states(&labels) + .unwrap(); + for (idx, (cpu_state, gpu_state)) in + cpu_label_states.iter().zip(&gpu_label_states).enumerate() { let name = if idx == 0 { "label_post_embeddings".to_string() @@ -762,9 +777,13 @@ mod gpu_parity_tests { } let (cpu_label_layer0_attention, cpu_label_layer0_intermediate, cpu_label_layer0_output) = - cpu.label_encoder.debug_sentence_first_layer(&labels).unwrap(); + cpu.label_encoder + .debug_sentence_first_layer(&labels) + .unwrap(); let (gpu_label_layer0_attention, gpu_label_layer0_intermediate, gpu_label_layer0_output) = - gpu.label_encoder.debug_sentence_first_layer(&labels).unwrap(); + gpu.label_encoder + .debug_sentence_first_layer(&labels) + .unwrap(); let _ = print_diff( "label_layer0_attention_output", &cpu_label_layer0_attention, @@ -886,9 +905,13 @@ mod gpu_parity_tests { let cpu_label_embeddings = cpu.label_encoder.encode_labels(&labels).await.unwrap(); let gpu_label_embeddings = gpu.label_encoder.encode_labels(&labels).await.unwrap(); - let _ = print_diff("label_embeddings", &cpu_label_embeddings, &gpu_label_embeddings) - .await - .unwrap(); + let _ = print_diff( + "label_embeddings", + &cpu_label_embeddings, + &gpu_label_embeddings, + ) + .await + .unwrap(); let (cpu_token_ids, cpu_attention_mask) = build_text_inputs(&cpu_tokenized, &cpu_device); let (gpu_token_ids, gpu_attention_mask) = build_text_inputs(&gpu_tokenized, &gpu_device); @@ -906,7 +929,8 @@ mod gpu_parity_tests { .text_encoder .debug_hidden_states(&gpu_token_ids, Some(&gpu_attention_mask)); - for (idx, (cpu_state, gpu_state)) in cpu_text_states.iter().zip(&gpu_text_states).enumerate() + for (idx, (cpu_state, gpu_state)) in + cpu_text_states.iter().zip(&gpu_text_states).enumerate() { let name = if idx + 1 == cpu_text_states.len() { "text_final_norm_output".to_string() @@ -983,32 +1007,62 @@ mod gpu_parity_tests { let [b_sz, seq_len, _] = cpu_layer0_qkv_projection.shape(); let cpu_query_states = cpu_layer0_qkv_projection .narrow(2, 0, hidden_size) - .reshape([b_sz, seq_len, text_config.num_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let cpu_key_states = cpu_layer0_qkv_projection .narrow(2, hidden_size, hidden_size) - .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_kv_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let cpu_value_states = cpu_layer0_qkv_projection .narrow(2, 2 * hidden_size, hidden_size) - .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_kv_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let gpu_query_states = gpu_layer0_qkv_projection .narrow(2, 0, hidden_size) - .reshape([b_sz, seq_len, text_config.num_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let gpu_key_states = gpu_layer0_qkv_projection .narrow(2, hidden_size, hidden_size) - .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_kv_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let gpu_value_states = gpu_layer0_qkv_projection .narrow(2, 2 * hidden_size, hidden_size) - .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_kv_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); @@ -1217,32 +1271,62 @@ mod gpu_parity_tests { let cpu_layer2_query_states = cpu_layer2_qkv_projection .narrow(2, 0, hidden_size) - .reshape([b_sz, seq_len, text_config.num_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let cpu_layer2_key_states = cpu_layer2_qkv_projection .narrow(2, hidden_size, hidden_size) - .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_kv_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let cpu_layer2_value_states = cpu_layer2_qkv_projection .narrow(2, 2 * hidden_size, hidden_size) - .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_kv_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let gpu_layer2_query_states = gpu_layer2_qkv_projection .narrow(2, 0, hidden_size) - .reshape([b_sz, seq_len, text_config.num_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let gpu_layer2_key_states = gpu_layer2_qkv_projection .narrow(2, hidden_size, hidden_size) - .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_kv_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); let gpu_layer2_value_states = gpu_layer2_qkv_projection .narrow(2, 2 * hidden_size, hidden_size) - .reshape([b_sz, seq_len, text_config.num_kv_heads, text_config.head_dimension]) + .reshape([ + b_sz, + seq_len, + text_config.num_kv_heads, + text_config.head_dimension, + ]) .transpose(1, 2) .to_concrete(); @@ -1381,10 +1465,12 @@ mod gpu_parity_tests { "ffn_gate_up.weight", &gpu_device, ); - let cpu_layer2_ffn_gate_up_proj = - cpu_layer2_ffn_input.q_mat_mul(&cpu_layer2_ffn_gate_up).to_concrete(); - let gpu_layer2_ffn_gate_up_proj = - gpu_layer2_ffn_input.q_mat_mul(&gpu_layer2_ffn_gate_up).to_concrete(); + let cpu_layer2_ffn_gate_up_proj = cpu_layer2_ffn_input + .q_mat_mul(&cpu_layer2_ffn_gate_up) + .to_concrete(); + let gpu_layer2_ffn_gate_up_proj = gpu_layer2_ffn_input + .q_mat_mul(&gpu_layer2_ffn_gate_up) + .to_concrete(); let _ = print_diff( "text_layer2_ffn_gate_up_projection", &cpu_layer2_ffn_gate_up_proj, @@ -1397,13 +1483,19 @@ mod gpu_parity_tests { let cpu_layer2_gate = cpu_layer2_ffn_gate_up_proj .narrow(2, 0, layer2_intermediate_size) .to_concrete(); - let cpu_layer2_up = cpu_layer2_ffn_gate_up_proj - .narrow(2, layer2_intermediate_size, layer2_intermediate_size); + let cpu_layer2_up = cpu_layer2_ffn_gate_up_proj.narrow( + 2, + layer2_intermediate_size, + layer2_intermediate_size, + ); let gpu_layer2_gate = gpu_layer2_ffn_gate_up_proj .narrow(2, 0, layer2_intermediate_size) .to_concrete(); - let gpu_layer2_up = gpu_layer2_ffn_gate_up_proj - .narrow(2, layer2_intermediate_size, layer2_intermediate_size); + let gpu_layer2_up = gpu_layer2_ffn_gate_up_proj.narrow( + 2, + layer2_intermediate_size, + layer2_intermediate_size, + ); let cpu_layer2_ffn_activated = cpu_layer2_gate.gelu().mul_(&cpu_layer2_up); let gpu_layer2_ffn_activated = gpu_layer2_gate.gelu().mul_(&gpu_layer2_up); let _ = print_diff( @@ -1442,26 +1534,38 @@ mod gpu_parity_tests { .await .unwrap(); - let _ = print_diff("token_embeddings", &cpu_token_embeddings, &gpu_token_embeddings) - .await - .unwrap(); + let _ = print_diff( + "token_embeddings", + &cpu_token_embeddings, + &gpu_token_embeddings, + ) + .await + .unwrap(); let (cpu_word_embeddings, _) = first_subtoken_pooling(&cpu_token_embeddings, &[cpu_tokenized.clone()], &cpu_device); let (gpu_word_embeddings, _) = first_subtoken_pooling(&gpu_token_embeddings, &[gpu_tokenized.clone()], &gpu_device); - let _ = print_diff("word_embeddings", &cpu_word_embeddings, &gpu_word_embeddings) - .await - .unwrap(); + let _ = print_diff( + "word_embeddings", + &cpu_word_embeddings, + &gpu_word_embeddings, + ) + .await + .unwrap(); let (cpu_span_embeddings, cpu_span_indices) = cpu.span_layer.forward(&cpu_word_embeddings, &cpu_device); let (gpu_span_embeddings, gpu_span_indices) = gpu.span_layer.forward(&gpu_word_embeddings, &gpu_device); assert_eq!(cpu_span_indices, gpu_span_indices); - let _ = print_diff("span_embeddings", &cpu_span_embeddings, &gpu_span_embeddings) - .await - .unwrap(); + let _ = print_diff( + "span_embeddings", + &cpu_span_embeddings, + &gpu_span_embeddings, + ) + .await + .unwrap(); let cpu_scores = Scorer::forward(&cpu_span_embeddings, &cpu_label_embeddings); let gpu_scores = Scorer::forward(&gpu_span_embeddings, &gpu_label_embeddings); diff --git a/models/rgliner/src/raw/bilstm.rs b/models/rgliner/src/raw/bilstm.rs index f7283956e..85b0b6125 100644 --- a/models/rgliner/src/raw/bilstm.rs +++ b/models/rgliner/src/raw/bilstm.rs @@ -11,10 +11,10 @@ use fusor::{Device, Result, Tensor, VarBuilder}; /// and concatenates the outputs. pub struct BiLstm { // Forward LSTM weights - weight_ih_f: Tensor<2, f32>, // [4*hidden, input_size] - weight_hh_f: Tensor<2, f32>, // [4*hidden, hidden_size] - bias_ih_f: Tensor<1, f32>, // [4*hidden] - bias_hh_f: Tensor<1, f32>, // [4*hidden] + weight_ih_f: Tensor<2, f32>, // [4*hidden, input_size] + weight_hh_f: Tensor<2, f32>, // [4*hidden, hidden_size] + bias_ih_f: Tensor<1, f32>, // [4*hidden] + bias_hh_f: Tensor<1, f32>, // [4*hidden] // Backward LSTM weights weight_ih_b: Tensor<2, f32>, weight_hh_b: Tensor<2, f32>, @@ -191,7 +191,11 @@ impl BiLstm { } // Store output in correct position - let store_pos = if reverse { seq_len - 1 - out_idx } else { out_idx }; + let store_pos = if reverse { + seq_len - 1 - out_idx + } else { + out_idx + }; for i in 0..hidden_size { outputs[store_pos * hidden_size + i] = h[i]; } diff --git a/models/rgliner/src/raw/joint_scorer.rs b/models/rgliner/src/raw/joint_scorer.rs index 6ff22ac90..0c621ca94 100644 --- a/models/rgliner/src/raw/joint_scorer.rs +++ b/models/rgliner/src/raw/joint_scorer.rs @@ -17,7 +17,7 @@ use fusor::{Device, Result, Tensor, VarBuilder}; /// raw token embeddings with projected labels for the MLP input. pub struct JointScorer { #[allow(dead_code)] - proj_token: Linear, // Not used in main scoring path + proj_token: Linear, // Not used in main scoring path proj_label: Linear, out_fc1: Linear, out_fc2: Linear, @@ -34,14 +34,29 @@ impl JointScorer { #[cfg(debug_assertions)] { eprintln!("[DEBUG] JointScorer loaded:"); - eprintln!(" proj_label: in={}, out={}", proj_label.in_features(), proj_label.out_features()); - eprintln!(" out_fc1: in={}, out={}", out_fc1.in_features(), out_fc1.out_features()); - eprintln!(" out_fc2: in={}, out={}", out_fc2.in_features(), out_fc2.out_features()); + eprintln!( + " proj_label: in={}, out={}", + proj_label.in_features(), + proj_label.out_features() + ); + eprintln!( + " out_fc1: in={}, out={}", + out_fc1.in_features(), + out_fc1.out_features() + ); + eprintln!( + " out_fc2: in={}, out={}", + out_fc2.in_features(), + out_fc2.out_features() + ); // Print fc2 bias values (these are the biases for O, B, I classes) if let Some(bias) = out_fc2.bias() { let bias_data = pollster::block_on(bias.clone().as_slice()).unwrap(); let b = bias_data.as_slice(); - eprintln!(" out_fc2 bias: O={:.6}, B={:.6}, I={:.6}", b[0], b[1], b[2]); + eprintln!( + " out_fc2 bias: O={:.6}, B={:.6}, I={:.6}", + b[0], b[1], b[2] + ); } } @@ -77,8 +92,10 @@ impl JointScorer { let [n_labels, _] = label_embs.shape(); #[cfg(debug_assertions)] - eprintln!("[DEBUG] scorer.forward: batch={}, seq_len={}, hidden_dim={}, n_labels={}", - batch_size, seq_len, hidden_dim, n_labels); + eprintln!( + "[DEBUG] scorer.forward: batch={}, seq_len={}, hidden_dim={}, n_labels={}", + batch_size, seq_len, hidden_dim, n_labels + ); // Project both token and label embeddings // token: [batch, seq, hidden] -> [batch, seq, hidden*2] @@ -94,8 +111,14 @@ impl JointScorer { let output_data = proj_tokens.clone().as_slice().await.unwrap(); let output_slice = output_data.as_slice(); eprintln!("[DEBUG] proj_token input[0,0,:5]: {:?}", &input_slice[0..5]); - eprintln!("[DEBUG] proj_token output[0,0,:5]: {:?}", &output_slice[0..5]); - eprintln!("[DEBUG] proj_token output[0,0,768:773]: {:?}", &output_slice[768..773]); + eprintln!( + "[DEBUG] proj_token output[0,0,:5]: {:?}", + &output_slice[0..5] + ); + eprintln!( + "[DEBUG] proj_token output[0,0,768:773]: {:?}", + &output_slice[768..773] + ); } // label: [n_labels, hidden] -> [n_labels, hidden*2] @@ -104,15 +127,20 @@ impl JointScorer { let proj_labels: Tensor<2, f32> = proj_labels.squeeze(0).to_concrete(); #[cfg(debug_assertions)] - eprintln!("[DEBUG] proj_tokens shape: [{}, {}, {}], proj_labels shape: [{}, {}], half_proj={}", - batch_size, seq_len, proj_dim, n_labels, proj_dim, half_proj); + eprintln!( + "[DEBUG] proj_tokens shape: [{}, {}, {}], proj_labels shape: [{}, {}], half_proj={}", + batch_size, seq_len, proj_dim, n_labels, proj_dim, half_proj + ); // Split and combine: token_first + label_first + (token_second * label_second) // MLP input dimension = half_proj + half_proj + half_proj = 3 * half_proj let mlp_input_dim = 3 * half_proj; #[cfg(debug_assertions)] - eprintln!("[DEBUG] mlp_input_dim={} (3 * {})", mlp_input_dim, half_proj); + eprintln!( + "[DEBUG] mlp_input_dim={} (3 * {})", + mlp_input_dim, half_proj + ); // Get raw data slices (without expansion - we'll handle broadcast manually) // proj_tokens shape: [batch, seq, proj_dim] @@ -120,8 +148,8 @@ impl JointScorer { let tokens_data = proj_tokens.clone().as_slice().await.unwrap(); let labels_data = proj_labels.clone().as_slice().await.unwrap(); - let tokens_slice = tokens_data.as_slice(); // [batch * seq * proj_dim] - let labels_slice = labels_data.as_slice(); // [n_labels * proj_dim] + let tokens_slice = tokens_data.as_slice(); // [batch * seq * proj_dim] + let labels_slice = labels_data.as_slice(); // [n_labels * proj_dim] #[cfg(debug_assertions)] { @@ -138,7 +166,9 @@ impl JointScorer { for t in 0..5.min(seq_len) { let start = t * proj_dim; let vals: Vec = (0..5).map(|i| tokens_slice[start + i]).collect(); - let vals_second: Vec = (0..5).map(|i| tokens_slice[start + half_proj + i]).collect(); + let vals_second: Vec = (0..5) + .map(|i| tokens_slice[start + half_proj + i]) + .collect(); eprintln!(" token {}: first={:?}, second={:?}", t, vals, vals_second); } } @@ -216,8 +246,14 @@ impl JointScorer { for s in 0..3.min(seq_len) { for l in 0..n_labels { let idx = s * n_labels * num_classes + l * num_classes; - eprintln!(" token {} label {}: start={:.4}, end={:.4}, inside={:.4}", - s, l, data[idx], data[idx+1], data[idx+2]); + eprintln!( + " token {} label {}: start={:.4}, end={:.4}, inside={:.4}", + s, + l, + data[idx], + data[idx + 1], + data[idx + 2] + ); } } } @@ -250,8 +286,16 @@ impl PromptRepLayer { #[cfg(debug_assertions)] { eprintln!("[DEBUG] PromptRepLayer loaded:"); - eprintln!(" fc1: in={}, out={}", fc1.in_features(), fc1.out_features()); - eprintln!(" fc2: in={}, out={}", fc2.in_features(), fc2.out_features()); + eprintln!( + " fc1: in={}, out={}", + fc1.in_features(), + fc1.out_features() + ); + eprintln!( + " fc2: in={}, out={}", + fc2.in_features(), + fc2.out_features() + ); } Ok(Self { fc1, fc2 }) diff --git a/models/rgliner/src/raw/mdeberta/attention.rs b/models/rgliner/src/raw/mdeberta/attention.rs index b698e2d1c..3b07cb3e5 100644 --- a/models/rgliner/src/raw/mdeberta/attention.rs +++ b/models/rgliner/src/raw/mdeberta/attention.rs @@ -36,7 +36,10 @@ impl RelativePositionEmbedding { let embeddings = if dim0 > dim1 { // Shape is [hidden_size, positions] - need to transpose #[cfg(debug_assertions)] - eprintln!("[DEBUG] Transposing rel_pos_embd from [{}, {}] to [{}, {}]", dim0, dim1, dim1, dim0); + eprintln!( + "[DEBUG] Transposing rel_pos_embd from [{}, {}] to [{}, {}]", + dim0, dim1, dim1, dim0 + ); embeddings_raw.transpose(0, 1).to_concrete() } else { embeddings_raw @@ -56,7 +59,13 @@ impl RelativePositionEmbedding { /// Apply log-bucket position encoding (matches Python make_log_bucket_position). /// positions close to 0 use linear indexing, far positions are log-bucketed. fn make_log_bucket_position(rel_pos: i32, bucket_size: i32, max_position: i32) -> i32 { - let sign = if rel_pos > 0 { 1 } else if rel_pos < 0 { -1 } else { 0 }; + let sign = if rel_pos > 0 { + 1 + } else if rel_pos < 0 { + -1 + } else { + 0 + }; let mid = bucket_size / 2; let abs_pos = if rel_pos < mid && rel_pos > -mid { mid - 1 @@ -108,7 +117,9 @@ impl RelativePositionEmbedding { } } - Tensor::new(device, &indices).reshape([seq_len, seq_len]).to_concrete() + Tensor::new(device, &indices) + .reshape([seq_len, seq_len]) + .to_concrete() } /// Get the raw relative position embedding table (normalized). @@ -141,7 +152,9 @@ impl RelativePositionEmbedding { let gathered = normalized_embeddings.index_select(0, &flat_indices); // Reshape back to [seq_len, seq_len, hidden_size] - gathered.reshape([seq_len, seq_len, hidden_size]).to_concrete() + gathered + .reshape([seq_len, seq_len, hidden_size]) + .to_concrete() } /// Get the maximum relative positions setting. @@ -418,9 +431,18 @@ impl MDebertaAttention { let key = self.key.forward(hidden_states); let value = self.value.forward(hidden_states); - let query = query.reshape([b_sz, seq_len, self.num_heads, self.head_dim]).transpose(1, 2).to_concrete(); - let key = key.reshape([b_sz, seq_len, self.num_heads, self.head_dim]).transpose(1, 2).to_concrete(); - let value = value.reshape([b_sz, seq_len, self.num_heads, self.head_dim]).transpose(1, 2).to_concrete(); + let query = query + .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + let key = key + .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); + let value = value + .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) + .transpose(1, 2) + .to_concrete(); let c2c_scores = query.mat_mul(&key.transpose(2, 3)); let attn_scores = c2c_scores.mul_scalar(1.0 / (self.head_dim as f32).sqrt()); @@ -440,7 +462,11 @@ impl MDebertaAttention { let attn_probs = attn_scores.softmax_last_dim::<3>(); let context = attn_probs.mat_mul(&value); - let context = context.transpose(1, 2).to_concrete().reshape([b_sz, seq_len, hidden_size]).to_concrete(); + let context = context + .transpose(1, 2) + .to_concrete() + .reshape([b_sz, seq_len, hidden_size]) + .to_concrete(); self.output.forward(&context) } } @@ -469,7 +495,12 @@ impl DisentangledSelfAttention { rel_pos_indices: &Tensor<2, u32>, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { - self.attention.forward_with_indices(hidden_states, rel_pos_emb, rel_pos_indices, attention_mask) + self.attention.forward_with_indices( + hidden_states, + rel_pos_emb, + rel_pos_indices, + attention_mask, + ) } pub fn forward( @@ -478,6 +509,7 @@ impl DisentangledSelfAttention { rel_pos_emb: Option<&Tensor<3, f32>>, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { - self.attention.forward(hidden_states, rel_pos_emb, attention_mask) + self.attention + .forward(hidden_states, rel_pos_emb, attention_mask) } } diff --git a/models/rgliner/src/raw/mdeberta/config.rs b/models/rgliner/src/raw/mdeberta/config.rs index 86578db39..fce51a18b 100644 --- a/models/rgliner/src/raw/mdeberta/config.rs +++ b/models/rgliner/src/raw/mdeberta/config.rs @@ -46,14 +46,16 @@ impl MDebertaConfig { let num_layers = vb .get_metadata("gliner.block_count") .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: gliner.block_count"))? - as usize; + .ok_or_else(|| { + fusor::Error::msg("Missing required GGUF metadata: gliner.block_count") + })? as usize; let hidden_size = vb .get_metadata("gliner.embedding_length") .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: gliner.embedding_length"))? - as usize; + .ok_or_else(|| { + fusor::Error::msg("Missing required GGUF metadata: gliner.embedding_length") + })? as usize; if hidden_size % num_heads != 0 { return Err(fusor::Error::msg(format!( diff --git a/models/rgliner/src/raw/mdeberta/layer.rs b/models/rgliner/src/raw/mdeberta/layer.rs index 554bb934e..ed6d27311 100644 --- a/models/rgliner/src/raw/mdeberta/layer.rs +++ b/models/rgliner/src/raw/mdeberta/layer.rs @@ -28,12 +28,8 @@ impl MDebertaLayer { head_dim: usize, eps: f32, ) -> Result { - let attention = DisentangledSelfAttention::load( - device, - &mut vb.pp("attention"), - num_heads, - head_dim, - )?; + let attention = + DisentangledSelfAttention::load(device, &mut vb.pp("attention"), num_heads, head_dim)?; let attention_norm = LayerNorm::load(device, &mut vb.pp("attention_norm"), eps)?; let feed_forward = MDebertaFeedForward::load(device, &mut vb.pp("ffn"))?; let output_norm = LayerNorm::load(device, &mut vb.pp("output_norm"), eps)?; @@ -61,8 +57,15 @@ impl MDebertaLayer { attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { // Self-attention + residual + norm - let attn_output = self.attention.forward_with_rel(hidden_states, rel_pos_emb, rel_pos_indices, attention_mask); - let hidden_states = self.attention_norm.forward(&hidden_states.add_(&attn_output)); + let attn_output = self.attention.forward_with_rel( + hidden_states, + rel_pos_emb, + rel_pos_indices, + attention_mask, + ); + let hidden_states = self + .attention_norm + .forward(&hidden_states.add_(&attn_output)); // FFN + residual + norm let ffn_output = self.feed_forward.forward(&hidden_states); @@ -77,8 +80,12 @@ impl MDebertaLayer { attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { // Self-attention + residual + norm - let attn_output = self.attention.forward(hidden_states, rel_pos_emb, attention_mask); - let hidden_states = self.attention_norm.forward(&hidden_states.add_(&attn_output)); + let attn_output = self + .attention + .forward(hidden_states, rel_pos_emb, attention_mask); + let hidden_states = self + .attention_norm + .forward(&hidden_states.add_(&attn_output)); // FFN + residual + norm let ffn_output = self.feed_forward.forward(&hidden_states); diff --git a/models/rgliner/src/raw/mdeberta/model.rs b/models/rgliner/src/raw/mdeberta/model.rs index 0902fcf3c..115889ebc 100644 --- a/models/rgliner/src/raw/mdeberta/model.rs +++ b/models/rgliner/src/raw/mdeberta/model.rs @@ -115,8 +115,12 @@ impl MDebertaModel { let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); let slice = data.as_slice(); let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); - eprintln!("[DEBUG] After token_embeddings: mean={:.6}, std={:.6}", mean, std); + let std: f32 = + (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + eprintln!( + "[DEBUG] After token_embeddings: mean={:.6}, std={:.6}", + mean, std + ); } // Apply embedding LayerNorm @@ -128,8 +132,12 @@ impl MDebertaModel { let slice = data.as_slice(); let hidden_size = self.config.hidden_size; let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); - eprintln!("[DEBUG] After embedding_norm: mean={:.6}, std={:.6}", mean, std); + let std: f32 = + (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + eprintln!( + "[DEBUG] After embedding_norm: mean={:.6}, std={:.6}", + mean, std + ); // Print raw embeddings at <> positions (1, 3, 5) and others eprintln!("[DEBUG] Raw embeddings at positions (first 5 values):"); for pos in [0, 1, 2, 3, 4, 5, 10, 17] { @@ -142,20 +150,28 @@ impl MDebertaModel { } // Compute relative position indices and get embedding table - let rel_indices = self.rel_pos_embedding.compute_relative_indices(seq_len, &self.device); + let rel_indices = self + .rel_pos_embedding + .compute_relative_indices(seq_len, &self.device); let rel_pos_emb = self.rel_pos_embedding.get_embeddings(); // Pass through transformer layers with proper position attention for (i, layer) in self.layers.iter().enumerate() { - hidden_states = layer.forward_with_rel(&hidden_states, &rel_pos_emb, &rel_indices, attention_mask); + hidden_states = + layer.forward_with_rel(&hidden_states, &rel_pos_emb, &rel_indices, attention_mask); #[cfg(debug_assertions)] if i == 0 || i == 11 { let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); let slice = data.as_slice(); let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); - eprintln!("[DEBUG] After layer {}: mean={:.6}, std={:.6}", i, mean, std); + let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() + / slice.len() as f32) + .sqrt(); + eprintln!( + "[DEBUG] After layer {}: mean={:.6}, std={:.6}", + i, mean, std + ); } } @@ -164,8 +180,12 @@ impl MDebertaModel { let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); let slice = data.as_slice(); let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); - eprintln!("[DEBUG] Encoder output (pre-projection): mean={:.6}, std={:.6}", mean, std); + let std: f32 = + (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + eprintln!( + "[DEBUG] Encoder output (pre-projection): mean={:.6}, std={:.6}", + mean, std + ); } // Apply optional post-encoder projection (large variants). @@ -177,7 +197,9 @@ impl MDebertaModel { let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); let slice = data.as_slice(); let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); + let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() + / slice.len() as f32) + .sqrt(); eprintln!( "[DEBUG] Encoder output (post-projection): mean={:.6}, std={:.6}", mean, std diff --git a/models/rgliner/src/raw/pair_projector.rs b/models/rgliner/src/raw/pair_projector.rs index a43dee7dd..91f5cb18f 100644 --- a/models/rgliner/src/raw/pair_projector.rs +++ b/models/rgliner/src/raw/pair_projector.rs @@ -136,6 +136,8 @@ impl RelationScorer { let scores = flat_pairs.mat_mul(&rel_t); // Reshape back: [batch, num_pairs, num_relations] - scores.reshape([batch_size, num_pairs, num_relations]).to_concrete() + scores + .reshape([batch_size, num_pairs, num_relations]) + .to_concrete() } } diff --git a/models/rgliner/src/relation_decoding.rs b/models/rgliner/src/relation_decoding.rs index b2c11f481..e191f4e75 100644 --- a/models/rgliner/src/relation_decoding.rs +++ b/models/rgliner/src/relation_decoding.rs @@ -131,7 +131,11 @@ impl RelationDecoder { } // Sort by score descending - relations.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal)); + relations.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); relations } @@ -144,7 +148,11 @@ impl RelationDecoder { /// /// # Returns /// Vector of (head_idx, tail_idx) pairs above the adjacency threshold - pub fn filter_pairs(&self, adjacency_scores: &[f32], num_entities: usize) -> Vec<(usize, usize)> { + pub fn filter_pairs( + &self, + adjacency_scores: &[f32], + num_entities: usize, + ) -> Vec<(usize, usize)> { let mut pairs = Vec::new(); for i in 0..num_entities { @@ -250,8 +258,7 @@ mod tests { #[test] fn test_decode_relations() { - let decoder = RelationDecoder::new() - .with_relation_threshold(0.7); + let decoder = RelationDecoder::new().with_relation_threshold(0.7); let entities = vec![ make_entity(0, 1, "organization"), @@ -264,9 +271,9 @@ mod tests { // Relation scores: 3 pairs x 2 relations let relation_scores = vec![ - 0.85, 0.3, // pair (0,1): "founded by" = 0.85, "located in" = 0.3 - 0.2, 0.9, // pair (0,2): "founded by" = 0.2, "located in" = 0.9 - 0.1, 0.4, // pair (1,2): both below threshold + 0.85, 0.3, // pair (0,1): "founded by" = 0.85, "located in" = 0.3 + 0.2, 0.9, // pair (0,2): "founded by" = 0.2, "located in" = 0.9 + 0.1, 0.4, // pair (1,2): both below threshold ]; let relation_labels = &["founded by", "located in"]; diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index 780f3f729..d22e8127f 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -51,7 +51,9 @@ use tokenizers::Tokenizer; use crate::decoding::Entity; use crate::error::{GlinerError, GlinerLoadingError}; use crate::raw::mdeberta::MDebertaModel; -use crate::raw::{BiLstm, JointScorer, PairProjector, PromptRepLayer, RelationsRepLayer, SpanLayer}; +use crate::raw::{ + BiLstm, JointScorer, PairProjector, PromptRepLayer, RelationsRepLayer, SpanLayer, +}; use crate::relation_decoding::{Relation, RelationDecoder, RelationDecoderConfig}; use crate::relex_tokenization::{RelExTokenizer, SpecialTokenIds}; @@ -392,10 +394,8 @@ impl GlinerRelEx { let mut effective_config = config; effective_config.special_tokens = SpecialTokenIds::from_tokenizer(&tokenizer, effective_config.special_tokens); - let relex_tokenizer = RelExTokenizer::with_special_tokens( - tokenizer, - effective_config.special_tokens.clone(), - ); + let relex_tokenizer = + RelExTokenizer::with_special_tokens(tokenizer, effective_config.special_tokens.clone()); let config = effective_config; // Load encoder (mDeBERTa) @@ -462,17 +462,24 @@ impl GlinerRelEx { relation_labels: &[&str], ) -> Result<(Vec, Vec), GlinerError> { // 1. Tokenize with special tokens - let tokenized = self.tokenizer.tokenize(text, entity_labels, relation_labels)?; + let tokenized = self + .tokenizer + .tokenize(text, entity_labels, relation_labels)?; #[cfg(debug_assertions)] { - eprintln!("[DEBUG] Tokenized: {} tokens, {} words", - tokenized.token_ids.len(), tokenized.num_words); + eprintln!( + "[DEBUG] Tokenized: {} tokens, {} words", + tokenized.token_ids.len(), + tokenized.num_words + ); eprintln!("[DEBUG] ent_positions: {:?}", tokenized.ent_positions); eprintln!("[DEBUG] rel_positions: {:?}", tokenized.rel_positions); eprintln!("[DEBUG] text_positions: {:?}", tokenized.text_positions); - eprintln!("[DEBUG] token_ids (first 20): {:?}", - &tokenized.token_ids[..20.min(tokenized.token_ids.len())]); + eprintln!( + "[DEBUG] token_ids (first 20): {:?}", + &tokenized.token_ids[..20.min(tokenized.token_ids.len())] + ); } if tokenized.num_words == 0 { @@ -494,11 +501,15 @@ impl GlinerRelEx { let enc_data = encoder_output.clone().as_slice().await.unwrap(); let enc_slice = enc_data.as_slice(); let mean: f32 = enc_slice.iter().sum::() / enc_slice.len() as f32; - let variance: f32 = enc_slice.iter().map(|x| (x - mean).powi(2)).sum::() / enc_slice.len() as f32; - eprintln!("[DEBUG] Encoder output stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", - mean, variance, - enc_slice.iter().cloned().fold(f32::INFINITY, f32::min), - enc_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max)); + let variance: f32 = + enc_slice.iter().map(|x| (x - mean).powi(2)).sum::() / enc_slice.len() as f32; + eprintln!( + "[DEBUG] Encoder output stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", + mean, + variance, + enc_slice.iter().cloned().fold(f32::INFINITY, f32::min), + enc_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max) + ); // Check encoder output at specific positions (for <> tokens) let hidden_size = self.config.hidden_size; @@ -521,7 +532,8 @@ impl GlinerRelEx { // 4. Extract word-level embeddings from encoder output, THEN apply BiLSTM // (Python applies BiLSTM to word-level embeddings, not the full token sequence.) - let word_encoder_embs = self.gather_at_positions(&encoder_output, &tokenized.text_positions); + let word_encoder_embs = + self.gather_at_positions(&encoder_output, &tokenized.text_positions); let lstm_output = self.bilstm.forward(&word_encoder_embs).await; #[cfg(debug_assertions)] @@ -551,7 +563,10 @@ impl GlinerRelEx { for l in 0..tokenized.ent_positions.len() { let start = l * hidden_size; let vals: Vec = (0..5).map(|i| raw_slice[start + i]).collect(); - eprintln!(" label {} (pos {}): {:?}", l, tokenized.ent_positions[l], vals); + eprintln!( + " label {} (pos {}): {:?}", + l, tokenized.ent_positions[l], vals + ); } } @@ -563,11 +578,15 @@ impl GlinerRelEx { let ent_slice = ent_data.as_slice(); let hidden_size = self.config.hidden_size; let mean: f32 = ent_slice.iter().sum::() / ent_slice.len() as f32; - let variance: f32 = ent_slice.iter().map(|x| (x - mean).powi(2)).sum::() / ent_slice.len() as f32; - eprintln!("[DEBUG] Entity label embs stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", - mean, variance, - ent_slice.iter().cloned().fold(f32::INFINITY, f32::min), - ent_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max)); + let variance: f32 = + ent_slice.iter().map(|x| (x - mean).powi(2)).sum::() / ent_slice.len() as f32; + eprintln!( + "[DEBUG] Entity label embs stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", + mean, + variance, + ent_slice.iter().cloned().fold(f32::INFINITY, f32::min), + ent_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max) + ); // Print projected values per label (compare with Python) eprintln!("[DEBUG] Projected entity embeddings (first 5 values per label):"); for l in 0..entity_labels.len() { @@ -589,11 +608,15 @@ impl GlinerRelEx { let text_data = text_embs.clone().as_slice().await.unwrap(); let text_slice = text_data.as_slice(); let mean: f32 = text_slice.iter().sum::() / text_slice.len() as f32; - let variance: f32 = text_slice.iter().map(|x| (x - mean).powi(2)).sum::() / text_slice.len() as f32; - eprintln!("[DEBUG] Text token embs stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", - mean, variance, - text_slice.iter().cloned().fold(f32::INFINITY, f32::min), - text_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max)); + let variance: f32 = text_slice.iter().map(|x| (x - mean).powi(2)).sum::() + / text_slice.len() as f32; + eprintln!( + "[DEBUG] Text token embs stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", + mean, + variance, + text_slice.iter().cloned().fold(f32::INFINITY, f32::min), + text_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max) + ); } // 7–8. Decode entities using the mode matching the trained head. @@ -601,9 +624,7 @@ impl GlinerRelEx { let entities = match self.span_mode { SpanMode::TokenLevel => { let scorer = self.scorer.as_ref().expect("token_level requires scorer"); - let token_scores = scorer - .forward_entity_scores(&text_embs, &ent_embs_2d) - .await; + let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs_2d).await; self.decode_entities_from_tokens( &token_scores, entity_labels, @@ -636,7 +657,9 @@ impl GlinerRelEx { .iter() .map(|e| (e.start_word, e.end_word)) .collect(); - let span_reps = self.span_layer.forward_for_spans(&text_embs, &entity_spans, &self.device); + let span_reps = self + .span_layer + .forward_for_spans(&text_embs, &entity_spans, &self.device); // span_reps shape: [num_entities, hidden] #[cfg(debug_assertions)] @@ -648,7 +671,10 @@ impl GlinerRelEx { for (i, e) in entities.iter().enumerate() { let start = i * hidden; let vals: Vec = (0..5).map(|k| sr[start + k]).collect(); - eprintln!(" {} ({}, {}): {:?}", e.text, e.start_word, e.end_word, vals); + eprintln!( + " {} ({}, {}): {:?}", + e.text, e.start_word, e.end_word, vals + ); } // Print rel_embs let re_data = rel_embs.clone().as_slice().await?; @@ -710,14 +736,20 @@ impl GlinerRelEx { { eprintln!( "[DEBUG] Relation scoring: {} pairs, {} relations, threshold={}", - candidate_pairs.len(), n_rels, threshold); + candidate_pairs.len(), + n_rels, + threshold + ); for (pair_idx, &(h, t)) in candidate_pairs.iter().enumerate().take(6) { let base = pair_idx * n_rels; - let raw: Vec = (0..n_rels).map(|c| rel_scores_slice.as_slice()[base + c]).collect(); + let raw: Vec = (0..n_rels) + .map(|c| rel_scores_slice.as_slice()[base + c]) + .collect(); let sig: Vec = raw.iter().map(|x| 1.0 / (1.0 + (-x).exp())).collect(); eprintln!( " pair ({}->{}) [{} -> {}]: raw={:?}, sig={:?}", - h, t, entities[h].text, entities[t].text, raw, sig); + h, t, entities[h].text, entities[t].text, raw, sig + ); } } @@ -738,7 +770,11 @@ impl GlinerRelEx { } // Sort by score descending - relations.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal)); + relations.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); Ok((entities, relations)) } @@ -835,7 +871,11 @@ impl GlinerRelEx { } // Ensure output is sorted by score descending for presentation. - entities.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal)); + entities.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); let _ = hidden; Ok(entities) } @@ -860,16 +900,23 @@ impl GlinerRelEx { #[cfg(debug_assertions)] { - eprintln!("[DEBUG] Entity decoding: num_tokens={}, num_labels={}, threshold={}", - num_tokens, num_labels, threshold); + eprintln!( + "[DEBUG] Entity decoding: num_tokens={}, num_labels={}, threshold={}", + num_tokens, num_labels, threshold + ); eprintln!("[DEBUG] Sigmoid scores [start, end, inside] (first 5 tokens):"); for t in 0..5.min(num_tokens) { for l in 0..num_labels { let base = t * num_labels * 3 + l * 3; eprintln!( " token {} label {} ({}): start={:.4}, end={:.4}, inside={:.4}", - t, l, entity_labels[l], - scores[base], scores[base + 1], scores[base + 2]); + t, + l, + entity_labels[l], + scores[base], + scores[base + 1], + scores[base + 2] + ); } } } @@ -884,21 +931,32 @@ impl GlinerRelEx { for label_idx in 0..num_labels { for start_tok in 0..num_tokens { let start_score = score_at(start_tok, label_idx, 0); - if start_score < threshold { continue; } + if start_score < threshold { + continue; + } for end_tok in start_tok..num_tokens { let end_score = score_at(end_tok, label_idx, 1); - if end_score < threshold { continue; } + if end_score < threshold { + continue; + } // Check all inside scores from start_tok to end_tok let mut min_score = start_score.min(end_score); let mut valid = true; for t in start_tok..=end_tok { let inside = score_at(t, label_idx, 2); - if inside < threshold { valid = false; break; } - if inside < min_score { min_score = inside; } + if inside < threshold { + valid = false; + break; + } + if inside < min_score { + min_score = inside; + } + } + if !valid { + continue; } - if !valid { continue; } candidates.push((start_tok, end_tok, label_idx, min_score)); } @@ -913,7 +971,9 @@ impl GlinerRelEx { let mut entities = Vec::new(); for (start_tok, end_tok, label_idx, score) in candidates { let overlap = taken.iter().any(|&(a, b)| !(end_tok < a || start_tok > b)); - if overlap { continue; } + if overlap { + continue; + } taken.push((start_tok, end_tok)); if start_tok < word_offsets.len() && end_tok < word_offsets.len() { @@ -932,7 +992,11 @@ impl GlinerRelEx { } // Sort by score descending - entities.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal)); + entities.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); Ok(entities) } diff --git a/models/rgliner/src/relex_tokenization.rs b/models/rgliner/src/relex_tokenization.rs index d6bbdf9b8..2e429abfc 100644 --- a/models/rgliner/src/relex_tokenization.rs +++ b/models/rgliner/src/relex_tokenization.rs @@ -47,9 +47,8 @@ impl SpecialTokenIds { /// where `<>` is id 250102 vs 128001). Falls back to the corresponding /// field in `fallback` if the tokenizer doesn't contain a particular token. pub fn from_tokenizer(tokenizer: &tokenizers::Tokenizer, fallback: Self) -> Self { - let lookup = |tok: &str, default: u32| -> u32 { - tokenizer.token_to_id(tok).unwrap_or(default) - }; + let lookup = + |tok: &str, default: u32| -> u32 { tokenizer.token_to_id(tok).unwrap_or(default) }; Self { cls_id: lookup("[CLS]", fallback.cls_id), sep_id: lookup("[SEP]", fallback.sep_id), @@ -227,7 +226,10 @@ impl RelExTokenizer { } // Try to extend with (-|_)\w+ groups (greedy) loop { - if i + 1 < n && (bytes[i] == b'-' || bytes[i] == b'_') && is_word_char(bytes[i + 1]) { + if i + 1 < n + && (bytes[i] == b'-' || bytes[i] == b'_') + && is_word_char(bytes[i + 1]) + { i += 1; while i < n && is_word_char(bytes[i]) { i += 1; diff --git a/models/rgliner/src/source.rs b/models/rgliner/src/source.rs index 2a3b6560d..47ba2079d 100644 --- a/models/rgliner/src/source.rs +++ b/models/rgliner/src/source.rs @@ -130,11 +130,7 @@ impl GlinerSource { /// `Demonthos/gliner-gguf`. pub fn demonthos_edge() -> Self { Self { - model: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "gliner-edge.gguf", - ), + model: Self::huggingface_or_cached("Demonthos/gliner-gguf", "main", "gliner-edge.gguf"), label_encoder: Self::huggingface_or_cached( "Demonthos/gliner-gguf", "main", @@ -208,11 +204,7 @@ impl GlinerSource { /// `Demonthos/gliner-gguf`. pub fn demonthos_base() -> Self { Self { - model: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "gliner-base.gguf", - ), + model: Self::huggingface_or_cached("Demonthos/gliner-gguf", "main", "gliner-base.gguf"), label_encoder: Self::huggingface_or_cached( "Demonthos/gliner-gguf", "main", From 75579dccad87ae801e5bd1c6c2a4b1a0f35ca088 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 21:06:24 -0500 Subject: [PATCH 08/34] start moving models into rbert --- .claude/settings.local.json | 9 +- models/rbert/src/lib.rs | 3 +- .../src/raw/mdeberta/attention.rs | 0 .../src/raw/mdeberta/config.rs | 40 ++--- .../src/raw/mdeberta/feed_forward.rs | 0 .../src/raw/mdeberta/layer.rs | 0 .../src/raw/mdeberta/mod.rs | 1 + .../src/raw/mdeberta/model.rs | 10 +- models/rbert/src/raw/mod.rs | 4 + .../src/raw/modern_bert/attention.rs | 0 .../src/raw/modern_bert/config.rs | 0 .../src/raw/modern_bert/feed_forward.rs | 0 .../src/raw/modern_bert/layer.rs | 0 .../src/raw/modern_bert/mod.rs | 0 .../src/raw/modern_bert/model.rs | 11 +- models/rgliner/convert_to_gguf.py | 168 +++++++++++++++--- models/rgliner/examples/basic.rs | 6 +- models/rgliner/src/lib.rs | 2 +- models/rgliner/src/raw/mod.rs | 3 - models/rgliner/src/raw/text_encoder.rs | 44 ++++- models/rgliner/src/relex.rs | 2 +- 21 files changed, 232 insertions(+), 71 deletions(-) rename models/{rgliner => rbert}/src/raw/mdeberta/attention.rs (100%) rename models/{rgliner => rbert}/src/raw/mdeberta/config.rs (72%) rename models/{rgliner => rbert}/src/raw/mdeberta/feed_forward.rs (100%) rename models/{rgliner => rbert}/src/raw/mdeberta/layer.rs (100%) rename models/{rgliner => rbert}/src/raw/mdeberta/mod.rs (89%) rename models/{rgliner => rbert}/src/raw/mdeberta/model.rs (96%) rename models/{rgliner => rbert}/src/raw/modern_bert/attention.rs (100%) rename models/{rgliner => rbert}/src/raw/modern_bert/config.rs (100%) rename models/{rgliner => rbert}/src/raw/modern_bert/feed_forward.rs (100%) rename models/{rgliner => rbert}/src/raw/modern_bert/layer.rs (100%) rename models/{rgliner => rbert}/src/raw/modern_bert/mod.rs (100%) rename models/{rgliner => rbert}/src/raw/modern_bert/model.rs (89%) diff --git a/.claude/settings.local.json b/.claude/settings.local.json index 53905fd06..40fcd0c01 100644 --- a/.claude/settings.local.json +++ b/.claude/settings.local.json @@ -30,7 +30,14 @@ "Bash(git stash:*)", "Bash(FILTER_BRANCH_SQUELCH_WARNING=1 git filter-branch -f --index-filter 'git rm --cached --ignore-unmatch models/rgliner/weights/gliner-relex-multi-v1.0.gguf' d6d5c674..HEAD)", "Bash(git ls-tree:*)", - "Bash(git push:*)" + "Bash(git push:*)", + "WebFetch(domain:dioxuslabs.com)", + "Bash(cp /Users/evanalmloff/Desktop/Github/ner/models/rgliner/src/raw/modern_bert/*.rs /Users/evanalmloff/Desktop/Github/ner/models/rbert/src/raw/modern_bert/)", + "Bash(cp /Users/evanalmloff/Desktop/Github/ner/models/rgliner/src/raw/mdeberta/*.rs /Users/evanalmloff/Desktop/Github/ner/models/rbert/src/raw/mdeberta/)", + "Bash(cargo tree:*)", + "Bash(r\" | head -20)", + "Read(//tmp/**)", + "Bash(cargo fmt:*)" ] } } diff --git a/models/rbert/src/lib.rs b/models/rbert/src/lib.rs index d6665d733..637f71750 100644 --- a/models/rbert/src/lib.rs +++ b/models/rbert/src/lib.rs @@ -49,7 +49,8 @@ use std::sync::{Arc, RwLock}; use tokenizers::{Encoding, PaddingDirection, PaddingParams, Tokenizer}; mod language_model; -mod raw; +/// Low-level encoder implementations (standard BERT, Qwen, ModernBERT, mDeBERTa-v3). +pub mod raw; mod source; pub use crate::language_model::*; diff --git a/models/rgliner/src/raw/mdeberta/attention.rs b/models/rbert/src/raw/mdeberta/attention.rs similarity index 100% rename from models/rgliner/src/raw/mdeberta/attention.rs rename to models/rbert/src/raw/mdeberta/attention.rs diff --git a/models/rgliner/src/raw/mdeberta/config.rs b/models/rbert/src/raw/mdeberta/config.rs similarity index 72% rename from models/rgliner/src/raw/mdeberta/config.rs rename to models/rbert/src/raw/mdeberta/config.rs index fce51a18b..9c9c55d57 100644 --- a/models/rgliner/src/raw/mdeberta/config.rs +++ b/models/rbert/src/raw/mdeberta/config.rs @@ -31,31 +31,25 @@ pub struct MDebertaConfig { impl MDebertaConfig { /// Load configuration from GGUF metadata. - /// - /// Note: GGUF metadata keys use "gliner." prefix regardless of VarBuilder scope, - /// since metadata is stored globally (not per-tensor). pub fn from_gguf(vb: &VarBuilder) -> Result { - // Metadata keys use "gliner." prefix (not the tensor prefix) let num_heads = vb - .get_metadata("gliner.attention.head_count") + .get_metadata(".attention.head_count") .and_then(|v| v.to_u32().ok()) .ok_or_else(|| { - fusor::Error::msg("Missing required GGUF metadata: gliner.attention.head_count") + fusor::Error::msg("Missing required GGUF metadata: .attention.head_count") })? as usize; let num_layers = vb - .get_metadata("gliner.block_count") + .get_metadata(".block_count") .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| { - fusor::Error::msg("Missing required GGUF metadata: gliner.block_count") - })? as usize; + .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: .block_count"))? + as usize; let hidden_size = vb - .get_metadata("gliner.embedding_length") + .get_metadata(".embedding_length") .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| { - fusor::Error::msg("Missing required GGUF metadata: gliner.embedding_length") - })? as usize; + .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: .embedding_length"))? + as usize; if hidden_size % num_heads != 0 { return Err(fusor::Error::msg(format!( @@ -64,45 +58,43 @@ impl MDebertaConfig { } let head_dimension = vb - .get_metadata("gliner.attention.key_length") + .get_metadata(".attention.key_length") .and_then(|v| v.to_u32().ok()) .map(|x| x as usize) .unwrap_or_else(|| hidden_size / num_heads); let intermediate_size = vb - .get_metadata("gliner.feed_forward_length") + .get_metadata(".feed_forward_length") .and_then(|v| v.to_u32().ok()) .unwrap_or((hidden_size * 4) as u32) as usize; let context_length = vb - .get_metadata("gliner.context_length") + .get_metadata(".context_length") .and_then(|v| v.to_u32().ok()) .unwrap_or(512) as usize; - // DeBERTa-specific: maximum relative position distance let max_relative_positions = vb - .get_metadata("gliner.attention.max_relative_positions") + .get_metadata(".attention.max_relative_positions") .and_then(|v| v.to_u32().ok()) .unwrap_or(512) as usize; let norm_eps = vb - .get_metadata("gliner.attention.layer_norm_epsilon") + .get_metadata(".attention.layer_norm_epsilon") .and_then(|v| v.to_f32().ok()) .unwrap_or(1e-7); let vocab_size = vb - .get_metadata("gliner.vocab_size") + .get_metadata(".vocab_size") .and_then(|v| v.to_u32().ok()) .unwrap_or(250105) as usize; - // DeBERTa-v3 specific: position buckets for relative position encoding let position_buckets = vb - .get_metadata("gliner.attention.position_buckets") + .get_metadata(".attention.position_buckets") .and_then(|v| v.to_u32().ok()) .unwrap_or(256) as usize; let share_att_key = vb - .get_metadata("gliner.attention.share_att_key") + .get_metadata(".attention.share_att_key") .and_then(|v| v.to_bool().ok()) .unwrap_or(true); diff --git a/models/rgliner/src/raw/mdeberta/feed_forward.rs b/models/rbert/src/raw/mdeberta/feed_forward.rs similarity index 100% rename from models/rgliner/src/raw/mdeberta/feed_forward.rs rename to models/rbert/src/raw/mdeberta/feed_forward.rs diff --git a/models/rgliner/src/raw/mdeberta/layer.rs b/models/rbert/src/raw/mdeberta/layer.rs similarity index 100% rename from models/rgliner/src/raw/mdeberta/layer.rs rename to models/rbert/src/raw/mdeberta/layer.rs diff --git a/models/rgliner/src/raw/mdeberta/mod.rs b/models/rbert/src/raw/mdeberta/mod.rs similarity index 89% rename from models/rgliner/src/raw/mdeberta/mod.rs rename to models/rbert/src/raw/mdeberta/mod.rs index 18ee18d02..32cab81dc 100644 --- a/models/rgliner/src/raw/mdeberta/mod.rs +++ b/models/rbert/src/raw/mdeberta/mod.rs @@ -9,4 +9,5 @@ mod feed_forward; mod layer; mod model; +pub use config::MDebertaConfig; pub use model::MDebertaModel; diff --git a/models/rgliner/src/raw/mdeberta/model.rs b/models/rbert/src/raw/mdeberta/model.rs similarity index 96% rename from models/rgliner/src/raw/mdeberta/model.rs rename to models/rbert/src/raw/mdeberta/model.rs index 115889ebc..63159de62 100644 --- a/models/rgliner/src/raw/mdeberta/model.rs +++ b/models/rbert/src/raw/mdeberta/model.rs @@ -10,10 +10,9 @@ use super::attention::RelativePositionEmbedding; use super::config::MDebertaConfig; use super::layer::MDebertaLayer; -/// mDeBERTa-v3 encoder model for GLiNER-RelEx. -/// -/// This is a bidirectional transformer encoder using disentangled attention -/// with relative position embeddings. +/// A raw synchronous mDeBERTa-v3 encoder model. This is a bidirectional +/// transformer encoder using disentangled attention with relative position +/// embeddings (DeBERTa-v3 architecture). pub struct MDebertaModel { /// Token embeddings token_embeddings: Embedding, @@ -31,6 +30,7 @@ pub struct MDebertaModel { device: Device, /// Configuration config: MDebertaConfig, + span: tracing::Span, } impl MDebertaModel { @@ -89,6 +89,7 @@ impl MDebertaModel { output_proj, device: device.clone(), config, + span: tracing::span!(tracing::Level::TRACE, "mdeberta"), }) } @@ -105,6 +106,7 @@ impl MDebertaModel { input_ids: &Tensor<2, u32>, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { + let _enter = self.span.enter(); let [_batch_size, seq_len] = input_ids.shape(); // Get token embeddings diff --git a/models/rbert/src/raw/mod.rs b/models/rbert/src/raw/mod.rs index adf89b896..96220b401 100644 --- a/models/rbert/src/raw/mod.rs +++ b/models/rbert/src/raw/mod.rs @@ -16,8 +16,12 @@ mod self_output; use self_output::*; mod intermediate_layer; use intermediate_layer::*; +pub mod mdeberta; +pub mod modern_bert; pub mod qwen; +pub use mdeberta::{MDebertaConfig, MDebertaModel}; +pub use modern_bert::{ModernBertConfig, ModernBertModel}; pub use qwen::QwenEmbeddingModel; use fusor::{Device, Result, Tensor, VarBuilder}; diff --git a/models/rgliner/src/raw/modern_bert/attention.rs b/models/rbert/src/raw/modern_bert/attention.rs similarity index 100% rename from models/rgliner/src/raw/modern_bert/attention.rs rename to models/rbert/src/raw/modern_bert/attention.rs diff --git a/models/rgliner/src/raw/modern_bert/config.rs b/models/rbert/src/raw/modern_bert/config.rs similarity index 100% rename from models/rgliner/src/raw/modern_bert/config.rs rename to models/rbert/src/raw/modern_bert/config.rs diff --git a/models/rgliner/src/raw/modern_bert/feed_forward.rs b/models/rbert/src/raw/modern_bert/feed_forward.rs similarity index 100% rename from models/rgliner/src/raw/modern_bert/feed_forward.rs rename to models/rbert/src/raw/modern_bert/feed_forward.rs diff --git a/models/rgliner/src/raw/modern_bert/layer.rs b/models/rbert/src/raw/modern_bert/layer.rs similarity index 100% rename from models/rgliner/src/raw/modern_bert/layer.rs rename to models/rbert/src/raw/modern_bert/layer.rs diff --git a/models/rgliner/src/raw/modern_bert/mod.rs b/models/rbert/src/raw/modern_bert/mod.rs similarity index 100% rename from models/rgliner/src/raw/modern_bert/mod.rs rename to models/rbert/src/raw/modern_bert/mod.rs diff --git a/models/rgliner/src/raw/modern_bert/model.rs b/models/rbert/src/raw/modern_bert/model.rs similarity index 89% rename from models/rgliner/src/raw/modern_bert/model.rs rename to models/rbert/src/raw/modern_bert/model.rs index 5d1e88a71..58dfcc47c 100644 --- a/models/rgliner/src/raw/modern_bert/model.rs +++ b/models/rbert/src/raw/modern_bert/model.rs @@ -6,7 +6,9 @@ use fusor::{Device, Result, RopeCache, Tensor, VarBuilder}; use super::config::ModernBertConfig; use super::layer::ModernBertLayer; -/// ModernBERT encoder model (text encoder for GLiNER). +/// A raw synchronous ModernBERT (Ettin) encoder model. This is a bidirectional +/// transformer with RoPE positional embeddings, pre-normalization, and GeGLU +/// feed-forward blocks. pub struct ModernBertModel { token_embeddings: Embedding, /// Embedding norm applied after token embeddings, before first layer @@ -16,6 +18,7 @@ pub struct ModernBertModel { rope_cache: RopeCache, pub(crate) device: Device, config: ModernBertConfig, + span: tracing::Span, } impl ModernBertModel { @@ -63,6 +66,7 @@ impl ModernBertModel { rope_cache, device: device.clone(), config, + span: tracing::span!(tracing::Level::TRACE, "modern-bert"), }) } @@ -74,6 +78,7 @@ impl ModernBertModel { input_ids: &Tensor<2, u32>, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { + let _enter = self.span.enter(); // Get token embeddings let hidden_states = self.token_embeddings.forward(input_ids); @@ -104,12 +109,14 @@ impl ModernBertModel { &self.device } - #[cfg(test)] + /// Return the hidden state after each layer (for debugging / regression tests). + #[doc(hidden)] pub fn debug_hidden_states( &self, input_ids: &Tensor<2, u32>, attention_mask: Option<&Tensor<2, u32>>, ) -> Vec> { + let _enter = self.span.enter(); let mut states = Vec::with_capacity(self.layers.len() + 2); let hidden_states = self.token_embeddings.forward(input_ids); diff --git a/models/rgliner/convert_to_gguf.py b/models/rgliner/convert_to_gguf.py index 039156189..9c01c028c 100644 --- a/models/rgliner/convert_to_gguf.py +++ b/models/rgliner/convert_to_gguf.py @@ -289,6 +289,12 @@ def map_weight_name(pytorch_name: str) -> str: """Map PyTorch weight names to GGUF conventions.""" name = pytorch_name + # ===== Encoder output projection ===== + # Some bi-encoder variants (small/base/large v2.0) project the text-encoder + # hidden state down to the shared label-aligned dim (e.g. 512 -> 384). + name = name.replace("token_rep_layer.projection.weight", "text.output_proj.weight") + name = name.replace("token_rep_layer.projection.bias", "text.output_proj.bias") + # ===== Text Encoder (ModernBERT/Ettin) ===== # token_rep_layer.bert_layer.model.embeddings.tok_embeddings.weight -> text.token_embd.weight name = name.replace("token_rep_layer.bert_layer.model.embeddings.tok_embeddings.weight", "text.token_embd.weight") @@ -353,6 +359,52 @@ def map_weight_name(pytorch_name: str) -> str: return name +def _llama_quantize(f32_path: str, output_path: str, quant_type: str, keep_f32: List[str]) -> None: + """Run `llama-quantize` to convert an f32 GGUF to a k-quant type. + + `keep_f32` is a list of tensor names that should remain at F32 (tiny + classifier heads hit a fusor bug; see the gliner-relex project notes). + """ + binary = shutil.which("llama-quantize") + if binary is None: + raise RuntimeError( + "`llama-quantize` not found in PATH. Install llama.cpp to use k-quant types:\n" + " brew install llama.cpp # macOS\n" + " or build from https://github.com/ggml-org/llama.cpp" + ) + cmd = [binary] + for name in keep_f32: + cmd += ["--tensor-type", f"{name}=f32"] + cmd += [f32_path, output_path, quant_type.upper()] + print(f"\n$ {' '.join(cmd)}") + result = subprocess.run(cmd) + if result.returncode != 0: + raise RuntimeError( + f"llama-quantize failed (exit {result.returncode}) for {quant_type!r}. " + f"See output above. The f32 GGUF was kept at {f32_path} for inspection." + ) + + +def _quantize_tensor(data: np.ndarray, quant: str, default_ggml_type: int) -> Tuple[np.ndarray, int, Tuple[int, ...]]: + """Quantise a single tensor, falling back to F32 for unsupported shapes. + + Returns (packed_data, ggml_type, logical_shape). Applies the same rules as + the relex converter: only 2-D+ tensors with inner dim divisible by 32 and + outer dim >= 32 are block-quantised. The lower outer-dim threshold avoids + a fusor quantised-matmul bug on tiny classifier heads. + """ + logical_shape = tuple(data.shape) + MIN_OUT_DIM_FOR_QUANT = 32 + if ( + quant.startswith("q") + and data.ndim >= 2 + and data.shape[-1] % 32 == 0 + and data.shape[0] >= MIN_OUT_DIM_FOR_QUANT + ): + return _gguf_block_quant(data, quant), default_ggml_type, logical_shape + return np.ascontiguousarray(data, dtype=np.float32), GGML_TYPE_F32, logical_shape + + def convert_gliner_to_gguf( model_id: str, output_path: str, @@ -374,15 +426,39 @@ def convert_gliner_to_gguf( for name, tensor in state_dict.items(): print(f" {name}: {tensor.shape} {tensor.dtype}") - # Determine quantization type - if quantize == "f32": - ggml_type = GGML_TYPE_F32 - elif quantize == "f16": - ggml_type = GGML_TYPE_F16 - elif quantize == "bf16": - ggml_type = GGML_TYPE_BF16 - else: - raise ValueError(f"Unsupported quantization: {quantize}") + # Determine quantization type and post-processing + quantize = quantize.lower() + if quantize not in QUANT_TYPES: + raise ValueError( + f"Unsupported quantization: {quantize!r}. Supported: {', '.join(sorted(QUANT_TYPES))}" + ) + post_quant = None + if quantize in LLAMA_QUANT_TYPES: + post_quant = quantize + quantize = "f32" + default_ggml_type = _ggml_type_for(quantize) + + # When post-quantising, write the raw f32 file to a temp path first. + main_label_output = output_path.replace(".gguf", "-label-encoder.gguf") + final_main = output_path + final_label = main_label_output + if post_quant is not None: + main_tmp = tempfile.NamedTemporaryFile( + prefix=os.path.basename(output_path).rsplit(".", 1)[0] + ".f32.", + suffix=".gguf", + delete=False, + dir=os.path.dirname(os.path.abspath(output_path)) or None, + ) + label_tmp = tempfile.NamedTemporaryFile( + prefix=os.path.basename(main_label_output).rsplit(".", 1)[0] + ".f32.", + suffix=".gguf", + delete=False, + dir=os.path.dirname(os.path.abspath(main_label_output)) or None, + ) + output_path = main_tmp.name + main_label_output = label_tmp.name + main_tmp.close() + label_tmp.close() # Separate label encoder weights from main model weights label_encoder_weights = {} @@ -397,8 +473,11 @@ def convert_gliner_to_gguf( # ============ Main Model GGUF ============ writer = GGUFWriter(output_path) - # Add metadata - writer.add_metadata("general.architecture", "gliner") + # Add metadata. Masquerade as `bert` when we'll feed this to llama-quantize + # so its loader accepts the custom architecture. rgliner reads `gliner.*` + # keys, not `general.architecture`. + main_arch = "bert" if post_quant is not None else "gliner" + writer.add_metadata("general.architecture", main_arch) writer.add_metadata("general.name", model_id.split("/")[-1]) writer.add_metadata("general.quantization_version", 2) @@ -428,8 +507,17 @@ def convert_gliner_to_gguf( writer.add_metadata("gliner.attention.layer_norm_rms_epsilon", 1e-5) writer.add_metadata("gliner.vocab_size", vocab_size) + # When feeding to llama-quantize, mirror arch keys with u32 scalars. + if post_quant is not None: + writer.add_metadata("bert.context_length", _U32(context_length)) + writer.add_metadata("bert.embedding_length", _U32(hidden_size)) + writer.add_metadata("bert.feed_forward_length", _U32(intermediate_size)) + writer.add_metadata("bert.block_count", _U32(num_layers)) + writer.add_metadata("bert.attention.head_count", _U32(num_heads)) + writer.add_metadata("bert.attention.layer_norm_epsilon", 1e-5) + # Convert main model tensors - print(f"\nConverting {len(main_model_weights)} main model tensors to GGUF...") + print(f"\nConverting {len(main_model_weights)} main model tensors to GGUF ({quantize.upper()})...") for pytorch_name, tensor in main_model_weights.items(): gguf_name = map_weight_name(pytorch_name) @@ -445,8 +533,9 @@ def convert_gliner_to_gguf( dtype=np.float32 ).reshape(t.shape) - print(f" {pytorch_name} -> {gguf_name} {data.shape}") - writer.add_tensor(gguf_name, data, ggml_type) + packed, tensor_type, logical_shape = _quantize_tensor(data, quantize, default_ggml_type) + print(f" {pytorch_name} -> {gguf_name} {logical_shape} [{_name_for_type(tensor_type)}]") + writer.add_tensor(gguf_name, packed, tensor_type, shape=logical_shape) writer.write() print(f"Main model output: {output_path}") @@ -454,7 +543,7 @@ def convert_gliner_to_gguf( # ============ Label Encoder GGUF ============ # Create separate file for label encoder (without prefix, for rbert compatibility) - label_output_path = output_path.replace(".gguf", "-label-encoder.gguf") + label_output_path = main_label_output label_writer = GGUFWriter(label_output_path) # Add BERT metadata @@ -469,16 +558,18 @@ def convert_gliner_to_gguf( label_vocab = labels_config.get("vocab_size", 30522) label_max_pos = labels_config.get("max_position_embeddings", 512) - label_writer.add_metadata("bert.attention.head_count", label_heads) - label_writer.add_metadata("bert.block_count", label_layers) - label_writer.add_metadata("bert.embedding_length", label_hidden) - label_writer.add_metadata("bert.feed_forward_length", label_intermediate) - label_writer.add_metadata("bert.context_length", label_max_pos) + # Use u32 for arch-scoped ints when we'll feed this through llama-quantize. + int_ctor = _U32 if post_quant is not None else (lambda x: x) + label_writer.add_metadata("bert.attention.head_count", int_ctor(label_heads)) + label_writer.add_metadata("bert.block_count", int_ctor(label_layers)) + label_writer.add_metadata("bert.embedding_length", int_ctor(label_hidden)) + label_writer.add_metadata("bert.feed_forward_length", int_ctor(label_intermediate)) + label_writer.add_metadata("bert.context_length", int_ctor(label_max_pos)) label_writer.add_metadata("bert.attention.layer_norm_epsilon", 1e-12) - label_writer.add_metadata("bert.vocab_size", label_vocab) + label_writer.add_metadata("bert.vocab_size", int_ctor(label_vocab)) # Convert label encoder tensors (remove prefix so rbert can load them) - print(f"\nConverting {len(label_encoder_weights)} label encoder tensors to GGUF...") + print(f"\nConverting {len(label_encoder_weights)} label encoder tensors to GGUF ({quantize.upper()})...") for pytorch_name, tensor in label_encoder_weights.items(): # Map name but remove the "label." prefix for rbert compatibility @@ -496,13 +587,33 @@ def convert_gliner_to_gguf( dtype=np.float32 ).reshape(t.shape) - print(f" {pytorch_name} -> {gguf_name} {data.shape}") - label_writer.add_tensor(gguf_name, data, ggml_type) + packed, tensor_type, logical_shape = _quantize_tensor(data, quantize, default_ggml_type) + print(f" {pytorch_name} -> {gguf_name} {logical_shape} [{_name_for_type(tensor_type)}]") + label_writer.add_tensor(gguf_name, packed, tensor_type, shape=logical_shape) label_writer.write() print(f"Label encoder output: {label_output_path}") print(f"Size: {os.path.getsize(label_output_path) / 1024 / 1024:.2f} MB") + if post_quant is not None: + print(f"\nQuantizing main model to {post_quant.upper()} via llama-quantize...") + _llama_quantize(output_path, final_main, post_quant, keep_f32=[]) + try: + os.remove(output_path) + except OSError: + pass + print(f"Final main output: {final_main}") + print(f"Final main size: {os.path.getsize(final_main) / 1024 / 1024:.2f} MB") + + print(f"\nQuantizing label encoder to {post_quant.upper()} via llama-quantize...") + _llama_quantize(label_output_path, final_label, post_quant, keep_f32=[]) + try: + os.remove(label_output_path) + except OSError: + pass + print(f"Final label output: {final_label}") + print(f"Final label size: {os.path.getsize(final_label) / 1024 / 1024:.2f} MB") + print(f"\nConversion complete!") @@ -524,8 +635,13 @@ def main(): "--quantize", "-q", type=str, default="f32", - choices=["f32", "f16", "bf16"], - help="Quantization type (default: f32)" + choices=sorted(QUANT_TYPES), + help=( + "Quantization type (default: f32). " + "Block quants (q4_0, q5_0, q8_0, ...) are packed in-process via the " + "`gguf` Python package. K-quants (q4_k, q5_k, q6_k, ...) require " + "`llama-quantize` in PATH and are applied as a post-processing step." + ), ) parser.add_argument( "--cache-dir", diff --git a/models/rgliner/examples/basic.rs b/models/rgliner/examples/basic.rs index dd8af2f4a..40d09d55c 100644 --- a/models/rgliner/examples/basic.rs +++ b/models/rgliner/examples/basic.rs @@ -15,7 +15,11 @@ async fn main() -> anyhow::Result<()> { GlinerSource::edge() }; - let mut gliner = Gliner::builder().with_source(source).build().await?; + let mut gliner = Gliner::builder() + .with_source(source) + .with_threshold(0.01) + .build() + .await?; println!("Model loaded!"); diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs index 57a378387..083c28980 100644 --- a/models/rgliner/src/lib.rs +++ b/models/rgliner/src/lib.rs @@ -86,7 +86,7 @@ mod tokenization; pub use config::GlinerConfig; pub use decoding::{Decoder, DecodingMode, Entity}; pub use error::{GlinerError, GlinerLoadingError}; -pub use raw::modern_bert::{ModernBertConfig, ModernBertModel}; +pub use rbert::raw::{ModernBertConfig, ModernBertModel}; pub use source::GlinerSource; use fusor::{Device, Tensor, VarBuilder}; diff --git a/models/rgliner/src/raw/mod.rs b/models/rgliner/src/raw/mod.rs index 9804d6b5f..a1812477d 100644 --- a/models/rgliner/src/raw/mod.rs +++ b/models/rgliner/src/raw/mod.rs @@ -1,8 +1,5 @@ //! Raw model implementations for GLiNER. -pub mod mdeberta; -pub mod modern_bert; - mod bilstm; mod joint_scorer; mod label_encoder; diff --git a/models/rgliner/src/raw/text_encoder.rs b/models/rgliner/src/raw/text_encoder.rs index 19fd945a3..95207c6af 100644 --- a/models/rgliner/src/raw/text_encoder.rs +++ b/models/rgliner/src/raw/text_encoder.rs @@ -1,31 +1,56 @@ //! Text encoder wrapper for GLiNER. +use fusor::layers::Linear; use fusor::{Device, Result, Tensor, VarBuilder}; -use super::modern_bert::ModernBertModel; +use rbert::raw::ModernBertModel; /// Text encoder for GLiNER (ModernBERT/Ettin). pub struct TextEncoder { model: ModernBertModel, + /// Optional output projection. Some bi-encoder variants (e.g. + /// `gliner-bi-small-v2.0`) have a `token_rep_layer.projection` that maps + /// the encoder's native hidden size down to the dim shared with the label + /// encoder / downstream heads (e.g. 512 -> 384). Absent on `edge`. + output_proj: Option>, } impl TextEncoder { /// Load text encoder from GGUF weights. pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { // GLiNER GGUF uses "text." prefix for text encoder weights - let model = ModernBertModel::load(device, &mut vb.pp("text"))?; - Ok(Self { model }) + let mut text_vb = vb.pp("text"); + let model = ModernBertModel::load(device, &mut text_vb)?; + // Optional output projection (small/base/large v2.0 variants). + let output_proj = Linear::load(device, &mut text_vb.pp("output_proj")).ok(); + + #[cfg(debug_assertions)] + if let Some(ref p) = output_proj { + eprintln!( + "[DEBUG] TextEncoder output projection loaded: {} -> {}", + p.in_features(), + p.out_features() + ); + } + + Ok(Self { model, output_proj }) } /// Forward pass returning per-token embeddings. /// - /// Returns: [batch_size, seq_len, hidden_size] + /// Returns: [batch_size, seq_len, hidden_size] (hidden_size is the + /// projected dim if `output_proj` is present) pub fn forward( &self, input_ids: &Tensor<2, u32>, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { - self.model.forward(input_ids, attention_mask) + let hidden = self.model.forward(input_ids, attention_mask); + if let Some(ref proj) = self.output_proj { + proj.forward(&hidden) + } else { + hidden + } } /// Get the maximum sequence length. @@ -33,9 +58,14 @@ impl TextEncoder { self.model.max_seq_len() } - /// Get the embedding dimension. + /// Get the embedding dimension seen by downstream layers (post-projection + /// if this variant has one, otherwise the raw encoder hidden dim). pub fn embedding_dim(&self) -> usize { - self.model.embedding_dim() + if let Some(ref proj) = self.output_proj { + proj.out_features() + } else { + self.model.embedding_dim() + } } /// Get the device. diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index d22e8127f..9d48223e1 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -50,12 +50,12 @@ use tokenizers::Tokenizer; use crate::decoding::Entity; use crate::error::{GlinerError, GlinerLoadingError}; -use crate::raw::mdeberta::MDebertaModel; use crate::raw::{ BiLstm, JointScorer, PairProjector, PromptRepLayer, RelationsRepLayer, SpanLayer, }; use crate::relation_decoding::{Relation, RelationDecoder, RelationDecoderConfig}; use crate::relex_tokenization::{RelExTokenizer, SpecialTokenIds}; +use rbert::raw::MDebertaModel; /// Source configuration for GLiNER-RelEx models. /// From 724ade19bbeec32a85c4d24287844ba1160d97b3 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 21:14:00 -0500 Subject: [PATCH 09/34] dedup some code --- models/rbert/src/raw/mdeberta/attention.rs | 182 ++++-------------- models/rbert/src/raw/mdeberta/config.rs | 79 ++------ models/rbert/src/raw/mdeberta/layer.rs | 20 -- models/rbert/src/raw/mod.rs | 1 + models/rbert/src/raw/modern_bert/attention.rs | 54 ++---- models/rbert/src/raw/modern_bert/config.rs | 68 ++----- models/rbert/src/raw/self_attention.rs | 26 +-- models/rbert/src/raw/utils.rs | 73 +++++++ 8 files changed, 170 insertions(+), 333 deletions(-) create mode 100644 models/rbert/src/raw/utils.rs diff --git a/models/rbert/src/raw/mdeberta/attention.rs b/models/rbert/src/raw/mdeberta/attention.rs index 3b07cb3e5..cd82b80ec 100644 --- a/models/rbert/src/raw/mdeberta/attention.rs +++ b/models/rbert/src/raw/mdeberta/attention.rs @@ -136,31 +136,6 @@ impl RelativePositionEmbedding { self.embeddings.to_concrete() } } - - /// Get relative position embeddings for the given indices (legacy method). - /// Input: indices [seq_len, seq_len] - /// Output: embeddings [seq_len, seq_len, hidden_size] - pub fn forward(&self, indices: &Tensor<2, u32>) -> Tensor<3, f32> { - let [seq_len, _] = indices.shape(); - let [_num_positions, hidden_size] = self.embeddings.shape(); - - // Get normalized embeddings - let normalized_embeddings = self.get_embeddings(); - - // Flatten indices and gather - let flat_indices = indices.reshape([seq_len * seq_len]).to_concrete(); - let gathered = normalized_embeddings.index_select(0, &flat_indices); - - // Reshape back to [seq_len, seq_len, hidden_size] - gathered - .reshape([seq_len, seq_len, hidden_size]) - .to_concrete() - } - - /// Get the maximum relative positions setting. - pub fn max_relative_positions(&self) -> usize { - self.max_relative_positions - } } /// mDeBERTa disentangled self-attention with shared key attention (share_att_key=True). @@ -216,51 +191,42 @@ impl MDebertaAttention { rel_pos_indices: &Tensor<2, u32>, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { - let [b_sz, seq_len, _] = hidden_states.shape(); - let hidden_size = self.num_heads * self.head_dim; - let [num_positions, _] = rel_pos_emb.shape(); - - // Compute Q, K, V projections for content - let query = self.query.forward(hidden_states); - let key = self.key.forward(hidden_states); - let value = self.value.forward(hidden_states); - - // Reshape to [batch, num_heads, seq_len, head_dim] - let query = query - .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); - let key = key - .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); - let value = value - .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); + use super::super::utils::split_heads; + + // Compute Q, K, V projections for content and reshape to + // [batch, num_heads, seq_len, head_dim]. + let query = split_heads( + &self.query.forward(hidden_states), + self.num_heads, + self.head_dim, + ); + let key = split_heads( + &self.key.forward(hidden_states), + self.num_heads, + self.head_dim, + ); + let value = split_heads( + &self.value.forward(hidden_states), + self.num_heads, + self.head_dim, + ); // === Content-to-Content attention === - // c2c = Q @ K^T let c2c_scores = query.mat_mul(&key.transpose(2, 3)); // === Position attention with shared Q/K projections === - // Project position embeddings using the same Q and K projections // rel_pos_emb: [2*max_pos, hidden_size] -> [1, 2*max_pos, hidden_size] let rel_emb_3d: Tensor<3, f32> = rel_pos_emb.unsqueeze(0).to_concrete(); - - // pos_query = query_proj(rel_emb): [1, 2*max_pos, hidden] -> [1, heads, 2*max_pos, head_dim] - let pos_query = self.query.forward(&rel_emb_3d); - let pos_query = pos_query - .reshape([1, num_positions, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); - - // pos_key = key_proj(rel_emb): [1, 2*max_pos, hidden] -> [1, heads, 2*max_pos, head_dim] - let pos_key = self.key.forward(&rel_emb_3d); - let pos_key = pos_key - .reshape([1, num_positions, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); + let pos_query = split_heads( + &self.query.forward(&rel_emb_3d), + self.num_heads, + self.head_dim, + ); + let pos_key = split_heads( + &self.key.forward(&rel_emb_3d), + self.num_heads, + self.head_dim, + ); // === Content-to-Position attention === // c2p = Q @ pos_key^T -> [batch, heads, seq, 2*max_pos] @@ -280,14 +246,9 @@ impl MDebertaAttention { .add_(&p2c_scores) .mul_scalar(self.scale); - // Apply attention mask + // Apply attention mask (broadcast bias to [batch, 1, 1, seq_len]) let attn_scores = if let Some(mask) = attention_mask { - const MASK_NEG_VALUE: f32 = -10000.0; - let mask_f32: Tensor<2, f32> = mask.cast(); - let zeros = mask_f32.zeros_like(); - let ones = (zeros + 1.0f32).to_concrete(); - let mask_bias = ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete(); - // Broadcast mask to [batch, 1, 1, seq_len] + let mask_bias = super::super::utils::attention_mask_to_bias(mask); let mask_bias_3d: Tensor<3, f32> = mask_bias.unsqueeze(1).to_concrete(); let mask_bias_4d: Tensor<4, f32> = mask_bias_3d.unsqueeze(1).to_concrete(); attn_scores.add_(&mask_bias_4d) @@ -298,17 +259,9 @@ impl MDebertaAttention { // Softmax let attn_probs = attn_scores.softmax_last_dim::<3>(); - // Apply attention to values + // Apply attention to values and merge heads back to [batch, seq_len, hidden]. let context = attn_probs.mat_mul(&value); - - // Reshape back to [batch, seq_len, hidden_size] - let context = context - .transpose(1, 2) - .to_concrete() - .reshape([b_sz, seq_len, hidden_size]) - .to_concrete(); - - // Output projection + let context = super::super::utils::merge_heads(&context); self.output.forward(&context) } @@ -410,65 +363,6 @@ impl MDebertaAttention { .reshape([b_sz, num_heads, seq_len, seq_len]) .to_concrete() } - - /// Legacy forward pass (for compatibility). - pub fn forward( - &self, - hidden_states: &Tensor<3, f32>, - rel_pos_emb: Option<&Tensor<3, f32>>, - attention_mask: Option<&Tensor<2, u32>>, - ) -> Tensor<3, f32> { - // This method is kept for backward compatibility but shouldn't be used - // with the new architecture - if rel_pos_emb.is_some() { - panic!("Use forward_with_indices for proper position attention"); - } - - let [b_sz, seq_len, _] = hidden_states.shape(); - let hidden_size = self.num_heads * self.head_dim; - - let query = self.query.forward(hidden_states); - let key = self.key.forward(hidden_states); - let value = self.value.forward(hidden_states); - - let query = query - .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); - let key = key - .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); - let value = value - .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); - - let c2c_scores = query.mat_mul(&key.transpose(2, 3)); - let attn_scores = c2c_scores.mul_scalar(1.0 / (self.head_dim as f32).sqrt()); - - let attn_scores = if let Some(mask) = attention_mask { - const MASK_NEG_VALUE: f32 = -10000.0; - let mask_f32: Tensor<2, f32> = mask.cast(); - let zeros = mask_f32.zeros_like(); - let ones = (zeros + 1.0f32).to_concrete(); - let mask_bias = ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete(); - let mask_bias_3d: Tensor<3, f32> = mask_bias.unsqueeze(1).to_concrete(); - let mask_bias_4d: Tensor<4, f32> = mask_bias_3d.unsqueeze(1).to_concrete(); - attn_scores.add_(&mask_bias_4d) - } else { - attn_scores - }; - - let attn_probs = attn_scores.softmax_last_dim::<3>(); - let context = attn_probs.mat_mul(&value); - let context = context - .transpose(1, 2) - .to_concrete() - .reshape([b_sz, seq_len, hidden_size]) - .to_concrete(); - self.output.forward(&context) - } } /// Shared relative position embedding layer (used across all layers in DeBERTa). @@ -502,14 +396,4 @@ impl DisentangledSelfAttention { attention_mask, ) } - - pub fn forward( - &self, - hidden_states: &Tensor<3, f32>, - rel_pos_emb: Option<&Tensor<3, f32>>, - attention_mask: Option<&Tensor<2, u32>>, - ) -> Tensor<3, f32> { - self.attention - .forward(hidden_states, rel_pos_emb, attention_mask) - } } diff --git a/models/rbert/src/raw/mdeberta/config.rs b/models/rbert/src/raw/mdeberta/config.rs index 9c9c55d57..a101f6235 100644 --- a/models/rbert/src/raw/mdeberta/config.rs +++ b/models/rbert/src/raw/mdeberta/config.rs @@ -2,6 +2,8 @@ use fusor::{Result, VarBuilder}; +use super::super::utils::{load_bool_or, load_f32_or, load_u32, load_u32_or}; + /// Configuration for mDeBERTa-v3 loaded from GGUF metadata. #[derive(Debug, Clone)] pub struct MDebertaConfig { @@ -32,24 +34,9 @@ pub struct MDebertaConfig { impl MDebertaConfig { /// Load configuration from GGUF metadata. pub fn from_gguf(vb: &VarBuilder) -> Result { - let num_heads = vb - .get_metadata(".attention.head_count") - .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| { - fusor::Error::msg("Missing required GGUF metadata: .attention.head_count") - })? as usize; - - let num_layers = vb - .get_metadata(".block_count") - .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: .block_count"))? - as usize; - - let hidden_size = vb - .get_metadata(".embedding_length") - .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: .embedding_length"))? - as usize; + let num_heads = load_u32(vb, ".attention.head_count")? as usize; + let num_layers = load_u32(vb, ".block_count")? as usize; + let hidden_size = load_u32(vb, ".embedding_length")? as usize; if hidden_size % num_heads != 0 { return Err(fusor::Error::msg(format!( @@ -57,46 +44,22 @@ impl MDebertaConfig { ))); } - let head_dimension = vb - .get_metadata(".attention.key_length") - .and_then(|v| v.to_u32().ok()) - .map(|x| x as usize) - .unwrap_or_else(|| hidden_size / num_heads); - - let intermediate_size = vb - .get_metadata(".feed_forward_length") - .and_then(|v| v.to_u32().ok()) - .unwrap_or((hidden_size * 4) as u32) as usize; - - let context_length = vb - .get_metadata(".context_length") - .and_then(|v| v.to_u32().ok()) - .unwrap_or(512) as usize; - - let max_relative_positions = vb - .get_metadata(".attention.max_relative_positions") - .and_then(|v| v.to_u32().ok()) - .unwrap_or(512) as usize; - - let norm_eps = vb - .get_metadata(".attention.layer_norm_epsilon") - .and_then(|v| v.to_f32().ok()) - .unwrap_or(1e-7); - - let vocab_size = vb - .get_metadata(".vocab_size") - .and_then(|v| v.to_u32().ok()) - .unwrap_or(250105) as usize; - - let position_buckets = vb - .get_metadata(".attention.position_buckets") - .and_then(|v| v.to_u32().ok()) - .unwrap_or(256) as usize; - - let share_att_key = vb - .get_metadata(".attention.share_att_key") - .and_then(|v| v.to_bool().ok()) - .unwrap_or(true); + let head_dimension = load_u32_or(vb, ".attention.key_length", 0); + let head_dimension = if head_dimension == 0 { + hidden_size / num_heads + } else { + head_dimension as usize + }; + + let intermediate_size = + load_u32_or(vb, ".feed_forward_length", (hidden_size * 4) as u32) as usize; + let context_length = load_u32_or(vb, ".context_length", 512) as usize; + let max_relative_positions = + load_u32_or(vb, ".attention.max_relative_positions", 512) as usize; + let norm_eps = load_f32_or(vb, ".attention.layer_norm_epsilon", 1e-7); + let vocab_size = load_u32_or(vb, ".vocab_size", 250105) as usize; + let position_buckets = load_u32_or(vb, ".attention.position_buckets", 256) as usize; + let share_att_key = load_bool_or(vb, ".attention.share_att_key", true); Ok(Self { num_heads, diff --git a/models/rbert/src/raw/mdeberta/layer.rs b/models/rbert/src/raw/mdeberta/layer.rs index ed6d27311..d22bf6111 100644 --- a/models/rbert/src/raw/mdeberta/layer.rs +++ b/models/rbert/src/raw/mdeberta/layer.rs @@ -71,24 +71,4 @@ impl MDebertaLayer { let ffn_output = self.feed_forward.forward(&hidden_states); self.output_norm.forward(&hidden_states.add_(&ffn_output)) } - - /// Legacy forward pass (for compatibility). - pub fn forward( - &self, - hidden_states: &Tensor<3, f32>, - rel_pos_emb: Option<&Tensor<3, f32>>, - attention_mask: Option<&Tensor<2, u32>>, - ) -> Tensor<3, f32> { - // Self-attention + residual + norm - let attn_output = self - .attention - .forward(hidden_states, rel_pos_emb, attention_mask); - let hidden_states = self - .attention_norm - .forward(&hidden_states.add_(&attn_output)); - - // FFN + residual + norm - let ffn_output = self.feed_forward.forward(&hidden_states); - self.output_norm.forward(&hidden_states.add_(&ffn_output)) - } } diff --git a/models/rbert/src/raw/mod.rs b/models/rbert/src/raw/mod.rs index 96220b401..b5425aa1d 100644 --- a/models/rbert/src/raw/mod.rs +++ b/models/rbert/src/raw/mod.rs @@ -19,6 +19,7 @@ use intermediate_layer::*; pub mod mdeberta; pub mod modern_bert; pub mod qwen; +mod utils; pub use mdeberta::{MDebertaConfig, MDebertaModel}; pub use modern_bert::{ModernBertConfig, ModernBertModel}; diff --git a/models/rbert/src/raw/modern_bert/attention.rs b/models/rbert/src/raw/modern_bert/attention.rs index b2c6b2e2e..43eb58bf1 100644 --- a/models/rbert/src/raw/modern_bert/attention.rs +++ b/models/rbert/src/raw/modern_bert/attention.rs @@ -40,30 +40,28 @@ impl ModernBertAttention { rope_cache: &RopeCache, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { - let [b_sz, seq_len, _hidden_size] = hidden_states.shape(); let hidden_size = self.num_heads * self.head_dim; // Compute fused QKV projection: [batch, seq_len, 3 * hidden_size] let qkv = hidden_states.q_mat_mul(&self.wqkv).to_concrete(); - // Split into Q, K, V - each [batch, seq_len, hidden_size] - let query_states = qkv - .narrow(2, 0, hidden_size) - .reshape([b_sz, seq_len, self.num_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); - - let key_states = qkv - .narrow(2, hidden_size, hidden_size) - .reshape([b_sz, seq_len, self.num_kv_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); - - let value_states = qkv - .narrow(2, 2 * hidden_size, hidden_size) - .reshape([b_sz, seq_len, self.num_kv_heads, self.head_dim]) - .transpose(1, 2) - .to_concrete(); + // Split into Q, K, V - each [batch, num_heads (or kv_heads), seq_len, head_dim] + use super::super::utils::split_heads; + let query_states = split_heads( + &qkv.narrow(2, 0, hidden_size).to_concrete(), + self.num_heads, + self.head_dim, + ); + let key_states = split_heads( + &qkv.narrow(2, hidden_size, hidden_size).to_concrete(), + self.num_kv_heads, + self.head_dim, + ); + let value_states = split_heads( + &qkv.narrow(2, 2 * hidden_size, hidden_size).to_concrete(), + self.num_kv_heads, + self.head_dim, + ); // Apply RoPE to Q and K let (query_states, key_states) = rope_cache.forward(&query_states, &key_states, 0); @@ -71,14 +69,7 @@ impl ModernBertAttention { // Scaled dot-product attention let scale = 1.0 / (self.head_dim as f32).sqrt(); - // Convert attention mask for flash attention if provided - const MASK_NEG_VALUE: f32 = -10000.0; - let mask: Option> = attention_mask.map(|m| { - let mask_f32: Tensor<2, f32> = m.cast(); - let zeros = mask_f32.zeros_like(); - let ones = (zeros + 1.0f32).to_concrete(); - ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete() - }); + let mask = attention_mask.map(super::super::utils::attention_mask_to_bias); let attn_output = query_states.flash_attention( &key_states, @@ -87,13 +78,8 @@ impl ModernBertAttention { mask.as_ref().map(|m| (m, fusor::MaskKind::BatchKeyMask)), ); - // Reshape and project output - let attn_output = attn_output.transpose(1, 2); - let attn_output = attn_output - .to_concrete() - .reshape([b_sz, seq_len, hidden_size]) - .to_concrete(); - + // Merge heads and project output + let attn_output = super::super::utils::merge_heads(&attn_output); attn_output.q_mat_mul(&self.wo) } } diff --git a/models/rbert/src/raw/modern_bert/config.rs b/models/rbert/src/raw/modern_bert/config.rs index 5c6960bf5..5600d08dc 100644 --- a/models/rbert/src/raw/modern_bert/config.rs +++ b/models/rbert/src/raw/modern_bert/config.rs @@ -2,6 +2,8 @@ use fusor::{Result, VarBuilder}; +use super::super::utils::{load_f32_or, load_u32, load_u32_or}; + /// Configuration for ModernBERT loaded from GGUF metadata. #[derive(Debug, Clone)] pub struct ModernBertConfig { @@ -28,29 +30,10 @@ pub struct ModernBertConfig { impl ModernBertConfig { /// Load configuration from GGUF metadata. pub fn from_gguf(vb: &VarBuilder) -> Result { - let num_heads = vb - .get_metadata(".attention.head_count") - .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| { - fusor::Error::msg("Missing required GGUF metadata: .attention.head_count") - })? as usize; - - let num_kv_heads = vb - .get_metadata(".attention.head_count_kv") - .and_then(|v| v.to_u32().ok()) - .unwrap_or(num_heads as u32) as usize; - - let num_layers = vb - .get_metadata(".block_count") - .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: .block_count"))? - as usize; - - let hidden_size = vb - .get_metadata(".embedding_length") - .and_then(|v| v.to_u32().ok()) - .ok_or_else(|| fusor::Error::msg("Missing required GGUF metadata: .embedding_length"))? - as usize; + let num_heads = load_u32(vb, ".attention.head_count")? as usize; + let num_kv_heads = load_u32_or(vb, ".attention.head_count_kv", num_heads as u32) as usize; + let num_layers = load_u32(vb, ".block_count")? as usize; + let hidden_size = load_u32(vb, ".embedding_length")? as usize; if hidden_size % num_heads != 0 { return Err(fusor::Error::msg(format!( @@ -58,33 +41,20 @@ impl ModernBertConfig { ))); } - let intermediate_size = vb - .get_metadata(".feed_forward_length") - .and_then(|v| v.to_u32().ok()) - .unwrap_or((hidden_size * 4) as u32) as usize; - - let context_length = vb - .get_metadata(".context_length") - .and_then(|v| v.to_u32().ok()) - .unwrap_or(8192) as usize; - - let rope_theta = vb - .get_metadata(".rope.freq_base") - .and_then(|v| v.to_f32().ok()) - .unwrap_or(10000.0); - - let norm_eps = vb - .get_metadata(".attention.layer_norm_rms_epsilon") - .and_then(|v| v.to_f32().ok()) - .unwrap_or(1e-6); + let intermediate_size = + load_u32_or(vb, ".feed_forward_length", (hidden_size * 4) as u32) as usize; + let context_length = load_u32_or(vb, ".context_length", 8192) as usize; + let rope_theta = load_f32_or(vb, ".rope.freq_base", 10000.0); + let norm_eps = load_f32_or(vb, ".attention.layer_norm_rms_epsilon", 1e-6); - // Use attention.key_length for head dimension - // Fall back to hidden_size / num_heads if not present - let head_dimension = vb - .get_metadata(".attention.key_length") - .and_then(|v| v.to_u32().ok()) - .map(|x| x as usize) - .unwrap_or_else(|| hidden_size / num_heads); + // Use attention.key_length for head dimension; fall back to + // hidden_size / num_heads if not present. + let head_dimension = load_u32_or(vb, ".attention.key_length", 0); + let head_dimension = if head_dimension == 0 { + hidden_size / num_heads + } else { + head_dimension as usize + }; Ok(Self { num_heads, diff --git a/models/rbert/src/raw/self_attention.rs b/models/rbert/src/raw/self_attention.rs index bb6a0faa7..2ac816614 100644 --- a/models/rbert/src/raw/self_attention.rs +++ b/models/rbert/src/raw/self_attention.rs @@ -34,15 +34,7 @@ impl BertSelfAttention { } pub(crate) fn transpose_for_scores(&self, xs: &Tensor<3, f32>) -> Tensor<4, f32> { - let shape = xs.shape(); - let new_x_shape = [ - shape[0], - shape[1], - self.num_attention_heads, - self.attention_head_size, - ]; - - xs.reshape(new_x_shape).transpose(1, 2).to_concrete() + super::utils::split_heads(xs, self.num_attention_heads, self.attention_head_size) } pub(crate) fn forward( @@ -60,13 +52,7 @@ impl BertSelfAttention { let value_layer = self.transpose_for_scores(&value_layer); let scale = 1.0 / (self.attention_head_size as f32).sqrt(); - const MASK_NEG_VALUE: f32 = -10000.0; - let mask: Option> = attention_mask.map(|m| { - let mask_f32: Tensor<2, f32> = m.cast(); - let zeros = mask_f32.zeros_like(); - let ones = (zeros + 1.0f32).to_concrete(); - ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete() - }); + let mask = attention_mask.map(super::utils::attention_mask_to_bias); let context_layer = { let _enter_sm = self.span_softmax.enter(); @@ -101,13 +87,7 @@ impl BertSelfAttention { let value_layer = self.transpose_for_scores(&value_layer); let scale = 1.0 / (self.attention_head_size as f32).sqrt(); - const MASK_NEG_VALUE: f32 = -10000.0; - let mask: Option> = attention_mask.map(|m| { - let mask_f32: Tensor<2, f32> = m.cast(); - let zeros = mask_f32.zeros_like(); - let ones = (zeros + 1.0f32).to_concrete(); - ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete() - }); + let mask = attention_mask.map(super::utils::attention_mask_to_bias); let context_layer = { let _enter_sm = self.span_softmax.enter(); diff --git a/models/rbert/src/raw/utils.rs b/models/rbert/src/raw/utils.rs new file mode 100644 index 000000000..16b9b70a4 --- /dev/null +++ b/models/rbert/src/raw/utils.rs @@ -0,0 +1,73 @@ +//! Shared helpers used by the raw encoder implementations. + +use fusor::{Result, Tensor, VarBuilder}; + +/// Reshape `[batch, seq_len, num_heads * head_dim]` into +/// `[batch, num_heads, seq_len, head_dim]` for multi-head attention. +pub(crate) fn split_heads( + tensor: &Tensor<3, f32>, + num_heads: usize, + head_dim: usize, +) -> Tensor<4, f32> { + let [batch, seq_len, _] = tensor.shape(); + tensor + .reshape([batch, seq_len, num_heads, head_dim]) + .transpose(1, 2) + .to_concrete() +} + +/// Inverse of [`split_heads`]: collapse +/// `[batch, num_heads, seq_len, head_dim]` back to +/// `[batch, seq_len, num_heads * head_dim]`. +pub(crate) fn merge_heads(tensor: &Tensor<4, f32>) -> Tensor<3, f32> { + let [batch, num_heads, seq_len, head_dim] = tensor.shape(); + tensor + .transpose(1, 2) + .to_concrete() + .reshape([batch, seq_len, num_heads * head_dim]) + .to_concrete() +} + +/// Large negative bias applied to masked-out positions so that, after +/// softmax, their probability collapses to ~0. +pub(crate) const MASK_NEG_VALUE: f32 = -10000.0; + +/// Convert a `[batch, seq]` boolean attention mask (1 = attend, 0 = pad) +/// into an additive bias tensor of the same shape (0 for real tokens, +/// [`MASK_NEG_VALUE`] for padding). +pub(crate) fn attention_mask_to_bias(mask: &Tensor<2, u32>) -> Tensor<2, f32> { + let mask_f32: Tensor<2, f32> = mask.cast(); + let zeros = mask_f32.zeros_like(); + let ones = (zeros + 1.0f32).to_concrete(); + ((ones - mask_f32) * MASK_NEG_VALUE).to_concrete() +} + +/// Read a required `u32` GGUF metadata value. The `.`-prefix on keys is +/// interpreted by fusor as a suffix match, so callers pass architecture-agnostic +/// keys like `.attention.head_count`. +pub(crate) fn load_u32(vb: &VarBuilder, key: &str) -> Result { + vb.get_metadata(key) + .and_then(|v| v.to_u32().ok()) + .ok_or_else(|| fusor::Error::msg(format!("Missing required GGUF metadata: {key}"))) +} + +/// Read an optional `u32` GGUF metadata value, falling back to `default`. +pub(crate) fn load_u32_or(vb: &VarBuilder, key: &str, default: u32) -> u32 { + vb.get_metadata(key) + .and_then(|v| v.to_u32().ok()) + .unwrap_or(default) +} + +/// Read an optional `f32` GGUF metadata value, falling back to `default`. +pub(crate) fn load_f32_or(vb: &VarBuilder, key: &str, default: f32) -> f32 { + vb.get_metadata(key) + .and_then(|v| v.to_f32().ok()) + .unwrap_or(default) +} + +/// Read an optional `bool` GGUF metadata value, falling back to `default`. +pub(crate) fn load_bool_or(vb: &VarBuilder, key: &str, default: bool) -> bool { + vb.get_metadata(key) + .and_then(|v| v.to_bool().ok()) + .unwrap_or(default) +} From 2bb8ac678ff666e3f1cda060b02e082f7a50b2cc Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 21:20:19 -0500 Subject: [PATCH 10/34] remove some of the debug prints --- models/rbert/src/raw/mdeberta/attention.rs | 10 -- models/rbert/src/raw/mdeberta/model.rs | 110 +----------- models/rgliner/src/raw/bilstm.rs | 10 -- models/rgliner/src/raw/joint_scorer.rs | 125 +------------ models/rgliner/src/raw/text_encoder.rs | 9 - models/rgliner/src/relex.rs | 197 --------------------- 6 files changed, 6 insertions(+), 455 deletions(-) diff --git a/models/rbert/src/raw/mdeberta/attention.rs b/models/rbert/src/raw/mdeberta/attention.rs index cd82b80ec..157e1f2ac 100644 --- a/models/rbert/src/raw/mdeberta/attention.rs +++ b/models/rbert/src/raw/mdeberta/attention.rs @@ -34,21 +34,11 @@ impl RelativePositionEmbedding { // GGUF stores shape as [hidden_size, 2*max_pos] but we need [2*max_pos, hidden_size] let [dim0, dim1] = embeddings_raw.shape(); let embeddings = if dim0 > dim1 { - // Shape is [hidden_size, positions] - need to transpose - #[cfg(debug_assertions)] - eprintln!( - "[DEBUG] Transposing rel_pos_embd from [{}, {}] to [{}, {}]", - dim0, dim1, dim1, dim0 - ); embeddings_raw.transpose(0, 1).to_concrete() } else { embeddings_raw }; - #[cfg(debug_assertions)] - eprintln!("[DEBUG] RelativePositionEmbedding loaded: shape={:?}, max_relative_positions={}, has_layer_norm={}", - embeddings.shape(), max_relative_positions, layer_norm.is_some()); - Ok(Self { embeddings, layer_norm, diff --git a/models/rbert/src/raw/mdeberta/model.rs b/models/rbert/src/raw/mdeberta/model.rs index 63159de62..037fa7b7e 100644 --- a/models/rbert/src/raw/mdeberta/model.rs +++ b/models/rbert/src/raw/mdeberta/model.rs @@ -1,8 +1,5 @@ //! mDeBERTa-v3 encoder model. -#[cfg(debug_assertions)] -use pollster; - use fusor::layers::{Embedding, LayerNorm, Linear}; use fusor::{Device, Result, Tensor, VarBuilder}; @@ -38,15 +35,11 @@ impl MDebertaModel { pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { let config = MDebertaConfig::from_gguf(vb)?; - // Load token embeddings let token_embeddings = Embedding::load(device, &mut vb.pp("token_embd"))?; - - // Load embedding LayerNorm let embedding_norm = LayerNorm::load(device, &mut vb.pp("embd_norm"), config.norm_eps)?; - // Load relative position embeddings with LayerNorm - // The output_norm in GGUF is the LayerNorm for relative position embeddings - // Load norm first to avoid borrow issues + // The `output_norm` tensor in GGUF is the LayerNorm for the relative + // position embeddings. Load it first to avoid borrow issues. let rel_pos_norm = LayerNorm::load(device, &mut vb.pp("output_norm"), config.norm_eps).ok(); let rel_pos_embedding = RelativePositionEmbedding::load_with_norm( device, @@ -55,7 +48,6 @@ impl MDebertaModel { config.max_relative_positions, )?; - // Load transformer layers let mut layers = Vec::with_capacity(config.num_layers); for i in 0..config.num_layers { let layer = MDebertaLayer::load( @@ -72,15 +64,6 @@ impl MDebertaModel { // which use DeBERTa-v3-large at 1024-dim and project down to 768). let output_proj = Linear::load(device, &mut vb.pp("output_proj")).ok(); - #[cfg(debug_assertions)] - if let Some(ref p) = output_proj { - eprintln!( - "[DEBUG] Encoder output projection loaded: {} -> {}", - p.in_features(), - p.out_features() - ); - } - Ok(Self { token_embeddings, embedding_norm, @@ -109,104 +92,21 @@ impl MDebertaModel { let _enter = self.span.enter(); let [_batch_size, seq_len] = input_ids.shape(); - // Get token embeddings - let mut hidden_states = self.token_embeddings.forward(input_ids); - - #[cfg(debug_assertions)] - { - let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); - let slice = data.as_slice(); - let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = - (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); - eprintln!( - "[DEBUG] After token_embeddings: mean={:.6}, std={:.6}", - mean, std - ); - } + let hidden_states = self.token_embeddings.forward(input_ids); + let mut hidden_states = self.embedding_norm.forward(&hidden_states); - // Apply embedding LayerNorm - hidden_states = self.embedding_norm.forward(&hidden_states); - - #[cfg(debug_assertions)] - { - let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); - let slice = data.as_slice(); - let hidden_size = self.config.hidden_size; - let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = - (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); - eprintln!( - "[DEBUG] After embedding_norm: mean={:.6}, std={:.6}", - mean, std - ); - // Print raw embeddings at <> positions (1, 3, 5) and others - eprintln!("[DEBUG] Raw embeddings at positions (first 5 values):"); - for pos in [0, 1, 2, 3, 4, 5, 10, 17] { - if pos < seq_len { - let start = pos * hidden_size; - let vals: Vec = (0..5).map(|i| slice[start + i]).collect(); - eprintln!(" pos {}: {:?}", pos, vals); - } - } - } - - // Compute relative position indices and get embedding table let rel_indices = self .rel_pos_embedding .compute_relative_indices(seq_len, &self.device); let rel_pos_emb = self.rel_pos_embedding.get_embeddings(); - // Pass through transformer layers with proper position attention - for (i, layer) in self.layers.iter().enumerate() { + for layer in &self.layers { hidden_states = layer.forward_with_rel(&hidden_states, &rel_pos_emb, &rel_indices, attention_mask); - - #[cfg(debug_assertions)] - if i == 0 || i == 11 { - let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); - let slice = data.as_slice(); - let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() - / slice.len() as f32) - .sqrt(); - eprintln!( - "[DEBUG] After layer {}: mean={:.6}, std={:.6}", - i, mean, std - ); - } } - #[cfg(debug_assertions)] - { - let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); - let slice = data.as_slice(); - let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = - (slice.iter().map(|x| (x - mean).powi(2)).sum::() / slice.len() as f32).sqrt(); - eprintln!( - "[DEBUG] Encoder output (pre-projection): mean={:.6}, std={:.6}", - mean, std - ); - } - - // Apply optional post-encoder projection (large variants). if let Some(ref proj) = self.output_proj { hidden_states = proj.forward(&hidden_states); - - #[cfg(debug_assertions)] - { - let data = pollster::block_on(hidden_states.clone().as_slice()).unwrap(); - let slice = data.as_slice(); - let mean: f32 = slice.iter().sum::() / slice.len() as f32; - let std: f32 = (slice.iter().map(|x| (x - mean).powi(2)).sum::() - / slice.len() as f32) - .sqrt(); - eprintln!( - "[DEBUG] Encoder output (post-projection): mean={:.6}, std={:.6}", - mean, std - ); - } } hidden_states diff --git a/models/rgliner/src/raw/bilstm.rs b/models/rgliner/src/raw/bilstm.rs index 85b0b6125..4b8ee6978 100644 --- a/models/rgliner/src/raw/bilstm.rs +++ b/models/rgliner/src/raw/bilstm.rs @@ -39,16 +39,6 @@ impl BiLstm { // hidden_size is 4*hidden (for i,f,g,o gates), so actual hidden = shape[0]/4 let hidden_size = weight_ih_f.shape()[0] / 4; - #[cfg(debug_assertions)] - { - eprintln!("[DEBUG] BiLstm loaded:"); - eprintln!(" weight_ih_f shape: {:?}", weight_ih_f.shape()); - eprintln!(" weight_hh_f shape: {:?}", weight_hh_f.shape()); - eprintln!(" bias_ih_f shape: {:?}", bias_ih_f.shape()); - eprintln!(" computed hidden_size: {}", hidden_size); - eprintln!(" output_dim: {}", 2 * hidden_size); - } - Ok(Self { weight_ih_f, weight_hh_f, diff --git a/models/rgliner/src/raw/joint_scorer.rs b/models/rgliner/src/raw/joint_scorer.rs index 0c621ca94..a20c8ae2d 100644 --- a/models/rgliner/src/raw/joint_scorer.rs +++ b/models/rgliner/src/raw/joint_scorer.rs @@ -31,35 +31,6 @@ impl JointScorer { let out_fc1 = Linear::load(device, &mut vb.pp("out_mlp.0"))?; let out_fc2 = Linear::load(device, &mut vb.pp("out_mlp.3"))?; - #[cfg(debug_assertions)] - { - eprintln!("[DEBUG] JointScorer loaded:"); - eprintln!( - " proj_label: in={}, out={}", - proj_label.in_features(), - proj_label.out_features() - ); - eprintln!( - " out_fc1: in={}, out={}", - out_fc1.in_features(), - out_fc1.out_features() - ); - eprintln!( - " out_fc2: in={}, out={}", - out_fc2.in_features(), - out_fc2.out_features() - ); - // Print fc2 bias values (these are the biases for O, B, I classes) - if let Some(bias) = out_fc2.bias() { - let bias_data = pollster::block_on(bias.clone().as_slice()).unwrap(); - let b = bias_data.as_slice(); - eprintln!( - " out_fc2 bias: O={:.6}, B={:.6}, I={:.6}", - b[0], b[1], b[2] - ); - } - } - Ok(Self { proj_token, proj_label, @@ -88,60 +59,24 @@ impl JointScorer { token_embs: &Tensor<3, f32>, label_embs: &Tensor<2, f32>, ) -> Tensor<4, f32> { - let [batch_size, seq_len, hidden_dim] = token_embs.shape(); + let [batch_size, seq_len, _hidden_dim] = token_embs.shape(); let [n_labels, _] = label_embs.shape(); - #[cfg(debug_assertions)] - eprintln!( - "[DEBUG] scorer.forward: batch={}, seq_len={}, hidden_dim={}, n_labels={}", - batch_size, seq_len, hidden_dim, n_labels - ); - // Project both token and label embeddings // token: [batch, seq, hidden] -> [batch, seq, hidden*2] let proj_tokens = self.proj_token.forward(token_embs); let [_, _, proj_dim] = proj_tokens.shape(); let half_proj = proj_dim / 2; - #[cfg(debug_assertions)] - { - // Verify proj_token computation - let input_data = token_embs.clone().as_slice().await.unwrap(); - let input_slice = input_data.as_slice(); - let output_data = proj_tokens.clone().as_slice().await.unwrap(); - let output_slice = output_data.as_slice(); - eprintln!("[DEBUG] proj_token input[0,0,:5]: {:?}", &input_slice[0..5]); - eprintln!( - "[DEBUG] proj_token output[0,0,:5]: {:?}", - &output_slice[0..5] - ); - eprintln!( - "[DEBUG] proj_token output[0,0,768:773]: {:?}", - &output_slice[768..773] - ); - } - // label: [n_labels, hidden] -> [n_labels, hidden*2] let label_embs_3d: Tensor<3, f32> = label_embs.unsqueeze(0).to_concrete(); let proj_labels = self.proj_label.forward(&label_embs_3d); let proj_labels: Tensor<2, f32> = proj_labels.squeeze(0).to_concrete(); - #[cfg(debug_assertions)] - eprintln!( - "[DEBUG] proj_tokens shape: [{}, {}, {}], proj_labels shape: [{}, {}], half_proj={}", - batch_size, seq_len, proj_dim, n_labels, proj_dim, half_proj - ); - // Split and combine: token_first + label_first + (token_second * label_second) // MLP input dimension = half_proj + half_proj + half_proj = 3 * half_proj let mlp_input_dim = 3 * half_proj; - #[cfg(debug_assertions)] - eprintln!( - "[DEBUG] mlp_input_dim={} (3 * {})", - mlp_input_dim, half_proj - ); - // Get raw data slices (without expansion - we'll handle broadcast manually) // proj_tokens shape: [batch, seq, proj_dim] // proj_labels shape: [n_labels, proj_dim] @@ -151,28 +86,6 @@ impl JointScorer { let tokens_slice = tokens_data.as_slice(); // [batch * seq * proj_dim] let labels_slice = labels_data.as_slice(); // [n_labels * proj_dim] - #[cfg(debug_assertions)] - { - // Check if label projections are different for each label - eprintln!("[DEBUG] Label projection check (first 5 values per label):"); - for l in 0..n_labels { - let start = l * proj_dim; - let vals: Vec = (0..5).map(|i| labels_slice[start + i]).collect(); - eprintln!(" label {}: {:?}", l, vals); - } - - // Check token projections for different tokens - eprintln!("[DEBUG] Token projection check (first 5 tokens, first 5 values):"); - for t in 0..5.min(seq_len) { - let start = t * proj_dim; - let vals: Vec = (0..5).map(|i| tokens_slice[start + i]).collect(); - let vals_second: Vec = (0..5) - .map(|i| tokens_slice[start + half_proj + i]) - .collect(); - eprintln!(" token {}: first={:?}, second={:?}", t, vals, vals_second); - } - } - // Build combined features with manual broadcasting // Output: [batch, seq, n_labels, mlp_input_dim] let total_elements = batch_size * seq_len * n_labels; @@ -235,29 +148,8 @@ impl JointScorer { label_embs: &Tensor<2, f32>, ) -> Tensor<4, f32> { let logits = self.forward(token_embs, label_embs).await; - let [_batch_size, seq_len, n_labels, num_classes] = logits.shape(); - let logits_data = logits.clone().as_slice().await.unwrap(); - #[cfg(debug_assertions)] - { - let data = logits_data.as_slice(); - eprintln!("[DEBUG] Raw logits (first 3 tokens, all labels) [start, end, inside]:"); - for s in 0..3.min(seq_len) { - for l in 0..n_labels { - let idx = s * n_labels * num_classes + l * num_classes; - eprintln!( - " token {} label {}: start={:.4}, end={:.4}, inside={:.4}", - s, - l, - data[idx], - data[idx + 1], - data[idx + 2] - ); - } - } - } - // Apply sigmoid to each value independently (NOT softmax). let data = logits_data.as_slice(); let sigmoid_data: Vec = data.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect(); @@ -283,21 +175,6 @@ impl PromptRepLayer { let fc1 = Linear::load(device, &mut vb.pp("0"))?; let fc2 = Linear::load(device, &mut vb.pp("3"))?; - #[cfg(debug_assertions)] - { - eprintln!("[DEBUG] PromptRepLayer loaded:"); - eprintln!( - " fc1: in={}, out={}", - fc1.in_features(), - fc1.out_features() - ); - eprintln!( - " fc2: in={}, out={}", - fc2.in_features(), - fc2.out_features() - ); - } - Ok(Self { fc1, fc2 }) } diff --git a/models/rgliner/src/raw/text_encoder.rs b/models/rgliner/src/raw/text_encoder.rs index 95207c6af..136ca66f9 100644 --- a/models/rgliner/src/raw/text_encoder.rs +++ b/models/rgliner/src/raw/text_encoder.rs @@ -24,15 +24,6 @@ impl TextEncoder { // Optional output projection (small/base/large v2.0 variants). let output_proj = Linear::load(device, &mut text_vb.pp("output_proj")).ok(); - #[cfg(debug_assertions)] - if let Some(ref p) = output_proj { - eprintln!( - "[DEBUG] TextEncoder output projection loaded: {} -> {}", - p.in_features(), - p.out_features() - ); - } - Ok(Self { model, output_proj }) } diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index 9d48223e1..1adff5c94 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -466,22 +466,6 @@ impl GlinerRelEx { .tokenizer .tokenize(text, entity_labels, relation_labels)?; - #[cfg(debug_assertions)] - { - eprintln!( - "[DEBUG] Tokenized: {} tokens, {} words", - tokenized.token_ids.len(), - tokenized.num_words - ); - eprintln!("[DEBUG] ent_positions: {:?}", tokenized.ent_positions); - eprintln!("[DEBUG] rel_positions: {:?}", tokenized.rel_positions); - eprintln!("[DEBUG] text_positions: {:?}", tokenized.text_positions); - eprintln!( - "[DEBUG] token_ids (first 20): {:?}", - &tokenized.token_ids[..20.min(tokenized.token_ids.len())] - ); - } - if tokenized.num_words == 0 { return Ok((Vec::new(), Vec::new())); } @@ -496,106 +480,18 @@ impl GlinerRelEx { // 3. Forward pass through encoder let encoder_output = self.encoder.forward(&token_ids, Some(&attention_mask)); - #[cfg(debug_assertions)] - { - let enc_data = encoder_output.clone().as_slice().await.unwrap(); - let enc_slice = enc_data.as_slice(); - let mean: f32 = enc_slice.iter().sum::() / enc_slice.len() as f32; - let variance: f32 = - enc_slice.iter().map(|x| (x - mean).powi(2)).sum::() / enc_slice.len() as f32; - eprintln!( - "[DEBUG] Encoder output stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", - mean, - variance, - enc_slice.iter().cloned().fold(f32::INFINITY, f32::min), - enc_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max) - ); - - // Check encoder output at specific positions (for <> tokens) - let hidden_size = self.config.hidden_size; - eprintln!("[DEBUG] Encoder output at <> positions (first 5 values):"); - for &pos in &tokenized.ent_positions { - let start = pos * hidden_size; - let vals: Vec = (0..5).map(|i| enc_slice[start + i]).collect(); - eprintln!(" pos {}: {:?}", pos, vals); - } - // Also check a few other positions for comparison - eprintln!("[DEBUG] Encoder output at other positions:"); - for &pos in &[0, 2, 4, 10, 17] { - if pos < tokenized.token_ids.len() { - let start = pos * hidden_size; - let vals: Vec = (0..5).map(|i| enc_slice[start + i]).collect(); - eprintln!(" pos {}: {:?}", pos, vals); - } - } - } - // 4. Extract word-level embeddings from encoder output, THEN apply BiLSTM // (Python applies BiLSTM to word-level embeddings, not the full token sequence.) let word_encoder_embs = self.gather_at_positions(&encoder_output, &tokenized.text_positions); let lstm_output = self.bilstm.forward(&word_encoder_embs).await; - #[cfg(debug_assertions)] - { - let lstm_data = lstm_output.clone().as_slice().await.unwrap(); - let lstm_slice = lstm_data.as_slice(); - let hidden_size = self.config.hidden_size; - eprintln!("[DEBUG] Word-level BiLSTM output (first 5 values per word):"); - for w in 0..tokenized.num_words { - let start = w * hidden_size; - let vals: Vec = (0..5).map(|i| lstm_slice[start + i]).collect(); - eprintln!(" word {}: {:?}", w, vals); - } - } - // 5. Extract label embeddings at marker positions from ENCODER output and project them // (Labels are extracted from encoder output, text tokens from BiLSTM output) // Entity label embeddings: hidden states at <> positions let ent_embs_raw = self.gather_at_positions(&encoder_output, &tokenized.ent_positions); - - #[cfg(debug_assertions)] - { - let raw_data = ent_embs_raw.clone().as_slice().await.unwrap(); - let raw_slice = raw_data.as_slice(); - let hidden_size = self.config.hidden_size; - eprintln!("[DEBUG] Raw entity embeddings check (first 5 values per label):"); - for l in 0..tokenized.ent_positions.len() { - let start = l * hidden_size; - let vals: Vec = (0..5).map(|i| raw_slice[start + i]).collect(); - eprintln!( - " label {} (pos {}): {:?}", - l, tokenized.ent_positions[l], vals - ); - } - } - let ent_embs = self.prompt_rep_layer.forward_3d(&ent_embs_raw); - #[cfg(debug_assertions)] - { - let ent_data = ent_embs.clone().as_slice().await.unwrap(); - let ent_slice = ent_data.as_slice(); - let hidden_size = self.config.hidden_size; - let mean: f32 = ent_slice.iter().sum::() / ent_slice.len() as f32; - let variance: f32 = - ent_slice.iter().map(|x| (x - mean).powi(2)).sum::() / ent_slice.len() as f32; - eprintln!( - "[DEBUG] Entity label embs stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", - mean, - variance, - ent_slice.iter().cloned().fold(f32::INFINITY, f32::min), - ent_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max) - ); - // Print projected values per label (compare with Python) - eprintln!("[DEBUG] Projected entity embeddings (first 5 values per label):"); - for l in 0..entity_labels.len() { - let start = l * hidden_size; - let vals: Vec = (0..5).map(|i| ent_slice[start + i]).collect(); - eprintln!(" {}: {:?}", entity_labels[l], vals); - } - } - // Relation label embeddings: raw hidden states at <> positions // (unlike entity labels, relation labels are NOT projected through prompt_rep_layer) let rel_embs = self.gather_at_positions(&encoder_output, &tokenized.rel_positions); @@ -603,22 +499,6 @@ impl GlinerRelEx { // 6. Text embeddings = BiLSTM output (already at word level) let text_embs = lstm_output.clone(); - #[cfg(debug_assertions)] - { - let text_data = text_embs.clone().as_slice().await.unwrap(); - let text_slice = text_data.as_slice(); - let mean: f32 = text_slice.iter().sum::() / text_slice.len() as f32; - let variance: f32 = text_slice.iter().map(|x| (x - mean).powi(2)).sum::() - / text_slice.len() as f32; - eprintln!( - "[DEBUG] Text token embs stats: mean={:.6}, var={:.6}, min={:.6}, max={:.6}", - mean, - variance, - text_slice.iter().cloned().fold(f32::INFINITY, f32::min), - text_slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max) - ); - } - // 7–8. Decode entities using the mode matching the trained head. let ent_embs_2d: Tensor<2, f32> = ent_embs.squeeze(0).to_concrete(); let entities = match self.span_mode { @@ -662,31 +542,6 @@ impl GlinerRelEx { .forward_for_spans(&text_embs, &entity_spans, &self.device); // span_reps shape: [num_entities, hidden] - #[cfg(debug_assertions)] - { - let sr_data = span_reps.clone().as_slice().await?; - let sr = sr_data.as_slice(); - let hidden = self.config.hidden_size; - eprintln!("[DEBUG] Span reps (first 5 values per entity):"); - for (i, e) in entities.iter().enumerate() { - let start = i * hidden; - let vals: Vec = (0..5).map(|k| sr[start + k]).collect(); - eprintln!( - " {} ({}, {}): {:?}", - e.text, e.start_word, e.end_word, vals - ); - } - // Print rel_embs - let re_data = rel_embs.clone().as_slice().await?; - let re = re_data.as_slice(); - eprintln!("[DEBUG] Rel embs (first 5 values per label):"); - for (i, l) in relation_labels.iter().enumerate() { - let start = i * hidden; - let vals: Vec = (0..5).map(|k| re[start + k]).collect(); - eprintln!(" {}: {:?}", l, vals); - } - } - let num_entities = entities.len(); let hidden_size = self.config.hidden_size; @@ -732,27 +587,6 @@ impl GlinerRelEx { let mut relations = Vec::new(); let threshold = self.config.relation_threshold; - #[cfg(debug_assertions)] - { - eprintln!( - "[DEBUG] Relation scoring: {} pairs, {} relations, threshold={}", - candidate_pairs.len(), - n_rels, - threshold - ); - for (pair_idx, &(h, t)) in candidate_pairs.iter().enumerate().take(6) { - let base = pair_idx * n_rels; - let raw: Vec = (0..n_rels) - .map(|c| rel_scores_slice.as_slice()[base + c]) - .collect(); - let sig: Vec = raw.iter().map(|x| 1.0 / (1.0 + (-x).exp())).collect(); - eprintln!( - " pair ({}->{}) [{} -> {}]: raw={:?}, sig={:?}", - h, t, entities[h].text, entities[t].text, raw, sig - ); - } - } - for (pair_idx, &(head_idx, tail_idx)) in candidate_pairs.iter().enumerate() { let base = pair_idx * n_rels; for rel_idx in 0..n_rels { @@ -824,14 +658,6 @@ impl GlinerRelEx { let logits_data = logits.clone().as_slice().await?; let logits_slice = logits_data.as_slice(); - #[cfg(debug_assertions)] - eprintln!( - "[DEBUG] markerV0 decoding: {} spans x {} labels, threshold={}", - spans.len(), - n_labels, - threshold, - ); - // Candidate (start, end, label, score) above threshold. let mut candidates: Vec<(usize, usize, usize, f32)> = Vec::new(); for (span_idx, &(s, e)) in spans.iter().enumerate() { @@ -898,29 +724,6 @@ impl GlinerRelEx { let threshold = self.config.entity_threshold; - #[cfg(debug_assertions)] - { - eprintln!( - "[DEBUG] Entity decoding: num_tokens={}, num_labels={}, threshold={}", - num_tokens, num_labels, threshold - ); - eprintln!("[DEBUG] Sigmoid scores [start, end, inside] (first 5 tokens):"); - for t in 0..5.min(num_tokens) { - for l in 0..num_labels { - let base = t * num_labels * 3 + l * 3; - eprintln!( - " token {} label {} ({}): start={:.4}, end={:.4}, inside={:.4}", - t, - l, - entity_labels[l], - scores[base], - scores[base + 1], - scores[base + 2] - ); - } - } - } - // Candidate spans: (start, end, label, score) let mut candidates: Vec<(usize, usize, usize, f32)> = Vec::new(); From aee29dcb9583151ab76923e72ad5edef1ca55aec Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 13 Apr 2026 22:06:24 -0500 Subject: [PATCH 11/34] cut out some dead code --- models/rbert/src/raw/mdeberta/attention.rs | 240 +++++++--------- models/rbert/src/raw/mdeberta/layer.rs | 15 +- models/rbert/src/raw/mdeberta/model.rs | 13 +- models/rbert/src/raw/mod.rs | 7 + models/rgliner/examples/basic.rs | 9 +- models/rgliner/src/raw/bilstm.rs | 315 ++++++++------------- models/rgliner/src/raw/joint_scorer.rs | 134 +++------ models/rgliner/src/raw/label_encoder.rs | 20 -- models/rgliner/src/raw/mod.rs | 2 - models/rgliner/src/raw/pair_projector.rs | 82 ------ models/rgliner/src/raw/relations_layer.rs | 120 -------- models/rgliner/src/raw/scorer.rs | 17 -- models/rgliner/src/raw/span_layer.rs | 10 - models/rgliner/src/raw/text_encoder.rs | 20 -- models/rgliner/src/relex.rs | 40 +-- models/rgliner/src/tokenization.rs | 29 +- 16 files changed, 301 insertions(+), 772 deletions(-) delete mode 100644 models/rgliner/src/raw/relations_layer.rs diff --git a/models/rbert/src/raw/mdeberta/attention.rs b/models/rbert/src/raw/mdeberta/attention.rs index 157e1f2ac..dcb025ad2 100644 --- a/models/rbert/src/raw/mdeberta/attention.rs +++ b/models/rbert/src/raw/mdeberta/attention.rs @@ -10,6 +10,19 @@ use fusor::layers::{LayerNorm, Linear}; use fusor::{Device, Result, Tensor, VarBuilder}; +/// Precomputed flat index tensors for the c2p and p2c gathers, valid for a +/// single `(b_sz, num_heads, seq_len)` combination. Built once per forward in +/// [`MDebertaModel::forward`] and threaded through every layer. +pub struct GatherIndices { + /// Flat source indices for c2p, shape `[b_sz * num_heads * seq_len * seq_len]`. + pub(crate) c2p: Tensor<1, u32>, + /// Flat source indices for p2c, shape `[b_sz * num_heads * seq_len * seq_len]`. + pub(crate) p2c: Tensor<1, u32>, + pub(crate) b_sz: usize, + pub(crate) num_heads: usize, + pub(crate) seq_len: usize, +} + /// Relative position embeddings for disentangled attention. pub struct RelativePositionEmbedding { /// Relative position embedding table [2*max_pos, hidden_size] @@ -73,43 +86,77 @@ impl RelativePositionEmbedding { } } - /// Compute relative position indices for a sequence. - /// Returns indices [seq_len, seq_len] where each entry is the relative position - /// index into the embedding table. - /// - /// Matches Python: rel_pos_ids = q_ids[:,None] - k_ids[None,:] = i - j - /// Then applies log bucketing with bucket_size=2*max_relative_positions (pos_ebd_size*2), - /// max_position = 2*max_relative_positions... actually: - /// - bucket_size = position_buckets = 256 (pos_ebd_size) - /// - max_position = max_relative_positions = 512 - pub fn compute_relative_indices(&self, seq_len: usize, device: &Device) -> Tensor<2, u32> { - // Python: bucket_size = position_buckets = 256, max_position = max_relative_positions = 512 - // att_span = pos_ebd_size = 256 (= bucket_size) - // The position embedding table has 2*pos_ebd_size = 512 entries - // After bucketing, rel_pos ranges in [-(pos_ebd_size), pos_ebd_size-1] approximately - // c2p_pos = clamp(rel_pos + att_span, 0, 2*att_span-1) -> [0, 2*pos_ebd_size-1] - let bucket_size = self.max_relative_positions as i32; // 256 (pos_ebd_size) - let max_position = 2 * bucket_size; // 512 (2*pos_ebd_size = max_relative_positions) - let att_span = bucket_size; // 256 - let num_positions = (2 * att_span) as i32; // 512 - - let mut indices = vec![0u32; seq_len * seq_len]; + /// Number of entries (`2 * max_relative_positions`) in the relative + /// position embedding table — this is the per-head "position dimension" of + /// the `c2p_all` / `p2c_all` attention scores before gathering. + pub fn num_positions(&self) -> usize { + self.embeddings.shape()[0] + } + /// Build the flat index tensors used to gather `c2p_all` and `p2c_all` + /// along the last two dims in a single on-device `index_select`. + /// + /// The gather semantics are: + /// c2p_out[b, h, i, j] = c2p_all[b, h, i, indices[i, j]] + /// p2c_out[b, h, i, j] = p2c_all[b, h, j, indices[i, j]] + /// + /// `indices[i, j]` is the DeBERTa log-bucketed relative position index. We + /// bake the `(b, h, i|j)` outer offsets into a single 1D `u32` tensor of + /// length `b_sz * num_heads * seq_len * seq_len`, so each gather becomes: + /// source.reshape([b*h*s*p]).index_select(0, flat_idx).reshape([b,h,s,s]) + pub fn compute_gather_indices( + &self, + b_sz: usize, + num_heads: usize, + seq_len: usize, + device: &Device, + ) -> GatherIndices { + let num_pos = self.num_positions(); + let bucket_size = self.max_relative_positions as i32; + let max_position = 2 * bucket_size; + let att_span = bucket_size; + let num_positions_i = (2 * att_span) as i32; + + // Raw relative-position indices, [seq_len, seq_len]. + let mut rel = vec![0u32; seq_len * seq_len]; for i in 0..seq_len { for j in 0..seq_len { - // Python: rel_pos = q - k = i - j let rel_pos = i as i32 - j as i32; - // Apply log bucketing let bucketed = Self::make_log_bucket_position(rel_pos, bucket_size, max_position); - // Shift to positive index: c2p_pos = clamp(bucketed + att_span, 0, 2*att_span-1) - let idx = (bucketed + att_span).clamp(0, num_positions - 1) as u32; - indices[i * seq_len + j] = idx; + let idx = (bucketed + att_span).clamp(0, num_positions_i - 1) as u32; + rel[i * seq_len + j] = idx; } } - Tensor::new(device, &indices) - .reshape([seq_len, seq_len]) - .to_concrete() + let total = b_sz * num_heads * seq_len * seq_len; + let mut c2p = vec![0u32; total]; + let mut p2c = vec![0u32; total]; + for b in 0..b_sz { + for h in 0..num_heads { + let bh_offset = ((b * num_heads + h) * seq_len) * num_pos; + for i in 0..seq_len { + let row_offset_c2p = bh_offset + i * num_pos; + for j in 0..seq_len { + let rel_idx = rel[i * seq_len + j] as usize; + // c2p: key dim = i + c2p[((b * num_heads + h) * seq_len + i) * seq_len + j] = + (row_offset_c2p + rel_idx.min(num_pos - 1)) as u32; + // p2c: key dim = j (different row offset) + let row_offset_p2c = bh_offset + j * num_pos; + p2c[((b * num_heads + h) * seq_len + i) * seq_len + j] = + (row_offset_p2c + rel_idx.min(num_pos - 1)) as u32; + } + } + } + } + + GatherIndices { + c2p: Tensor::new(device, &c2p), + p2c: Tensor::new(device, &p2c), + b_sz, + num_heads, + seq_len, + } } /// Get the raw relative position embedding table (normalized). @@ -172,13 +219,13 @@ impl MDebertaAttention { /// # Arguments /// * `hidden_states` - Input [batch, seq_len, hidden_size] /// * `rel_pos_emb` - Relative position embedding table [2*max_pos, hidden_size] - /// * `rel_pos_indices` - Relative position indices [seq_len, seq_len] + /// * `gather_idx` - Precomputed flat indices for the c2p / p2c gathers. /// * `attention_mask` - Optional attention mask [batch, seq_len] pub fn forward_with_indices( &self, hidden_states: &Tensor<3, f32>, rel_pos_emb: &Tensor<2, f32>, - rel_pos_indices: &Tensor<2, u32>, + gather_idx: &GatherIndices, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { use super::super::utils::split_heads; @@ -222,13 +269,13 @@ impl MDebertaAttention { // c2p = Q @ pos_key^T -> [batch, heads, seq, 2*max_pos] // Then gather based on relative positions let c2p_all = query.mat_mul(&pos_key.transpose(2, 3)); - let c2p_scores = self.gather_c2p(&c2p_all, rel_pos_indices); + let c2p_scores = gather_by_flat_index(&c2p_all, gather_idx, &gather_idx.c2p); // === Position-to-Content attention === // p2c = K @ pos_query^T -> [batch, heads, seq, 2*max_pos] // Then gather based on transposed relative positions let p2c_all = key.mat_mul(&pos_query.transpose(2, 3)); - let p2c_scores = self.gather_p2c(&p2c_all, rel_pos_indices); + let p2c_scores = gather_by_flat_index(&p2c_all, gather_idx, &gather_idx.p2c); // Combine: attention = (c2c + c2p + p2c) * scale let attn_scores = c2c_scores @@ -254,105 +301,24 @@ impl MDebertaAttention { let context = super::super::utils::merge_heads(&context); self.output.forward(&context) } +} - /// Gather c2p attention scores based on relative position indices. - /// - /// Input: c2p_all [batch, heads, seq_len, 2*max_pos] - scores to all positions - /// rel_pos_indices: [seq_len, seq_len] - index into position embeddings - /// - /// Output: [batch, heads, seq_len, seq_len] - gathered scores - fn gather_c2p( - &self, - c2p_all: &Tensor<4, f32>, - rel_pos_indices: &Tensor<2, u32>, - ) -> Tensor<4, f32> { - let [b_sz, num_heads, seq_len, _num_pos] = c2p_all.shape(); - let device = c2p_all.device(); - - // Get data slices - let c2p_data = pollster::block_on(c2p_all.clone().as_slice()).unwrap(); - let indices_data = pollster::block_on(rel_pos_indices.clone().as_slice()).unwrap(); - let c2p = c2p_data.as_slice(); - let indices = indices_data.as_slice(); - let num_pos = _num_pos; - - let mut gathered = vec![0.0f32; b_sz * num_heads * seq_len * seq_len]; - - for b in 0..b_sz { - for h in 0..num_heads { - for i in 0..seq_len { - for j in 0..seq_len { - // Index into c2p_all: [b, h, i, rel_pos[i,j]] - let rel_idx = indices[i * seq_len + j] as usize; - let c2p_idx = b * num_heads * seq_len * num_pos - + h * seq_len * num_pos - + i * num_pos - + rel_idx; - let out_idx = b * num_heads * seq_len * seq_len - + h * seq_len * seq_len - + i * seq_len - + j; - gathered[out_idx] = c2p[c2p_idx]; - } - } - } - } - - Tensor::new(&device, &gathered) - .reshape([b_sz, num_heads, seq_len, seq_len]) - .to_concrete() - } - - /// Gather p2c attention scores. - /// - /// Python derivation: - /// - r_pos = relative_pos (since seq_q == seq_k) - /// - p2c_pos[i,j] = clamp(-r_pos[i,j] + att_span, 0, 2*att_span-1) - /// = clamp(-(i-j) + att_span) = clamp((j-i) + att_span) - /// - gather_out[b, m, n] = p2c_att[b, m, p2c_pos[m, n]] - /// - final[b, i, j] = gather_out[b, j, i] (after transpose) - /// = p2c_att[b, j, p2c_pos[j, i]] - /// = p2c_att[b, j, clamp(i - j + att_span)] - /// = p2c_all[b, j, indices[i, j]] (using indices[i,j] = bucketed(i-j) + att_span) - fn gather_p2c( - &self, - p2c_all: &Tensor<4, f32>, - rel_pos_indices: &Tensor<2, u32>, - ) -> Tensor<4, f32> { - let [b_sz, num_heads, seq_len, num_pos] = p2c_all.shape(); - let device = p2c_all.device(); - - let p2c_data = pollster::block_on(p2c_all.clone().as_slice()).unwrap(); - let indices_data = pollster::block_on(rel_pos_indices.clone().as_slice()).unwrap(); - let p2c = p2c_data.as_slice(); - let indices = indices_data.as_slice(); - - let mut gathered = vec![0.0f32; b_sz * num_heads * seq_len * seq_len]; - - for b in 0..b_sz { - for h in 0..num_heads { - for i in 0..seq_len { - for j in 0..seq_len { - // final[b, i, j] = p2c_all[b, j, indices[i, j]] - let rel_idx = (indices[i * seq_len + j] as usize).min(num_pos - 1); - let p2c_idx = b * num_heads * seq_len * num_pos - + h * seq_len * num_pos - + j * num_pos // key dim = j - + rel_idx; - let out_idx = b * num_heads * seq_len * seq_len - + h * seq_len * seq_len - + i * seq_len - + j; - gathered[out_idx] = p2c[p2c_idx]; - } - } - } - } - - Tensor::new(&device, &gathered) - .reshape([b_sz, num_heads, seq_len, seq_len]) - .to_concrete() - } +/// On-device gather used by both c2p and p2c. `flat_idx` encodes the full +/// `((b*H + h)*S + i_or_j) * P + rel[i, j]` source offset so that the gather +/// reduces to flatten → `index_select` → reshape. +fn gather_by_flat_index( + src: &Tensor<4, f32>, + shape: &GatherIndices, + flat_idx: &Tensor<1, u32>, +) -> Tensor<4, f32> { + let [b, h, s, p] = src.shape(); + debug_assert_eq!(b, shape.b_sz); + debug_assert_eq!(h, shape.num_heads); + debug_assert_eq!(s, shape.seq_len); + src.reshape([b * h * s * p]) + .index_select(0, flat_idx) + .reshape([b, h, s, s]) + .to_concrete() } /// Shared relative position embedding layer (used across all layers in DeBERTa). @@ -371,19 +337,15 @@ impl DisentangledSelfAttention { Ok(Self { attention }) } - /// Forward with relative position indices and embedding table. + /// Forward with precomputed gather indices and relative embedding table. pub fn forward_with_rel( &self, hidden_states: &Tensor<3, f32>, rel_pos_emb: &Tensor<2, f32>, - rel_pos_indices: &Tensor<2, u32>, + gather_idx: &GatherIndices, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { - self.attention.forward_with_indices( - hidden_states, - rel_pos_emb, - rel_pos_indices, - attention_mask, - ) + self.attention + .forward_with_indices(hidden_states, rel_pos_emb, gather_idx, attention_mask) } } diff --git a/models/rbert/src/raw/mdeberta/layer.rs b/models/rbert/src/raw/mdeberta/layer.rs index d22bf6111..0e098c594 100644 --- a/models/rbert/src/raw/mdeberta/layer.rs +++ b/models/rbert/src/raw/mdeberta/layer.rs @@ -3,7 +3,7 @@ use fusor::layers::LayerNorm; use fusor::{Device, Result, Tensor, VarBuilder}; -use super::attention::DisentangledSelfAttention; +use super::attention::{DisentangledSelfAttention, GatherIndices}; use super::feed_forward::MDebertaFeedForward; /// A single mDeBERTa transformer layer. @@ -47,22 +47,19 @@ impl MDebertaLayer { /// # Arguments /// * `hidden_states` - Input [batch, seq_len, hidden_size] /// * `rel_pos_emb` - Relative position embedding table [2*max_pos, hidden_size] - /// * `rel_pos_indices` - Relative position indices [seq_len, seq_len] + /// * `gather_idx` - Precomputed flat indices for the c2p / p2c gathers. /// * `attention_mask` - Optional attention mask [batch, seq_len] pub fn forward_with_rel( &self, hidden_states: &Tensor<3, f32>, rel_pos_emb: &Tensor<2, f32>, - rel_pos_indices: &Tensor<2, u32>, + gather_idx: &GatherIndices, attention_mask: Option<&Tensor<2, u32>>, ) -> Tensor<3, f32> { // Self-attention + residual + norm - let attn_output = self.attention.forward_with_rel( - hidden_states, - rel_pos_emb, - rel_pos_indices, - attention_mask, - ); + let attn_output = + self.attention + .forward_with_rel(hidden_states, rel_pos_emb, gather_idx, attention_mask); let hidden_states = self .attention_norm .forward(&hidden_states.add_(&attn_output)); diff --git a/models/rbert/src/raw/mdeberta/model.rs b/models/rbert/src/raw/mdeberta/model.rs index 037fa7b7e..fb6dc4486 100644 --- a/models/rbert/src/raw/mdeberta/model.rs +++ b/models/rbert/src/raw/mdeberta/model.rs @@ -92,17 +92,22 @@ impl MDebertaModel { let _enter = self.span.enter(); let [_batch_size, seq_len] = input_ids.shape(); + let [b_sz, _] = input_ids.shape(); let hidden_states = self.token_embeddings.forward(input_ids); let mut hidden_states = self.embedding_norm.forward(&hidden_states); - let rel_indices = self - .rel_pos_embedding - .compute_relative_indices(seq_len, &self.device); + // Compute the flat gather indices once per forward; every layer shares them. + let gather_idx = self.rel_pos_embedding.compute_gather_indices( + b_sz, + self.config.num_heads, + seq_len, + &self.device, + ); let rel_pos_emb = self.rel_pos_embedding.get_embeddings(); for layer in &self.layers { hidden_states = - layer.forward_with_rel(&hidden_states, &rel_pos_emb, &rel_indices, attention_mask); + layer.forward_with_rel(&hidden_states, &rel_pos_emb, &gather_idx, attention_mask); } if let Some(ref proj) = self.output_proj { diff --git a/models/rbert/src/raw/mod.rs b/models/rbert/src/raw/mod.rs index b5425aa1d..5e56add7c 100644 --- a/models/rbert/src/raw/mod.rs +++ b/models/rbert/src/raw/mod.rs @@ -16,8 +16,11 @@ mod self_output; use self_output::*; mod intermediate_layer; use intermediate_layer::*; +/// mDeBERTa-v3 raw encoder. pub mod mdeberta; +/// ModernBERT raw encoder. pub mod modern_bert; +/// Qwen embedding model raw implementation. pub mod qwen; mod utils; @@ -29,10 +32,14 @@ use fusor::{Device, Result, Tensor, VarBuilder}; use serde::Deserialize; use std::fmt::Debug; +/// Non-linear activation applied between the intermediate and output dense +/// layers of each [`BertLayer`]. #[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)] #[serde(rename_all = "lowercase")] pub enum HiddenAct { + /// Gaussian Error Linear Unit. Gelu, + /// Rectified Linear Unit. Relu, } diff --git a/models/rgliner/examples/basic.rs b/models/rgliner/examples/basic.rs index 40d09d55c..873e8ba51 100644 --- a/models/rgliner/examples/basic.rs +++ b/models/rgliner/examples/basic.rs @@ -23,14 +23,11 @@ async fn main() -> anyhow::Result<()> { println!("Model loaded!"); - let labels = ["person", "organization", "location"]; + let labels = ["person", "award", "date", "competitions", "teams"]; - // Test with multiple texts to see if the issue is consistent + // The Ronaldo paragraph from the v2.0 model card (short strings score poorly on v2.0). let texts = [ - "Apple Inc. was founded by Steve Jobs in California.", - "Microsoft Corporation is headquartered in Seattle.", - "Elon Musk is the CEO of Tesla.", - "Google was founded in Mountain View.", + "Cristiano Ronaldo dos Santos Aveiro (Portuguese pronunciation: [kɾiʃˈtjɐnu ʁɔˈnaldu]; born 5 February 1985) is a Portuguese professional footballer who plays as a forward for and captains both Saudi Pro League club Al Nassr and the Portugal national team.", ]; for text in texts { diff --git a/models/rgliner/src/raw/bilstm.rs b/models/rgliner/src/raw/bilstm.rs index 4b8ee6978..398b0442a 100644 --- a/models/rgliner/src/raw/bilstm.rs +++ b/models/rgliner/src/raw/bilstm.rs @@ -1,53 +1,69 @@ //! Bidirectional LSTM implementation for GLiNER token representation. //! //! The BiLSTM processes encoder output to capture bidirectional context -//! before span representation computation. +//! before span representation computation. The timestep loop is inherently +//! sequential, but every tensor operation inside the loop runs on the active +//! device — no `.as_slice()` round-trips or scalar Rust gate math. use fusor::{Device, Result, Tensor, VarBuilder}; +/// Per-direction LSTM parameters pre-arranged for matrix multiplication. +struct LstmDir { + // Shape [input_size, 4 * hidden] — transpose of the GGUF `weight_ih_l0*` layout. + w_ih_t: Tensor<2, f32>, + // Shape [hidden_size, 4 * hidden] — transpose of the GGUF `weight_hh_l0*` layout. + w_hh_t: Tensor<2, f32>, + // Shape [4 * hidden] — pre-summed `bias_ih + bias_hh`. + bias: Tensor<1, f32>, +} + +impl LstmDir { + fn load(device: &Device, vb: &mut VarBuilder, suffix: &str) -> Result { + let w_ih: Tensor<2, f32> = vb + .get(&format!("weight_ih_l0{suffix}"), device)? + .dequantize(); + let w_hh: Tensor<2, f32> = vb + .get(&format!("weight_hh_l0{suffix}"), device)? + .dequantize(); + let b_ih: Tensor<1, f32> = vb.get(&format!("bias_ih_l0{suffix}"), device)?.dequantize(); + let b_hh: Tensor<1, f32> = vb.get(&format!("bias_hh_l0{suffix}"), device)?.dequantize(); + + let w_ih_t = w_ih.transpose(0, 1).to_concrete(); + let w_hh_t = w_hh.transpose(0, 1).to_concrete(); + let bias = (b_ih + b_hh).to_concrete(); + + Ok(Self { + w_ih_t, + w_hh_t, + bias, + }) + } +} + /// Bidirectional LSTM layer for token representation. /// /// Processes transformer encoder output through forward and backward LSTMs /// and concatenates the outputs. pub struct BiLstm { - // Forward LSTM weights - weight_ih_f: Tensor<2, f32>, // [4*hidden, input_size] - weight_hh_f: Tensor<2, f32>, // [4*hidden, hidden_size] - bias_ih_f: Tensor<1, f32>, // [4*hidden] - bias_hh_f: Tensor<1, f32>, // [4*hidden] - // Backward LSTM weights - weight_ih_b: Tensor<2, f32>, - weight_hh_b: Tensor<2, f32>, - bias_ih_b: Tensor<1, f32>, - bias_hh_b: Tensor<1, f32>, + forward: LstmDir, + backward: LstmDir, hidden_size: usize, } impl BiLstm { /// Load BiLSTM weights from GGUF. pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { - let weight_ih_f: Tensor<2, f32> = vb.get("weight_ih_l0", device)?.dequantize(); - let weight_hh_f: Tensor<2, f32> = vb.get("weight_hh_l0", device)?.dequantize(); - let bias_ih_f: Tensor<1, f32> = vb.get("bias_ih_l0", device)?.dequantize(); - let bias_hh_f: Tensor<1, f32> = vb.get("bias_hh_l0", device)?.dequantize(); - - let weight_ih_b: Tensor<2, f32> = vb.get("weight_ih_l0_reverse", device)?.dequantize(); - let weight_hh_b: Tensor<2, f32> = vb.get("weight_hh_l0_reverse", device)?.dequantize(); - let bias_ih_b: Tensor<1, f32> = vb.get("bias_ih_l0_reverse", device)?.dequantize(); - let bias_hh_b: Tensor<1, f32> = vb.get("bias_hh_l0_reverse", device)?.dequantize(); + let forward = LstmDir::load(device, vb, "")?; + let backward = LstmDir::load(device, vb, "_reverse")?; - // hidden_size is 4*hidden (for i,f,g,o gates), so actual hidden = shape[0]/4 - let hidden_size = weight_ih_f.shape()[0] / 4; + // The forward weight matrix is [4*hidden, input_size] in GGUF layout, which + // after transpose becomes [input_size, 4*hidden]. The hidden dim is the + // last axis divided by 4. + let hidden_size = forward.w_ih_t.shape()[1] / 4; Ok(Self { - weight_ih_f, - weight_hh_f, - bias_ih_f, - bias_hh_f, - weight_ih_b, - weight_hh_b, - bias_ih_b, - bias_hh_b, + forward, + backward, hidden_size, }) } @@ -60,176 +76,89 @@ impl BiLstm { /// # Returns /// Output tensor [batch, seq_len, 2*hidden_size] pub async fn forward(&self, input: &Tensor<3, f32>) -> Tensor<3, f32> { - let [batch_size, seq_len, input_size] = input.shape(); + let [batch, seq_len, _input_size] = input.shape(); let device = input.device(); - let output_size = 2 * self.hidden_size; - - // Get all weight data upfront - let input_data = input.clone().as_slice().await.unwrap(); - let w_ih_f = self.weight_ih_f.clone().as_slice().await.unwrap(); - let w_hh_f = self.weight_hh_f.clone().as_slice().await.unwrap(); - let b_ih_f = self.bias_ih_f.clone().as_slice().await.unwrap(); - let b_hh_f = self.bias_hh_f.clone().as_slice().await.unwrap(); - let w_ih_b = self.weight_ih_b.clone().as_slice().await.unwrap(); - let w_hh_b = self.weight_hh_b.clone().as_slice().await.unwrap(); - let b_ih_b = self.bias_ih_b.clone().as_slice().await.unwrap(); - let b_hh_b = self.bias_hh_b.clone().as_slice().await.unwrap(); - - let mut output_data = vec![0.0f32; batch_size * seq_len * output_size]; - - for b in 0..batch_size { - // Forward LSTM - let forward_out = self.lstm_direction( - input_data.as_slice(), - b, - seq_len, - input_size, - w_ih_f.as_slice(), - w_hh_f.as_slice(), - b_ih_f.as_slice(), - b_hh_f.as_slice(), - false, - ); - - // Backward LSTM - let backward_out = self.lstm_direction( - input_data.as_slice(), - b, - seq_len, - input_size, - w_ih_b.as_slice(), - w_hh_b.as_slice(), - b_ih_b.as_slice(), - b_hh_b.as_slice(), - true, - ); - - // Concatenate forward and backward outputs - for t in 0..seq_len { - for i in 0..self.hidden_size { - let out_idx = b * seq_len * output_size + t * output_size; - output_data[out_idx + i] = forward_out[t * self.hidden_size + i]; - output_data[out_idx + self.hidden_size + i] = - backward_out[t * self.hidden_size + i]; - } - } - } - - Tensor::new(&device, &output_data) - .reshape([batch_size, seq_len, output_size]) - .to_concrete() - } - /// Single direction LSTM pass. - fn lstm_direction( - &self, - input_data: &[f32], - batch_idx: usize, - seq_len: usize, - input_size: usize, - w_ih: &[f32], - w_hh: &[f32], - b_ih: &[f32], - b_hh: &[f32], - reverse: bool, - ) -> Vec { - let hidden_size = self.hidden_size; - let mut h = vec![0.0f32; hidden_size]; - let mut c = vec![0.0f32; hidden_size]; - let mut outputs = vec![0.0f32; seq_len * hidden_size]; - - // Process sequence in order (or reverse) - let indices: Vec = if reverse { - (0..seq_len).rev().collect() - } else { - (0..seq_len).collect() - }; - - for (out_idx, &t) in indices.iter().enumerate() { - // Get input at time t for this batch - let x_start = batch_idx * seq_len * input_size + t * input_size; - let x = &input_data[x_start..x_start + input_size]; - - // Compute gates: i, f, g, o - let mut gates = vec![0.0f32; 4 * hidden_size]; - - for g in 0..(4 * hidden_size) { - let mut sum = b_ih[g] + b_hh[g]; - - // Input contribution: x @ W_ih^T - for i in 0..input_size { - sum += x[i] * w_ih[g * input_size + i]; - } - - // Hidden contribution: h @ W_hh^T - for j in 0..hidden_size { - sum += h[j] * w_hh[g * hidden_size + j]; - } - - gates[g] = sum; - } - - // Apply activations and compute new h, c - for i in 0..hidden_size { - let i_gate = sigmoid(gates[i]); - let f_gate = sigmoid(gates[hidden_size + i]); - let g_gate = tanh(gates[2 * hidden_size + i]); - let o_gate = sigmoid(gates[3 * hidden_size + i]); - - c[i] = f_gate * c[i] + i_gate * g_gate; - h[i] = o_gate * tanh(c[i]); - } - - // Store output in correct position - let store_pos = if reverse { - seq_len - 1 - out_idx - } else { - out_idx - }; - for i in 0..hidden_size { - outputs[store_pos * hidden_size + i] = h[i]; - } - } - - outputs - } - - /// Get output dimension (2 * hidden_size for bidirectional). - pub fn output_dim(&self) -> usize { - 2 * self.hidden_size - } + let fwd_out = run_direction(input, &self.forward, self.hidden_size, &device, false); + let bwd_out = run_direction(input, &self.backward, self.hidden_size, &device, true); - /// Get the hidden size of a single direction. - pub fn hidden_size(&self) -> usize { - self.hidden_size + // Concatenate forward and backward along the feature dim -> [batch, seq, 2*hidden] + Tensor::cat([fwd_out, bwd_out], 2) + .reshape([batch, seq_len, 2 * self.hidden_size]) + .to_concrete() } } -#[inline] -fn sigmoid(x: f32) -> f32 { - 1.0 / (1.0 + (-x).exp()) -} +/// Run one direction of the LSTM. Sequential over time; every timestep's gate +/// math stays on-device. +fn run_direction( + input: &Tensor<3, f32>, + dir: &LstmDir, + hidden_size: usize, + device: &Device, + reverse: bool, +) -> Tensor<3, f32> { + let [batch, seq_len, _] = input.shape(); + + let mut h: Tensor<2, f32> = Tensor::zeros(device, [batch, hidden_size]); + let mut c: Tensor<2, f32> = Tensor::zeros(device, [batch, hidden_size]); + + // outputs[t] holds the hidden state at timestep t, already unsqueezed on + // dim 1 so that a final cat along dim 1 yields [batch, seq_len, hidden]. + let mut outputs: Vec> = Vec::with_capacity(seq_len); + outputs.resize_with(seq_len, || Tensor::zeros(device, [batch, 1, hidden_size])); + + let bias_broadcast: Tensor<2, f32> = dir + .bias + .unsqueeze(0) + .broadcast_as([batch, 4 * hidden_size]) + .to_concrete(); + + let iter: Box> = if reverse { + Box::new((0..seq_len).rev()) + } else { + Box::new(0..seq_len) + }; + + for t in iter { + let x_t: Tensor<2, f32> = input + .narrow(1, t, 1) + .reshape([batch, input.shape()[2]]) + .to_concrete(); + + // gates_pre = x_t @ W_ih^T + h @ W_hh^T + bias, shape [batch, 4*hidden] + let gates_pre: Tensor<2, f32> = + (x_t.mat_mul(&dir.w_ih_t) + h.mat_mul(&dir.w_hh_t) + bias_broadcast.clone()) + .to_concrete(); + + let i_raw: Tensor<2, f32> = gates_pre.narrow(1, 0, hidden_size).to_concrete(); + let f_raw: Tensor<2, f32> = gates_pre.narrow(1, hidden_size, hidden_size).to_concrete(); + let g_raw: Tensor<2, f32> = gates_pre + .narrow(1, 2 * hidden_size, hidden_size) + .to_concrete(); + let o_raw: Tensor<2, f32> = gates_pre + .narrow(1, 3 * hidden_size, hidden_size) + .to_concrete(); + + let i_gate = sigmoid_2d(&i_raw); + let f_gate = sigmoid_2d(&f_raw); + let g_gate = g_raw.tanh(); + let o_gate = sigmoid_2d(&o_raw); + + c = (f_gate * c + i_gate * g_gate).to_concrete(); + h = (o_gate * c.clone().tanh()).to_concrete(); + + outputs[t] = h.clone().unsqueeze(1).to_concrete(); + } -#[inline] -fn tanh(x: f32) -> f32 { - x.tanh() + Tensor::cat(outputs, 1) + .reshape([batch, seq_len, hidden_size]) + .to_concrete() } -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_sigmoid() { - assert!((sigmoid(0.0) - 0.5).abs() < 1e-6); - assert!(sigmoid(10.0) > 0.99); - assert!(sigmoid(-10.0) < 0.01); - } - - #[test] - fn test_tanh() { - assert!(tanh(0.0).abs() < 1e-6); - assert!(tanh(10.0) > 0.99); - assert!(tanh(-10.0) < -0.99); - } +/// sigmoid via `0.5 * (tanh(x / 2) + 1)` — avoids needing scalar-left division +/// or a `recip` primitive, and keeps the computation on-device. +fn sigmoid_2d(x: &Tensor<2, f32>) -> Tensor<2, f32> { + let half = (x.clone() * 0.5f32).to_concrete(); + ((half.tanh() + 1.0f32) * 0.5f32).to_concrete() } diff --git a/models/rgliner/src/raw/joint_scorer.rs b/models/rgliner/src/raw/joint_scorer.rs index a20c8ae2d..04e477859 100644 --- a/models/rgliner/src/raw/joint_scorer.rs +++ b/models/rgliner/src/raw/joint_scorer.rs @@ -54,7 +54,7 @@ impl JointScorer { /// 2. Split each projection into two halves (first, second) /// 3. MLP input = concat(token_first, label_first, token_second * label_second) /// 4. This enables complex token-label interactions through the element-wise product - pub async fn forward( + pub fn forward( &self, token_embs: &Tensor<3, f32>, label_embs: &Tensor<2, f32>, @@ -62,74 +62,55 @@ impl JointScorer { let [batch_size, seq_len, _hidden_dim] = token_embs.shape(); let [n_labels, _] = label_embs.shape(); - // Project both token and label embeddings - // token: [batch, seq, hidden] -> [batch, seq, hidden*2] + // Project tokens: [batch, seq, hidden] -> [batch, seq, 2*half] let proj_tokens = self.proj_token.forward(token_embs); let [_, _, proj_dim] = proj_tokens.shape(); - let half_proj = proj_dim / 2; + let half = proj_dim / 2; - // label: [n_labels, hidden] -> [n_labels, hidden*2] + // Project labels: [n_labels, hidden] -> [n_labels, 2*half] let label_embs_3d: Tensor<3, f32> = label_embs.unsqueeze(0).to_concrete(); let proj_labels = self.proj_label.forward(&label_embs_3d); let proj_labels: Tensor<2, f32> = proj_labels.squeeze(0).to_concrete(); - // Split and combine: token_first + label_first + (token_second * label_second) - // MLP input dimension = half_proj + half_proj + half_proj = 3 * half_proj - let mlp_input_dim = 3 * half_proj; - - // Get raw data slices (without expansion - we'll handle broadcast manually) - // proj_tokens shape: [batch, seq, proj_dim] - // proj_labels shape: [n_labels, proj_dim] - let tokens_data = proj_tokens.clone().as_slice().await.unwrap(); - let labels_data = proj_labels.clone().as_slice().await.unwrap(); - - let tokens_slice = tokens_data.as_slice(); // [batch * seq * proj_dim] - let labels_slice = labels_data.as_slice(); // [n_labels * proj_dim] - - // Build combined features with manual broadcasting - // Output: [batch, seq, n_labels, mlp_input_dim] - let total_elements = batch_size * seq_len * n_labels; - let mut combined_data = vec![0.0f32; total_elements * mlp_input_dim]; - - for b in 0..batch_size { - for s in 0..seq_len { - for l in 0..n_labels { - // Token features for (b, s): at index (b * seq_len + s) * proj_dim - let tok_base = (b * seq_len + s) * proj_dim; - // Label features for l: at index l * proj_dim - let lab_base = l * proj_dim; - // Output index for (b, s, l) - let out_idx = (b * seq_len * n_labels + s * n_labels + l) * mlp_input_dim; - - // token_first (first half of token projection) - for i in 0..half_proj { - combined_data[out_idx + i] = tokens_slice[tok_base + i]; - } - // label_first (first half of label projection) - for i in 0..half_proj { - combined_data[out_idx + half_proj + i] = labels_slice[lab_base + i]; - } - // element-wise product of second halves - for i in 0..half_proj { - let tok_second = tokens_slice[tok_base + half_proj + i]; - let lab_second = labels_slice[lab_base + half_proj + i]; - combined_data[out_idx + 2 * half_proj + i] = tok_second * lab_second; - } - } - } - } - - let device = token_embs.device(); - let combined: Tensor<3, f32> = Tensor::new(&device, &combined_data) - .reshape([1, total_elements, mlp_input_dim]) + // Split projections into first/second halves along the feature dim. + let tokens_first = proj_tokens.narrow(2, 0, half).to_concrete(); // [b, s, half] + let tokens_second = proj_tokens.narrow(2, half, half).to_concrete(); // [b, s, half] + let labels_first = proj_labels.narrow(1, 0, half).to_concrete(); // [n, half] + let labels_second = proj_labels.narrow(1, half, half).to_concrete(); // [n, half] + + // Broadcast to [batch, seq, n_labels, half] and build the three concatenation parts: + // [token_first, label_first, token_second * label_second] + let target = [batch_size, seq_len, n_labels, half]; + let tok_first_4d: Tensor<4, f32> = + tokens_first.unsqueeze(2).broadcast_as(target).to_concrete(); + let tok_second_4d: Tensor<4, f32> = tokens_second + .unsqueeze(2) + .broadcast_as(target) + .to_concrete(); + let lab_first_4d: Tensor<4, f32> = labels_first + .unsqueeze(0) + .unsqueeze(0) + .broadcast_as(target) + .to_concrete(); + let lab_second_4d: Tensor<4, f32> = labels_second + .unsqueeze(0) + .unsqueeze(0) + .broadcast_as(target) .to_concrete(); - // Apply MLP: fc1 -> ReLU -> fc2 - let hidden = self.out_fc1.forward(&combined); - let hidden = hidden.relu(); - let output = self.out_fc2.forward(&hidden); + let prod_4d: Tensor<4, f32> = (tok_second_4d * lab_second_4d).to_concrete(); + + // Concat along the last dim -> [batch, seq, n_labels, 3*half] + let combined: Tensor<4, f32> = Tensor::cat([tok_first_4d, lab_first_4d, prod_4d], 3); - // Reshape back: [batch, seq, n_labels, 3] + // Linear::forward takes 3D input, so fold (batch, seq, n_labels) into one dim. + let mlp_in_dim = 3 * half; + let flat: Tensor<3, f32> = combined + .reshape([1, batch_size * seq_len * n_labels, mlp_in_dim]) + .to_concrete(); + + let hidden = self.out_fc1.forward(&flat).relu(); + let output = self.out_fc2.forward(&hidden); output .reshape([batch_size, seq_len, n_labels, 3]) .to_concrete() @@ -142,22 +123,17 @@ impl JointScorer { /// /// The 3 channels are: [start, end, inside] (NOT OBI). /// Each channel is passed through independent sigmoid. - pub async fn forward_entity_scores( + pub fn forward_entity_scores( &self, token_embs: &Tensor<3, f32>, label_embs: &Tensor<2, f32>, ) -> Tensor<4, f32> { - let logits = self.forward(token_embs, label_embs).await; - let logits_data = logits.clone().as_slice().await.unwrap(); - - // Apply sigmoid to each value independently (NOT softmax). - let data = logits_data.as_slice(); - let sigmoid_data: Vec = data.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect(); - - let device = logits.device(); - Tensor::new(&device, &sigmoid_data) - .reshape(logits.shape()) - .to_concrete() + let logits = self.forward(token_embs, label_embs); + // sigmoid(x) = 0.5 * (tanh(x / 2) + 1); stays on-device and avoids needing + // scalar-left division or a `recip` primitive. + let half_logits: Tensor<4, f32> = (logits * 0.5f32).to_concrete(); + let tanh = half_logits.tanh(); + ((tanh + 1.0f32) * 0.5f32).to_concrete() } } @@ -178,22 +154,6 @@ impl PromptRepLayer { Ok(Self { fc1, fc2 }) } - /// Project label embeddings. - /// - /// # Arguments - /// * `label_embs` - Label embeddings from encoder [n_labels, hidden] - /// - /// # Returns - /// Projected embeddings [n_labels, hidden] - pub fn forward(&self, label_embs: &Tensor<2, f32>) -> Tensor<2, f32> { - // Wrap as 3D for Linear::forward - let label_3d: Tensor<3, f32> = label_embs.unsqueeze(0).to_concrete(); - let hidden = self.fc1.forward(&label_3d); - let hidden = hidden.relu(); - let output = self.fc2.forward(&hidden); - output.squeeze(0).to_concrete() - } - /// Forward for 3D tensor [batch, n_labels, hidden]. pub fn forward_3d(&self, label_embs: &Tensor<3, f32>) -> Tensor<3, f32> { let hidden = self.fc1.forward(label_embs); diff --git a/models/rgliner/src/raw/label_encoder.rs b/models/rgliner/src/raw/label_encoder.rs index 3f96eaabe..dd7515b06 100644 --- a/models/rgliner/src/raw/label_encoder.rs +++ b/models/rgliner/src/raw/label_encoder.rs @@ -141,11 +141,6 @@ impl LabelEncoder { }) } - /// Get the output dimension. - pub fn output_dim(&self) -> usize { - self.output_dim - } - #[cfg(test)] pub async fn debug_sentence_embeddings( &self, @@ -290,19 +285,4 @@ impl CachedLabels { pub fn new(labels: Vec, embeddings: Tensor<2, f32>) -> Self { Self { labels, embeddings } } - - /// Get the number of labels. - pub fn len(&self) -> usize { - self.labels.len() - } - - /// Check if empty. - pub fn is_empty(&self) -> bool { - self.labels.is_empty() - } - - /// Get label at index. - pub fn get_label(&self, idx: usize) -> Option<&str> { - self.labels.get(idx).map(|s| s.as_str()) - } } diff --git a/models/rgliner/src/raw/mod.rs b/models/rgliner/src/raw/mod.rs index a1812477d..b8557b656 100644 --- a/models/rgliner/src/raw/mod.rs +++ b/models/rgliner/src/raw/mod.rs @@ -4,7 +4,6 @@ mod bilstm; mod joint_scorer; mod label_encoder; mod pair_projector; -mod relations_layer; mod scorer; mod span_layer; mod text_encoder; @@ -13,7 +12,6 @@ pub use bilstm::BiLstm; pub use joint_scorer::{JointScorer, PromptRepLayer}; pub use label_encoder::{CachedLabels, LabelEncoder}; pub use pair_projector::PairProjector; -pub use relations_layer::RelationsRepLayer; pub use scorer::Scorer; pub use span_layer::SpanLayer; pub use text_encoder::TextEncoder; diff --git a/models/rgliner/src/raw/pair_projector.rs b/models/rgliner/src/raw/pair_projector.rs index 91f5cb18f..536d30d8a 100644 --- a/models/rgliner/src/raw/pair_projector.rs +++ b/models/rgliner/src/raw/pair_projector.rs @@ -58,86 +58,4 @@ impl PairProjector { let result = self.linear2.forward(&hidden); result.squeeze(0).to_concrete() } - - /// Project entity pairs for batched processing. - /// - /// # Arguments - /// * `head_embeddings` - Head entity embeddings [batch, num_pairs, hidden_size] - /// * `tail_embeddings` - Tail entity embeddings [batch, num_pairs, hidden_size] - /// - /// # Returns - /// Pair representations [batch, num_pairs, hidden_size] - pub fn forward_batched( - &self, - head_embeddings: &Tensor<3, f32>, - tail_embeddings: &Tensor<3, f32>, - ) -> Tensor<3, f32> { - let [_batch_size, _num_pairs, _hidden_size] = head_embeddings.shape(); - - // Concatenate head and tail: [batch, num_pairs, hidden_size * 2] - let concatenated = Tensor::cat( - [head_embeddings.to_concrete(), tail_embeddings.to_concrete()], - 2, - ); - - // First layer: Linear -> ReLU - let hidden = self.linear1.forward(&concatenated).relu(); - - // Second layer: Linear - self.linear2.forward(&hidden) - } -} - -/// Scorer for relation classification. -/// -/// Computes scores between pair representations and relation label embeddings. -pub struct RelationScorer; - -impl RelationScorer { - /// Score pairs against relation labels. - /// - /// # Arguments - /// * `pair_embeddings` - Pair representations [num_pairs, hidden_size] - /// * `relation_embeddings` - Relation label embeddings [num_relations, hidden_size] - /// - /// # Returns - /// Scores [num_pairs, num_relations] (logits, apply sigmoid for probabilities) - pub fn forward( - pair_embeddings: &Tensor<2, f32>, - relation_embeddings: &Tensor<2, f32>, - ) -> Tensor<2, f32> { - // Dot product: pairs @ relations.T - let rel_t = relation_embeddings.transpose(0, 1); - pair_embeddings.mat_mul(&rel_t) - } - - /// Score pairs against relation labels (batched). - /// - /// # Arguments - /// * `pair_embeddings` - Pair representations [batch, num_pairs, hidden_size] - /// * `relation_embeddings` - Relation label embeddings [num_relations, hidden_size] - /// - /// # Returns - /// Scores [batch, num_pairs, num_relations] - pub fn forward_batched( - pair_embeddings: &Tensor<3, f32>, - relation_embeddings: &Tensor<2, f32>, - ) -> Tensor<3, f32> { - let [batch_size, num_pairs, hidden_size] = pair_embeddings.shape(); - let [num_relations, _] = relation_embeddings.shape(); - - // Flatten pairs: [batch * num_pairs, hidden_size] - let flat_pairs = pair_embeddings - .reshape([batch_size * num_pairs, hidden_size]) - .to_concrete(); - - // Dot product: [batch * num_pairs, hidden_size] @ [hidden_size, num_relations] - let rel_t = relation_embeddings.transpose(0, 1); - let scores = flat_pairs.mat_mul(&rel_t); - - // Reshape back: [batch, num_pairs, num_relations] - scores - .reshape([batch_size, num_pairs, num_relations]) - .to_concrete() - } } diff --git a/models/rgliner/src/raw/relations_layer.rs b/models/rgliner/src/raw/relations_layer.rs deleted file mode 100644 index 3589e9fe6..000000000 --- a/models/rgliner/src/raw/relations_layer.rs +++ /dev/null @@ -1,120 +0,0 @@ -//! Relation representation layer for adjacency matrix computation. -//! -//! Computes an adjacency matrix between entity spans to filter -//! candidate pairs for relation classification. - -use fusor::layers::Linear; -use fusor::{Device, Result, Tensor, VarBuilder}; - -/// Relation representation layer - can be learned (bilinear) or simple dot-product. -pub enum RelationsRepLayer { - /// Learned bilinear projection - Bilinear(BilinearRelationsLayer), - /// Simple dot-product similarity (no learned weights) - DotProduct, -} - -impl RelationsRepLayer { - /// Load the relations layer from GGUF weights. - /// Falls back to dot-product if weights don't exist. - pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { - match BilinearRelationsLayer::load(device, vb) { - Ok(bilinear) => Ok(Self::Bilinear(bilinear)), - Err(_) => Ok(Self::DotProduct), - } - } - - /// Create a dot-product based relations layer (no learned weights). - pub fn identity(_device: &Device, _hidden_size: usize) -> Self { - Self::DotProduct - } - - /// Compute adjacency matrix for entity spans. - /// - /// # Arguments - /// * `entity_embeddings` - Entity span embeddings [batch, num_entities, hidden_size] - /// - /// # Returns - /// Adjacency logits [batch, num_entities, num_entities] (apply sigmoid externally) - pub fn forward(&self, entity_embeddings: &Tensor<3, f32>) -> Tensor<3, f32> { - match self { - Self::Bilinear(layer) => layer.forward(entity_embeddings), - Self::DotProduct => { - // Simple dot product: embeddings @ embeddings.T - let entity_t = entity_embeddings.transpose(1, 2); - entity_embeddings.mat_mul(&entity_t) - } - } - } - - /// Apply sigmoid to logits (for use after forward). - pub fn apply_sigmoid(logits: &[f32]) -> Vec { - logits.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect() - } - - /// Filter entity pairs based on adjacency threshold. - /// - /// # Arguments - /// * `adjacency_scores` - Adjacency matrix [batch, num_entities, num_entities] - /// * `threshold` - Minimum score for a pair to be considered - /// - /// # Returns - /// Vector of (batch_idx, head_idx, tail_idx, score) tuples for pairs above threshold - pub async fn filter_pairs( - &self, - adjacency_scores: &Tensor<3, f32>, - threshold: f32, - ) -> Result> { - let [batch_size, num_entities, _] = adjacency_scores.shape(); - - let scores_slice = adjacency_scores.clone().as_slice().await?; - let scores_data = scores_slice.as_slice(); - - let mut pairs = Vec::new(); - for b in 0..batch_size { - for i in 0..num_entities { - for j in 0..num_entities { - if i == j { - continue; // Skip self-relations - } - let idx = b * num_entities * num_entities + i * num_entities + j; - let score = scores_data[idx]; - if score >= threshold { - pairs.push((b, i, j, score)); - } - } - } - } - - Ok(pairs) - } -} - -/// Learned bilinear relation representation layer. -/// -/// Computes adjacency scores between entity pairs: -/// `adj_score[i,j] = sigmoid(entity_i @ W @ entity_j.T)` -pub struct BilinearRelationsLayer { - /// Bilinear projection weight [hidden_size, hidden_size] - projection: Linear, -} - -impl BilinearRelationsLayer { - /// Load from GGUF weights. - pub fn load(device: &Device, vb: &mut VarBuilder) -> Result { - let projection = Linear::load(device, &mut vb.pp("projection"))?; - Ok(Self { projection }) - } - - /// Compute adjacency matrix for entity spans. - pub fn forward(&self, entity_embeddings: &Tensor<3, f32>) -> Tensor<3, f32> { - // Project entity embeddings: [batch, num_entities, hidden_size] - let projected = self.projection.forward(entity_embeddings); - - // Compute bilinear scores: projected @ entity_embeddings.T - // [batch, num_entities, hidden_size] @ [batch, hidden_size, num_entities] - // = [batch, num_entities, num_entities] - let entity_t = entity_embeddings.transpose(1, 2); - projected.mat_mul(&entity_t) - } -} diff --git a/models/rgliner/src/raw/scorer.rs b/models/rgliner/src/raw/scorer.rs index 966649c9c..6f6259580 100644 --- a/models/rgliner/src/raw/scorer.rs +++ b/models/rgliner/src/raw/scorer.rs @@ -43,20 +43,3 @@ impl Scorer { .to_concrete() } } - -/// Apply sigmoid to raw scores in Rust (not on GPU). -/// -/// # Arguments -/// * `logits` - Raw logit scores -/// -/// # Returns -/// Probability scores (0.0 to 1.0) -#[inline] -pub fn sigmoid(x: f32) -> f32 { - 1.0 / (1.0 + (-x).exp()) -} - -/// Apply sigmoid to a slice of logits. -pub fn apply_sigmoid(logits: &[f32]) -> Vec { - logits.iter().map(|&x| sigmoid(x)).collect() -} diff --git a/models/rgliner/src/raw/span_layer.rs b/models/rgliner/src/raw/span_layer.rs index 033493ca3..31f9aee8f 100644 --- a/models/rgliner/src/raw/span_layer.rs +++ b/models/rgliner/src/raw/span_layer.rs @@ -24,8 +24,6 @@ pub struct SpanLayer { out_fc2: Linear, /// Maximum span width max_width: usize, - /// Hidden dimension - hidden_dim: usize, } impl SpanLayer { @@ -71,8 +69,6 @@ impl SpanLayer { ) })?; - let hidden_dim = out_fc2.out_features(); - Ok(Self { start_fc1, start_fc2, @@ -81,15 +77,9 @@ impl SpanLayer { out_fc1, out_fc2, max_width, - hidden_dim, }) } - /// Get the maximum span width. - pub fn max_width(&self) -> usize { - self.max_width - } - /// Enumerate all valid spans up to max_width. /// /// Returns Vec of (start_word, end_word) pairs. diff --git a/models/rgliner/src/raw/text_encoder.rs b/models/rgliner/src/raw/text_encoder.rs index 136ca66f9..4f55e0939 100644 --- a/models/rgliner/src/raw/text_encoder.rs +++ b/models/rgliner/src/raw/text_encoder.rs @@ -44,26 +44,6 @@ impl TextEncoder { } } - /// Get the maximum sequence length. - pub fn max_seq_len(&self) -> usize { - self.model.max_seq_len() - } - - /// Get the embedding dimension seen by downstream layers (post-projection - /// if this variant has one, otherwise the raw encoder hidden dim). - pub fn embedding_dim(&self) -> usize { - if let Some(ref proj) = self.output_proj { - proj.out_features() - } else { - self.model.embedding_dim() - } - } - - /// Get the device. - pub fn device(&self) -> &Device { - self.model.device() - } - #[cfg(test)] pub fn debug_hidden_states( &self, diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index 1adff5c94..3ef546da7 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -50,10 +50,8 @@ use tokenizers::Tokenizer; use crate::decoding::Entity; use crate::error::{GlinerError, GlinerLoadingError}; -use crate::raw::{ - BiLstm, JointScorer, PairProjector, PromptRepLayer, RelationsRepLayer, SpanLayer, -}; -use crate::relation_decoding::{Relation, RelationDecoder, RelationDecoderConfig}; +use crate::raw::{BiLstm, JointScorer, PairProjector, PromptRepLayer, SpanLayer}; +use crate::relation_decoding::Relation; use crate::relex_tokenization::{RelExTokenizer, SpecialTokenIds}; use rbert::raw::MDebertaModel; @@ -281,14 +279,10 @@ pub struct GlinerRelEx { scorer: Option, /// Span representation layer span_layer: SpanLayer, - /// Relations representation layer (adjacency scoring) - relations_layer: RelationsRepLayer, /// Entity pair projector pair_projector: PairProjector, /// Tokenizer with special token handling tokenizer: Arc, - /// Relation decoder - relation_decoder: RelationDecoder, /// How entities are scored (derived from `gliner.span_mode` metadata). span_mode: SpanMode, /// Device @@ -416,30 +410,17 @@ impl GlinerRelEx { // Load span layer let span_layer = SpanLayer::load(&device, &mut vb, config.max_width)?; - // Load relations layer (may not exist, use projection from pair_proj) - let relations_layer = RelationsRepLayer::load(&device, &mut vb.pp("relations")) - .unwrap_or_else(|_| RelationsRepLayer::identity(&device, config.hidden_size)); - // Load pair projector let pair_projector = PairProjector::load(&device, &mut vb.pp("pair_proj"))?; - // Create relation decoder - let relation_decoder = RelationDecoder::with_config(RelationDecoderConfig { - entity_threshold: config.entity_threshold, - adjacency_threshold: config.adjacency_threshold, - relation_threshold: config.relation_threshold, - }); - Ok(Self { encoder, bilstm, prompt_rep_layer, scorer, span_layer, - relations_layer, pair_projector, tokenizer: Arc::new(relex_tokenizer), - relation_decoder, span_mode, device, config, @@ -504,7 +485,7 @@ impl GlinerRelEx { let entities = match self.span_mode { SpanMode::TokenLevel => { let scorer = self.scorer.as_ref().expect("token_level requires scorer"); - let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs_2d).await; + let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs_2d); self.decode_entities_from_tokens( &token_scores, entity_labels, @@ -828,21 +809,6 @@ impl GlinerRelEx { gathered.unsqueeze(0).to_concrete() } - /// Build a tensor from entity embeddings. - fn build_entity_tensor(&self, embeddings: &[Vec], device: &Device) -> Tensor<3, f32> { - let num_entities = embeddings.len(); - if num_entities == 0 || embeddings[0].is_empty() { - return Tensor::zeros(device, [1, 1, self.config.hidden_size]); - } - - let hidden_size = embeddings[0].len(); - let flat: Vec = embeddings.iter().flatten().copied().collect(); - - Tensor::new(device, &flat) - .reshape([1, num_entities, hidden_size]) - .to_concrete() - } - /// Get the device. pub fn device(&self) -> &Device { &self.device diff --git a/models/rgliner/src/tokenization.rs b/models/rgliner/src/tokenization.rs index 1d11ca38a..ad702cb4a 100644 --- a/models/rgliner/src/tokenization.rs +++ b/models/rgliner/src/tokenization.rs @@ -12,8 +12,6 @@ pub struct TokenizedText { pub token_ids: Vec, /// Attention mask (1 for real tokens, 0 for padding). pub attention_mask: Vec, - /// Maps each token position to its word index (-1 for special tokens). - pub token_to_word: Vec, /// Index of the first token for each word. pub word_first_token: Vec, /// Number of words in the input. @@ -33,12 +31,6 @@ impl WordTokenizer { Self { tokenizer } } - /// Load tokenizer from JSON bytes. - pub fn from_bytes(bytes: &[u8]) -> Result { - let tokenizer = Tokenizer::from_bytes(bytes)?; - Ok(Self::new(tokenizer)) - } - /// Tokenize text and track word boundaries. pub fn tokenize(&self, text: &str) -> Result { let split_words = split_words(text); @@ -58,21 +50,12 @@ impl WordTokenizer { .map(|&x| x as u32) .collect(); - // Build token-to-word mapping - // word_ids() returns Option for each token - let token_to_word: Vec = encoding - .get_word_ids() - .iter() - .map(|opt| opt.map(|w| w as i32).unwrap_or(-1)) - .collect(); - + // Build token-to-word mapping to find the first token for each word. let num_words = word_offsets.len(); - - // Find first token index for each word let mut word_first_token = vec![0usize; num_words]; let mut seen_words = vec![false; num_words]; - for (token_idx, &word_id) in token_to_word.iter().enumerate() { - if word_id >= 0 { + for (token_idx, opt) in encoding.get_word_ids().iter().enumerate() { + if let Some(word_id) = *opt { let word_id = word_id as usize; if !seen_words[word_id] { word_first_token[word_id] = token_idx; @@ -84,17 +67,11 @@ impl WordTokenizer { Ok(TokenizedText { token_ids, attention_mask, - token_to_word, word_first_token, num_words, word_offsets, }) } - - /// Tokenize a batch of texts. - pub fn tokenize_batch(&self, texts: &[&str]) -> Result, GlinerError> { - texts.iter().map(|text| self.tokenize(text)).collect() - } } fn split_words(text: &str) -> Vec<(String, (usize, usize))> { From 25614b869c5f1d70ca199c33ce02aa06b653c473 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Tue, 14 Apr 2026 07:24:31 -0500 Subject: [PATCH 12/34] fix rgliner --- models/rgliner/examples/debug_v2.rs | 50 +++++++++++++++++++++++++++++ models/rgliner/src/config.rs | 18 +++++++++++ models/rgliner/src/lib.rs | 2 +- models/rgliner/src/tokenization.rs | 18 ++++++++--- 4 files changed, 83 insertions(+), 5 deletions(-) create mode 100644 models/rgliner/examples/debug_v2.rs diff --git a/models/rgliner/examples/debug_v2.rs b/models/rgliner/examples/debug_v2.rs new file mode 100644 index 000000000..8916304ae --- /dev/null +++ b/models/rgliner/examples/debug_v2.rs @@ -0,0 +1,50 @@ +use kalosm_model_types::FileSource; +use rgliner::*; +use std::path::PathBuf; + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + let source = GlinerSource::custom( + FileSource::local(PathBuf::from("./models/rgliner/weights/test-small-f32.gguf")), + FileSource::local(PathBuf::from( + "./models/rgliner/weights/test-small-f32-label-encoder.gguf", + )), + FileSource::huggingface( + "sentence-transformers/all-MiniLM-L12-v2".to_string(), + "main".to_string(), + "config.json".to_string(), + ), + FileSource::huggingface( + "sentence-transformers/all-MiniLM-L12-v2".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + FileSource::huggingface( + "knowledgator/gliner-bi-small-v2.0".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + FileSource::huggingface( + "knowledgator/gliner-bi-small-v2.0".to_string(), + "main".to_string(), + "gliner_config.json".to_string(), + ), + ); + let mut gliner = Gliner::builder() + .with_source(source) + .with_threshold(0.0) + .build() + .await?; + + let text = "Cristiano Ronaldo dos Santos Aveiro (Portuguese pronunciation: [kɾiʃˈtjɐnu ʁɔˈnaldu]; born 5 February 1985) is a Portuguese professional footballer who plays as a forward for and captains both Saudi Pro League club Al Nassr and the Portugal national team."; + let labels = ["person", "award", "date", "competitions", "teams"]; + + let entities = gliner.extract(text, &labels).await?; + let mut ents = entities; + ents.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap()); + ents.iter() + .take(10) + .for_each(|e| println!(" {} '{}' {:.4}", e.label, e.text, e.score)); + + Ok(()) +} diff --git a/models/rgliner/src/config.rs b/models/rgliner/src/config.rs index 1f17a9ff6..44ac3832d 100644 --- a/models/rgliner/src/config.rs +++ b/models/rgliner/src/config.rs @@ -122,6 +122,24 @@ impl GlinerConfig { pub fn uses_first_subtoken(&self) -> bool { self.subtoken_pooling == "first" } + + /// Whether the tokenizer should add [CLS]/[SEP] special tokens around text. + /// + /// Matches Python GLiNER's `_set_tokenizer_spec_tokens` behavior: + /// ModernBERT/ettin-style encoders (which have `add_bos_token=False` + /// semantics because they have no bos token) are fed raw text without + /// [CLS]/[SEP] wrappers. DeBERTa/RoBERTa/XLM-R family keep them. + pub fn should_add_special_tokens(&self) -> bool { + match self.model_name.as_deref() { + Some(name) => { + let lower = name.to_ascii_lowercase(); + !(lower.contains("ettin") + || lower.contains("modernbert") + || lower.contains("modern-bert")) + } + None => true, + } + } } impl Default for GlinerConfig { diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs index 083c28980..4f3974ea3 100644 --- a/models/rgliner/src/lib.rs +++ b/models/rgliner/src/lib.rs @@ -243,7 +243,7 @@ impl Gliner { let tokenizer = Tokenizer::from_bytes(&tokenizer_bytes).map_err(GlinerLoadingError::LoadTokenizer)?; - let word_tokenizer = WordTokenizer::new(tokenizer); + let word_tokenizer = WordTokenizer::new(tokenizer, config.should_add_special_tokens()); // Download main model weights let model_source = format!("Text Encoder ({})", source.model); diff --git a/models/rgliner/src/tokenization.rs b/models/rgliner/src/tokenization.rs index ad702cb4a..10f5ec3ae 100644 --- a/models/rgliner/src/tokenization.rs +++ b/models/rgliner/src/tokenization.rs @@ -23,12 +23,22 @@ pub struct TokenizedText { /// Word-level tokenizer wrapper. pub struct WordTokenizer { tokenizer: Tokenizer, + add_special_tokens: bool, } impl WordTokenizer { - /// Create a new word tokenizer from a HuggingFace tokenizer. - pub fn new(tokenizer: Tokenizer) -> Self { - Self { tokenizer } + /// Create a tokenizer. + /// + /// `add_special_tokens` controls whether the tokenizer's post-processor is + /// applied. Set to `false` for encoders whose Python counterpart strips + /// [CLS]/[SEP] from the post-processor (ModernBERT/ettin have + /// `add_bos_token=False` because they lack bos/eos tokens — see GLiNER's + /// `_set_tokenizer_spec_tokens`). + pub fn new(tokenizer: Tokenizer, add_special_tokens: bool) -> Self { + Self { + tokenizer, + add_special_tokens, + } } /// Tokenize text and track word boundaries. @@ -40,7 +50,7 @@ impl WordTokenizer { let encoding = self .tokenizer - .encode(words, true) + .encode(words, self.add_special_tokens) .map_err(GlinerError::Tokenizer)?; let token_ids = encoding.get_ids().to_vec(); From be018f5b099befae5e6425e881ea96807de6d583 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Tue, 14 Apr 2026 16:29:37 -0500 Subject: [PATCH 13/34] ui demo --- .cargo/config.toml | 3 + .claude/settings.local.json | 3 +- Cargo.lock | 1114 +++++++++++++++++++++++++++- Cargo.toml | 1 + demos/rgliner-web/Cargo.toml | 16 + demos/rgliner-web/Dioxus.toml | 5 + demos/rgliner-web/assets/style.css | 234 ++++++ demos/rgliner-web/src/main.rs | 567 ++++++++++++++ fusor-ml/core/src/tensor.rs | 2 +- models/rgliner/src/lib.rs | 18 +- models/rgliner/src/relex.rs | 24 +- models/rgliner/src/source.rs | 208 +----- 12 files changed, 1968 insertions(+), 227 deletions(-) create mode 100644 demos/rgliner-web/Cargo.toml create mode 100644 demos/rgliner-web/Dioxus.toml create mode 100644 demos/rgliner-web/assets/style.css create mode 100644 demos/rgliner-web/src/main.rs diff --git a/.cargo/config.toml b/.cargo/config.toml index d2b00e6c7..6ff324e91 100644 --- a/.cargo/config.toml +++ b/.cargo/config.toml @@ -4,6 +4,9 @@ [target.x86_64-apple-darwin] rustflags = ["-C", "target-feature=-avx,-avx2"] +[target.wasm32-unknown-unknown] +rustflags = ["--cfg", "getrandom_backend=\"wasm_js\""] + # [unstable] # build-std = ["std", "core", "alloc"] # build-std-features = ["panic_immediate_abort"] diff --git a/.claude/settings.local.json b/.claude/settings.local.json index 40fcd0c01..269f0fdf3 100644 --- a/.claude/settings.local.json +++ b/.claude/settings.local.json @@ -37,7 +37,8 @@ "Bash(cargo tree:*)", "Bash(r\" | head -20)", "Read(//tmp/**)", - "Bash(cargo fmt:*)" + "Bash(cargo fmt:*)", + "Bash(dx:*)" ] } } diff --git a/Cargo.lock b/Cargo.lock index 72d84b8e1..12b60aafa 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -501,6 +501,28 @@ dependencies = [ "pin-project-lite", ] +[[package]] +name = "async-stream" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b5a71a6f37880a80d1d7f19efd781e4b5de42c88f0722cc13bcb6cc2cfe8476" +dependencies = [ + "async-stream-impl", + "futures-core", + "pin-project-lite", +] + +[[package]] +name = "async-stream-impl" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "async-task" version = "4.7.1" @@ -518,6 +540,22 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "async-tungstenite" +version = "0.31.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee88b4c88ac8c9ea446ad43498955750a4bbe64c4392f21ccfe5d952865e318f" +dependencies = [ + "atomic-waker", + "futures-core", + "futures-io", + "futures-task", + "futures-util", + "log", + "pin-project-lite", + "tungstenite 0.27.0", +] + [[package]] name = "async_io_stream" version = "0.3.3" @@ -648,7 +686,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "edca88bc138befd0323b20752846e6587272d3b03b0343c8ea28a6f819e6e71f" dependencies = [ "async-trait", - "axum-core", + "axum-core 0.4.5", "bytes", "futures-util", "http 1.4.0", @@ -657,7 +695,7 @@ dependencies = [ "hyper 1.8.1", "hyper-util", "itoa", - "matchit", + "matchit 0.7.3", "memchr", "mime", "percent-encoding", @@ -675,6 +713,36 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90" +dependencies = [ + "axum-core 0.5.6", + "bytes", + "form_urlencoded", + "futures-util", + "http 1.4.0", + "http-body 1.0.1", + "http-body-util", + "itoa", + "matchit 0.8.4", + "memchr", + "mime", + "multer", + "percent-encoding", + "pin-project-lite", + "serde_core", + "serde_json", + "serde_path_to_error", + "serde_urlencoded", + "sync_wrapper 1.0.2", + "tower", + "tower-layer", + "tower-service", +] + [[package]] name = "axum-core" version = "0.4.5" @@ -696,6 +764,24 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum-core" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1" +dependencies = [ + "bytes", + "futures-core", + "http 1.4.0", + "http-body 1.0.1", + "http-body-util", + "mime", + "pin-project-lite", + "sync_wrapper 1.0.2", + "tower-layer", + "tower-service", +] + [[package]] name = "base16ct" version = "0.2.0" @@ -1291,6 +1377,17 @@ dependencies = [ "nom 7.1.3", ] +[[package]] +name = "cfb" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d38f2da7a0a2c4ccf0065be06397cc26a81f4e528be095826eee9d4adbb8c60f" +dependencies = [ + "byteorder", + "fnv", + "uuid", +] + [[package]] name = "cfg-if" version = "1.0.4" @@ -1303,6 +1400,16 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +[[package]] +name = "charset" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1f927b07c74ba84c7e5fe4db2baeb3e996ab2688992e39ac68ce3220a677c7e" +dependencies = [ + "base64 0.22.1", + "encoding_rs", +] + [[package]] name = "chrono" version = "0.4.44" @@ -1537,12 +1644,99 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "console_error_panic_hook" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a06aeb73f470f66dcdbf7223caeebb85984942f22f1adb2a088cf9668146bbbc" +dependencies = [ + "cfg-if", + "wasm-bindgen", +] + +[[package]] +name = "const-serialize" +version = "0.7.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ad7154afa56de2f290e3c82c2c6dc4f5b282b6870903f56ef3509aba95866edc" +dependencies = [ + "const-serialize-macro 0.7.2", +] + +[[package]] +name = "const-serialize" +version = "0.8.0-alpha.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e42cd5aabba86f128b3763da1fec1491c0f728ce99245062cd49b6f9e6d235b" +dependencies = [ + "const-serialize 0.7.2", + "const-serialize-macro 0.8.0-alpha.0", + "serde", +] + +[[package]] +name = "const-serialize-macro" +version = "0.7.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4f160aad86b4343e8d4e261fee9965c3005b2fd6bc117d172ab65948779e4acf" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "const-serialize-macro" +version = "0.8.0-alpha.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42571ed01eb46d2e1adcf99c8ca576f081e46f2623d13500eba70d1d99a4c439" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "const-str" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b0664d2867b4a32697dfe655557f5c3b187e9b605b38612a748e5ec99811d160" + +[[package]] +name = "const_format" +version = "0.2.35" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7faa7469a93a566e9ccc1c73fe783b4a65c274c5ace346038dca9c39fe0030ad" +dependencies = [ + "const_format_proc_macros", +] + +[[package]] +name = "const_format_proc_macros" +version = "0.2.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d57c2eccfb16dbac1f4e61e206105db5820c9d26c3c472bc17c774259ef7744" +dependencies = [ + "proc-macro2", + "quote", + "unicode-xid", +] + [[package]] name = "constant_time_eq" version = "0.4.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3d52eff69cd5e647efe296129160853a42795992097e8af39800e1060caeea9b" +[[package]] +name = "content_disposition" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc14a88e1463ddd193906285abe5c360c7e8564e05ccc5d501755f7fbc9ca9c" +dependencies = [ + "charset", +] + [[package]] name = "convert_case" version = "0.6.0" @@ -1561,6 +1755,44 @@ dependencies = [ "unicode-segmentation", ] +[[package]] +name = "convert_case" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "633458d4ef8c78b72454de2d54fd6ab2e60f9e02be22f3c6104cdc8a4e0fceb9" +dependencies = [ + "unicode-segmentation", +] + +[[package]] +name = "cookie" +version = "0.18.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ddef33a339a91ea89fb53151bd0a4689cfce27055c291dfa69945475d22c747" +dependencies = [ + "percent-encoding", + "time", + "version_check", +] + +[[package]] +name = "cookie_store" +version = "0.22.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "15b2c103cf610ec6cae3da84a766285b42fd16aad564758459e6ecf128c75206" +dependencies = [ + "cookie", + "document-features", + "idna", + "log", + "publicsuffix", + "serde", + "serde_derive", + "serde_json", + "time", + "url", +] + [[package]] name = "core-foundation" version = "0.9.4" @@ -1641,7 +1873,7 @@ dependencies = [ "js-sys", "libc", "mach2", - "ndk", + "ndk 0.8.0", "ndk-context", "oboe", "wasm-bindgen", @@ -2299,6 +2531,29 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "derive_more" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d751e9e49156b02b44f9c1815bcb94b984cdcc4396ecc32521c739452808b134" +dependencies = [ + "derive_more-impl", +] + +[[package]] +name = "derive_more-impl" +version = "2.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "799a97264921d8623a957f6c3b9011f3b5492f557bbb7a5a19b7fa6d06ba8dcb" +dependencies = [ + "convert_case 0.10.0", + "proc-macro2", + "quote", + "rustc_version", + "syn 2.0.117", + "unicode-xid", +] + [[package]] name = "deunicode" version = "1.6.2" @@ -2331,6 +2586,443 @@ dependencies = [ "chrono", ] +[[package]] +name = "dioxus" +version = "0.7.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a115f9dbe5900c6044ee6a791e1b160c29989c6a8721eec099e01a964e5dae4" +dependencies = [ + "dioxus-asset-resolver", + "dioxus-cli-config", + "dioxus-config-macro", + "dioxus-config-macros", + "dioxus-core", + "dioxus-core-macro", + "dioxus-devtools", + "dioxus-document", + "dioxus-fullstack", + "dioxus-history", + "dioxus-hooks", + "dioxus-html", + "dioxus-logger", + "dioxus-signals", + "dioxus-stores", + "dioxus-web", + "manganis", + "subsecond", + "warnings", +] + +[[package]] +name = "dioxus-asset-resolver" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c240c4f092024b26e200ecd64723009173cf5bc2e5083c9feb778c077eb5741b" +dependencies = [ + "dioxus-cli-config", + "http 1.4.0", + "infer", + "jni", + "js-sys", + "ndk 0.9.0", + "ndk-context", + "ndk-sys 0.6.0+11769913", + "percent-encoding", + "thiserror 2.0.18", + "tokio", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "dioxus-cli-config" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "86a13d42c5defcea333bdbae1dc5d64d078acd0fda1d8a1441c37e06be5146e3" +dependencies = [ + "wasm-bindgen", +] + +[[package]] +name = "dioxus-config-macro" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ba1d68a05a8a15293ba65d45c7a3263356f3eedf1a3e599440683f3eb014637" +dependencies = [ + "proc-macro2", + "quote", +] + +[[package]] +name = "dioxus-config-macros" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f43f2d511d3c3c439a2fb7f863668b84caf8e0d2440cbfbcbb28521e26ba7f44" + +[[package]] +name = "dioxus-core" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fb3dd61889e6a09daec93d44db86047fb8e6603beedcf9351b8528582254e075" +dependencies = [ + "anyhow", + "const_format", + "dioxus-core-types", + "futures-channel", + "futures-util", + "generational-box", + "longest-increasing-subsequence", + "rustc-hash 2.1.1", + "rustversion", + "serde", + "slab", + "slotmap", + "subsecond", + "tracing", +] + +[[package]] +name = "dioxus-core-macro" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8577c4d9a8cc23423c4d2137319044b03ab940e4b2790dd25f4f06601bd32d9a" +dependencies = [ + "convert_case 0.8.0", + "dioxus-rsx", + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "dioxus-core-types" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b99d7d199aad72431b549759550002e7d72c8a257eba500dca9fbdb2122de103" + +[[package]] +name = "dioxus-devtools" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d27e7212436a581ce058d7554f1383916bd18a68ebd6015b0b4c2e9ecb0d5535" +dependencies = [ + "dioxus-cli-config", + "dioxus-core", + "dioxus-devtools-types", + "dioxus-signals", + "serde", + "serde_json", + "subsecond", + "thiserror 2.0.18", + "tracing", + "tungstenite 0.28.0", +] + +[[package]] +name = "dioxus-devtools-types" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6aa24ed651b97e0b423270bf07a0f1b7dc0e0fa1f1dc26407cd2a118d6bf9de5" +dependencies = [ + "dioxus-core", + "serde", + "subsecond-types", +] + +[[package]] +name = "dioxus-document" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24685cb51cc6227ea606c49dfe531836f362c49183d3007241afcd8827498401" +dependencies = [ + "dioxus-core", + "dioxus-core-macro", + "dioxus-core-types", + "dioxus-html", + "futures-channel", + "futures-util", + "generational-box", + "lazy-js-bundle", + "serde", + "serde_json", + "tracing", +] + +[[package]] +name = "dioxus-fullstack" +version = "0.7.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e90a04f9bfbdb42801efbb329f7d0a5be79530a856302a2090cefb2db99185b0" +dependencies = [ + "anyhow", + "async-stream", + "async-tungstenite", + "axum 0.8.9", + "axum-core 0.5.6", + "base64 0.22.1", + "bytes", + "ciborium", + "const-str", + "const_format", + "content_disposition", + "derive_more 2.1.1", + "dioxus-asset-resolver", + "dioxus-cli-config", + "dioxus-core", + "dioxus-fullstack-core", + "dioxus-fullstack-macro", + "dioxus-hooks", + "dioxus-html", + "dioxus-signals", + "form_urlencoded", + "futures", + "futures-channel", + "futures-util", + "gloo-net", + "headers", + "http 1.4.0", + "http-body 1.0.1", + "http-body-util", + "js-sys", + "mime", + "pin-project", + "reqwest 0.12.28", + "rustversion", + "send_wrapper", + "serde", + "serde_json", + "serde_qs", + "serde_urlencoded", + "thiserror 2.0.18", + "tokio-util", + "tracing", + "tungstenite 0.27.0", + "url", + "wasm-bindgen", + "wasm-bindgen-futures", + "wasm-streams", + "web-sys", + "xxhash-rust", +] + +[[package]] +name = "dioxus-fullstack-core" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "28333274cfc8e5fe547ab04258c2511350c4930a07af9616d365dc4ba7b22d8f" +dependencies = [ + "anyhow", + "axum-core 0.5.6", + "base64 0.22.1", + "ciborium", + "dioxus-core", + "dioxus-document", + "dioxus-history", + "dioxus-hooks", + "dioxus-signals", + "futures-channel", + "futures-util", + "generational-box", + "http 1.4.0", + "inventory", + "parking_lot", + "serde", + "serde_json", + "thiserror 2.0.18", + "tokio", + "tracing", +] + +[[package]] +name = "dioxus-fullstack-macro" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "53f7e5a9fa7f657aa519a07aced8b8936f3ae8a246d94855d497d8cce59b9533" +dependencies = [ + "const_format", + "convert_case 0.8.0", + "proc-macro2", + "quote", + "syn 2.0.117", + "xxhash-rust", +] + +[[package]] +name = "dioxus-history" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "010b446322b3f9176476579fa61c7552f0430abbeec418cab543482da6ca4363" +dependencies = [ + "dioxus-core", + "tracing", +] + +[[package]] +name = "dioxus-hooks" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09e7a6ba279050cc161e1215c6db0bd15915c9314ec2916d7b22c113a3039536" +dependencies = [ + "dioxus-core", + "dioxus-signals", + "futures-channel", + "futures-util", + "generational-box", + "rustversion", + "slab", + "tracing", +] + +[[package]] +name = "dioxus-html" +version = "0.7.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6edc3f7b6cc88092fdf0181ee372ed0f8d0ea507c648a6d44abae341b79b1dee" +dependencies = [ + "async-trait", + "bytes", + "dioxus-core", + "dioxus-core-macro", + "dioxus-core-types", + "dioxus-hooks", + "dioxus-html-internal-macro", + "enumset", + "euclid", + "futures-channel", + "futures-util", + "generational-box", + "keyboard-types", + "lazy-js-bundle", + "rustversion", + "tracing", +] + +[[package]] +name = "dioxus-html-internal-macro" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff6b7918b0908c8719a6165b4e3c362da4fd311fc7cb48720eddd8a45b2ddfc6" +dependencies = [ + "convert_case 0.8.0", + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "dioxus-interpreter-js" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a8ce1cf487007f90d0ec4ec87dff111d74ac04fca0918f9dcc4e80dc3b0531b2" +dependencies = [ + "js-sys", + "lazy-js-bundle", + "rustc-hash 2.1.1", + "sledgehammer_bindgen", + "sledgehammer_utils", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "dioxus-logger" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d4742b16791a71eb4db2d0747f15c50b278b27369b3d93e5a4d6ec2570bcb9bc" +dependencies = [ + "dioxus-cli-config", + "tracing", + "tracing-subscriber 0.3.23", + "tracing-wasm", +] + +[[package]] +name = "dioxus-rsx" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "344621f6dc435e76fbe272da09988d0118cf35cc2aa88ebb5ae7c1317a36e57c" +dependencies = [ + "proc-macro2", + "proc-macro2-diagnostics", + "quote", + "rustversion", + "syn 2.0.117", +] + +[[package]] +name = "dioxus-signals" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "409bf65d243443416650945f22cd6caf2a6bb13ae0347a50ec5852adb1961072" +dependencies = [ + "dioxus-core", + "futures-channel", + "futures-util", + "generational-box", + "parking_lot", + "rustc-hash 2.1.1", + "tracing", + "warnings", +] + +[[package]] +name = "dioxus-stores" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "245ec4f84348e5be77451bd204181998b8bc0995b48ff3adb2db0e0ec430dab4" +dependencies = [ + "dioxus-core", + "dioxus-signals", + "dioxus-stores-macro", + "generational-box", +] + +[[package]] +name = "dioxus-stores-macro" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dd9da8e9a1cc2d8bff387e0b99f09f2590b71f67d5d73ab343b2cc9d17990d92" +dependencies = [ + "convert_case 0.8.0", + "proc-macro2", + "quote", + "syn 2.0.117", +] + +[[package]] +name = "dioxus-web" +version = "0.7.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2506d4bf5edbefb886105f7126695be048f8b5c12fc8529a71e77764906b054a" +dependencies = [ + "dioxus-cli-config", + "dioxus-core", + "dioxus-core-types", + "dioxus-devtools", + "dioxus-document", + "dioxus-history", + "dioxus-html", + "dioxus-interpreter-js", + "dioxus-signals", + "futures-channel", + "futures-util", + "generational-box", + "gloo-timers", + "js-sys", + "lazy-js-bundle", + "rustc-hash 2.1.1", + "send_wrapper", + "serde", + "serde-wasm-bindgen", + "serde_json", + "tracing", + "wasm-bindgen", + "wasm-bindgen-futures", + "wasm-streams", + "web-sys", +] + [[package]] name = "directories" version = "5.0.1" @@ -2497,6 +3189,12 @@ dependencies = [ "dtoa", ] +[[package]] +name = "dunce" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813" + [[package]] name = "dyn-clone" version = "1.0.20" @@ -2714,6 +3412,15 @@ version = "0.1.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d817e038c30374a4bcb22f94d0a8a0e216958d4c3dcde369b1439fec4bdda6e6" +[[package]] +name = "euclid" +version = "0.22.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1a05365e3b1c6d1650318537c7460c6923f1abdd272ad6842baa2b509957a06" +dependencies = [ + "num-traits", +] + [[package]] name = "event-listener" version = "5.4.1" @@ -3573,6 +4280,16 @@ dependencies = [ "seq-macro", ] +[[package]] +name = "generational-box" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ede46ff252793f9b6ef752c506ba8600c69d73cad2ef9bbf2e6dee85019a3bc" +dependencies = [ + "parking_lot", + "tracing", +] + [[package]] name = "generativity" version = "1.1.0" @@ -3759,6 +4476,52 @@ version = "0.3.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280" +[[package]] +name = "gloo-net" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c06f627b1a58ca3d42b45d6104bf1e1a03799df472df00988b6ba21accc10580" +dependencies = [ + "futures-channel", + "futures-core", + "futures-sink", + "gloo-utils", + "http 1.4.0", + "js-sys", + "pin-project", + "serde", + "serde_json", + "thiserror 1.0.69", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "gloo-timers" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbb143cf96099802033e0d4f4963b19fd2e0b728bcf076cd9cf7f6634f092994" +dependencies = [ + "futures-channel", + "futures-core", + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "gloo-utils" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b5555354113b18c547c1d3a98fbf7fb32a9ff4f6fa112ce823a21641a0ba3aa" +dependencies = [ + "js-sys", + "serde", + "serde_json", + "wasm-bindgen", + "web-sys", +] + [[package]] name = "glow" version = "0.16.0" @@ -3957,6 +4720,30 @@ dependencies = [ "num-traits", ] +[[package]] +name = "headers" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b3314d5adb5d94bcdf56771f2e50dbbc80bb4bdf88967526706205ac9eff24eb" +dependencies = [ + "base64 0.22.1", + "bytes", + "headers-core", + "http 1.4.0", + "httpdate", + "mime", + "sha1", +] + +[[package]] +name = "headers-core" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "54b4a22553d4242c49fddb9ba998a99962b5cc6f22cb5a3482bec22522403ce4" +dependencies = [ + "http 1.4.0", +] + [[package]] name = "headless_chrome" version = "1.0.21" @@ -4605,6 +5392,15 @@ dependencies = [ "web-time", ] +[[package]] +name = "infer" +version = "0.19.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a588916bfdfd92e71cacef98a63d9b1f0d74d6599980d11894290e7ddefffcf7" +dependencies = [ + "cfb", +] + [[package]] name = "inout" version = "0.1.4" @@ -4648,6 +5444,15 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "inventory" +version = "0.3.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4f0c30c76f2f4ccee3fe55a2435f691ca00c0e4bd87abe4f4a851b1d4dac39b" +dependencies = [ + "rustversion", +] + [[package]] name = "ipnet" version = "2.12.0" @@ -4795,7 +5600,7 @@ version = "0.4.0" dependencies = [ "anyhow", "arroy", - "axum", + "axum 0.7.9", "comfy-table", "futures-util", "hdrhistogram", @@ -5088,6 +5893,15 @@ dependencies = [ name = "kalosm-workspace" version = "0.4.0" +[[package]] +name = "keyboard-types" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b750dcadc39a09dbadd74e118f6dd6598df77fa01df0cfcdc52c28dece74528a" +dependencies = [ + "bitflags 2.11.0", +] + [[package]] name = "khronos-egl" version = "6.0.0" @@ -5136,6 +5950,12 @@ dependencies = [ "regex-automata 0.4.14", ] +[[package]] +name = "lazy-js-bundle" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "60d7adc10cb9440d17fa67e467febdfc98931338773d11bfee81809af54d0697" + [[package]] name = "lazy_static" version = "1.5.0" @@ -5295,6 +6115,12 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "longest-increasing-subsequence" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b3bd0dd2cd90571056fdb71f6275fada10131182f84899f4b2a916e565d81d86" + [[package]] name = "loop9" version = "0.1.5" @@ -5374,6 +6200,17 @@ dependencies = [ "libc", ] +[[package]] +name = "macro-string" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1b27834086c65ec3f9387b096d66e99f221cf081c2b738042aa252bcd41204e3" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "macro_rules_attribute" version = "0.2.2" @@ -5399,6 +6236,50 @@ dependencies = [ "libc", ] +[[package]] +name = "manganis" +version = "0.7.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3271e4d23afe07537293d1d44cf4a1793102c35713bdd04a547fa5c0f6eef373" +dependencies = [ + "const-serialize 0.7.2", + "const-serialize 0.8.0-alpha.0", + "jni", + "manganis-core", + "manganis-macro", + "ndk-context", + "objc2", + "thiserror 2.0.18", +] + +[[package]] +name = "manganis-core" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1b84cc2951f3b119702fab499b9b1aec3f454929c62feca55b895b82c628308" +dependencies = [ + "const-serialize 0.7.2", + "const-serialize 0.8.0-alpha.0", + "dioxus-cli-config", + "dioxus-core-types", + "serde", + "winnow", +] + +[[package]] +name = "manganis-macro" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6d2e60d36758b201b6ebb8a31aff6b013e58924eeb6d3cbf19aea764f51d69e4" +dependencies = [ + "dunce", + "macro-string", + "manganis-core", + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "maplit" version = "1.0.2" @@ -5476,12 +6357,27 @@ dependencies = [ "regex-automata 0.1.10", ] +[[package]] +name = "matchers" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1525a2a28c7f4fa0fc98bb91ae755d1e2d1505079e05539e35bc876b5d65ae9" +dependencies = [ + "regex-automata 0.4.14", +] + [[package]] name = "matchit" version = "0.7.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0e7465ac9959cc2b1404e8e2367b43684a6d13790fe23056cc8c6c5a6b7bcb94" +[[package]] +name = "matchit" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" + [[package]] name = "matrixmultiply" version = "0.3.10" @@ -5518,6 +6414,15 @@ version = "2.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" +[[package]] +name = "memfd" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ad38eb12aea514a0466ea40a80fd8cc83637065948eb4a426e4aa46261175227" +dependencies = [ + "rustix", +] + [[package]] name = "memmap2" version = "0.9.10" @@ -5786,6 +6691,21 @@ dependencies = [ "thiserror 1.0.69", ] +[[package]] +name = "ndk" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3f42e7bbe13d351b6bead8286a43aac9534b82bd3cc43e47037f012ebfd62d4" +dependencies = [ + "bitflags 2.11.0", + "jni-sys", + "log", + "ndk-sys 0.6.0+11769913", + "num_enum", + "raw-window-handle", + "thiserror 1.0.69", +] + [[package]] name = "ndk-context" version = "0.1.1" @@ -6162,7 +7082,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e8b61bebd49e5d43f5f8cc7ee2891c16e0f41ec7954d36bcb6c14c5e0de867fb" dependencies = [ "jni", - "ndk", + "ndk 0.8.0", "ndk-context", "num-derive", "num-traits", @@ -6870,6 +7790,18 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "proc-macro2-diagnostics" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af066a9c399a26e020ada66a034357a868728e72cd426f3adcd35f80d88d88c8" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", + "version_check", +] + [[package]] name = "profiling" version = "1.0.17" @@ -6925,6 +7857,16 @@ dependencies = [ "syn 1.0.109", ] +[[package]] +name = "publicsuffix" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f42ea446cab60335f76979ec15e12619a2165b5ae2c12166bef27d283a9fadf" +dependencies = [ + "idna", + "psl-types", +] + [[package]] name = "pulldown-cmark" version = "0.9.6" @@ -7562,6 +8504,8 @@ checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" dependencies = [ "base64 0.22.1", "bytes", + "cookie", + "cookie_store", "encoding_rs", "futures-core", "futures-util", @@ -7692,6 +8636,19 @@ dependencies = [ "tracing", ] +[[package]] +name = "rgliner-web" +version = "0.1.0" +dependencies = [ + "console_error_panic_hook", + "dioxus", + "getrandom 0.3.4", + "rgliner", + "tracing", + "tracing-wasm", + "wasm-bindgen-futures", +] + [[package]] name = "ring" version = "0.17.14" @@ -8226,7 +9183,7 @@ checksum = "4eb30575f3638fc8f6815f448d50cb1a2e255b0897985c8c59f4d37b72a07b06" dependencies = [ "bitflags 2.11.0", "cssparser 0.31.2", - "derive_more", + "derive_more 0.99.20", "fxhash", "log", "new_debug_unreachable", @@ -8258,6 +9215,9 @@ name = "send_wrapper" version = "0.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cd0b0ec5f1c1ca621c432a25813d8d60c88abe6d3e08a3eb9cf37d97a0fe3d73" +dependencies = [ + "futures-core", +] [[package]] name = "seq-macro" @@ -8284,6 +9244,17 @@ dependencies = [ "serde", ] +[[package]] +name = "serde-wasm-bindgen" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8302e169f0eddcc139c70f139d19d6467353af16f9fce27e8c30158036a1e16b" +dependencies = [ + "js-sys", + "serde", + "wasm-bindgen", +] + [[package]] name = "serde-xml-rs" version = "0.4.1" @@ -8360,6 +9331,17 @@ dependencies = [ "serde", ] +[[package]] +name = "serde_qs" +version = "0.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f3faaf9e727533a19351a43cc5a8de957372163c7d35cc48c90b75cdda13c352" +dependencies = [ + "percent-encoding", + "serde", + "thiserror 2.0.18", +] + [[package]] name = "serde_regex" version = "1.1.0" @@ -8532,12 +9514,42 @@ dependencies = [ "serde", ] +[[package]] +name = "sledgehammer_bindgen" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "49e83e178d176459c92bc129cfd0958afac3ced925471b889b3a75546cfc4133" +dependencies = [ + "sledgehammer_bindgen_macro", + "wasm-bindgen", +] + +[[package]] +name = "sledgehammer_bindgen_macro" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb251b407f50028476a600541542b605bb864d35d9ee1de4f6cab45d88475e6d" +dependencies = [ + "quote", + "syn 2.0.117", +] + +[[package]] +name = "sledgehammer_utils" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "debdd4b83524961983cea3c55383b3910fd2f24fd13a188f5b091d2d504a61ae" +dependencies = [ + "rustc-hash 1.1.0", +] + [[package]] name = "slotmap" version = "1.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bdd58c3c93c3d278ca835519292445cb4b0d4dc59ccfdf7ceadaab3f8aeb4038" dependencies = [ + "serde", "version_check", ] @@ -8785,6 +9797,34 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "subsecond" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5dbb9f2928b6654ccc28d4ddfef5213e97ed66afed4907774d049b376c62a838" +dependencies = [ + "js-sys", + "libc", + "libloading 0.8.9", + "memfd", + "memmap2", + "serde", + "subsecond-types", + "thiserror 2.0.18", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "subsecond-types" +version = "0.7.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "388bb28e6ddbee717745963b8932d9a6e24a5d3c93350655f733e938de04d81f" +dependencies = [ + "serde", +] + [[package]] name = "subtle" version = "2.6.1" @@ -9690,7 +10730,7 @@ dependencies = [ "ansi_term", "chrono", "lazy_static", - "matchers", + "matchers 0.0.1", "regex", "serde", "serde_json", @@ -9709,14 +10749,29 @@ version = "0.3.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319" dependencies = [ + "matchers 0.2.0", "nu-ansi-term", + "once_cell", + "regex-automata 0.4.14", "sharded-slab", "smallvec 1.15.1", "thread_local", + "tracing", "tracing-core", "tracing-log 0.2.0", ] +[[package]] +name = "tracing-wasm" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4575c663a174420fa2d78f4108ff68f65bf2fbb7dd89f33749b6e826b3626e07" +dependencies = [ + "tracing", + "tracing-subscriber 0.3.23", + "wasm-bindgen", +] + [[package]] name = "transpose" version = "0.2.3" @@ -9765,6 +10820,23 @@ dependencies = [ "utf-8", ] +[[package]] +name = "tungstenite" +version = "0.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eadc29d668c91fcc564941132e17b28a7ceb2f3ebf0b9dae3e03fd7a6748eb0d" +dependencies = [ + "bytes", + "data-encoding", + "http 1.4.0", + "httparse", + "log", + "rand 0.9.2", + "sha1", + "thiserror 2.0.18", + "utf-8", +] + [[package]] name = "tungstenite" version = "0.28.0" @@ -10158,6 +11230,28 @@ dependencies = [ "try-lock", ] +[[package]] +name = "warnings" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "64f68998838dab65727c9b30465595c6f7c953313559371ca8bf31759b3680ad" +dependencies = [ + "pin-project", + "tracing", + "warnings-macro", +] + +[[package]] +name = "warnings-macro" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59195a1db0e95b920366d949ba5e0d3fc0e70b67c09be15ce5abb790106b0571" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "wasi" version = "0.11.1+wasi-snapshot-preview1" @@ -11246,6 +12340,12 @@ dependencies = [ "markup5ever 0.11.0", ] +[[package]] +name = "xxhash-rust" +version = "0.8.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fdd20c5420375476fbd4394763288da7eb0cc0b8c11deed431a91562af7335d3" + [[package]] name = "y4m" version = "0.8.0" diff --git a/Cargo.toml b/Cargo.toml index 208e09fb4..294a3863e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -37,6 +37,7 @@ members = [ "fusor-ml/cpu", "fusor-ml/fusor", "fusor-ml/types", + "demos/rgliner-web", ] [workspace.dependencies] diff --git a/demos/rgliner-web/Cargo.toml b/demos/rgliner-web/Cargo.toml new file mode 100644 index 000000000..b9edaa640 --- /dev/null +++ b/demos/rgliner-web/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "rgliner-web" +version = "0.1.0" +edition = "2021" +publish = false + +[dependencies] +dioxus = { version = "=0.7.2", features = ["web"] } +rgliner = { path = "../../models/rgliner", default-features = false } +getrandom = { version = "0.3", features = ["wasm_js"] } +tracing = "0.1" + +[target.'cfg(target_arch = "wasm32")'.dependencies] +console_error_panic_hook = "0.1" +tracing-wasm = "0.2" +wasm-bindgen-futures = "0.4" diff --git a/demos/rgliner-web/Dioxus.toml b/demos/rgliner-web/Dioxus.toml new file mode 100644 index 000000000..0ef3fe34e --- /dev/null +++ b/demos/rgliner-web/Dioxus.toml @@ -0,0 +1,5 @@ +[application] +name = "rgliner-web" + +[web.app] +title = "rgliner — GLiNER NER + Relation Extraction in your browser" diff --git a/demos/rgliner-web/assets/style.css b/demos/rgliner-web/assets/style.css new file mode 100644 index 000000000..0cc1f6fe2 --- /dev/null +++ b/demos/rgliner-web/assets/style.css @@ -0,0 +1,234 @@ +:root { + --bg: #0b0d12; + --panel: #141821; + --panel-border: #232833; + --text: #e6e9ef; + --muted: #8b92a3; + --accent: #7aa7ff; + --accent-strong: #4a7fff; + --err: #ff6b6b; + --ok: #4ade80; + --radius: 8px; +} + +* { + box-sizing: border-box; +} + +body { + margin: 0; + font-family: ui-sans-serif, system-ui, -apple-system, "Segoe UI", Roboto, sans-serif; + background: var(--bg); + color: var(--text); + line-height: 1.5; +} + +.app { + max-width: 960px; + margin: 0 auto; + padding: 2rem 1.5rem 4rem; +} + +header.site-header { + display: flex; + align-items: baseline; + justify-content: space-between; + gap: 1rem; + margin-bottom: 1.5rem; +} + +header.site-header h1 { + font-size: 1.5rem; + margin: 0; +} + +header.site-header .tag { + color: var(--muted); + font-size: 0.9rem; +} + +header.site-header a { + color: var(--accent); + text-decoration: none; + font-size: 0.9rem; +} + +.tabs { + display: flex; + gap: 0.25rem; + border-bottom: 1px solid var(--panel-border); + margin-bottom: 1.5rem; +} + +.tab { + background: none; + border: none; + color: var(--muted); + padding: 0.6rem 1rem; + cursor: pointer; + font-size: 0.95rem; + border-bottom: 2px solid transparent; +} + +.tab:hover { + color: var(--text); +} + +.tab.active { + color: var(--text); + border-bottom-color: var(--accent); +} + +.panel { + background: var(--panel); + border: 1px solid var(--panel-border); + border-radius: var(--radius); + padding: 1rem 1.25rem; + margin-bottom: 1rem; +} + +.row { + display: flex; + gap: 0.75rem; + align-items: center; + flex-wrap: wrap; +} + +label { + color: var(--muted); + font-size: 0.85rem; + display: block; + margin-bottom: 0.35rem; +} + +input[type="text"], +textarea, +select { + background: #0f1218; + color: var(--text); + border: 1px solid var(--panel-border); + border-radius: 6px; + padding: 0.55rem 0.7rem; + font-family: inherit; + font-size: 0.95rem; + width: 100%; +} + +textarea { + min-height: 110px; + resize: vertical; + font-family: ui-monospace, SFMono-Regular, "SF Mono", Menlo, monospace; + font-size: 0.9rem; +} + +select { + width: auto; + min-width: 220px; +} + +button.primary, +button.secondary { + background: var(--accent-strong); + color: white; + border: none; + border-radius: 6px; + padding: 0.55rem 1rem; + cursor: pointer; + font-size: 0.95rem; + font-weight: 500; +} + +button.primary:hover:not(:disabled) { + background: var(--accent); +} + +button.secondary { + background: transparent; + color: var(--text); + border: 1px solid var(--panel-border); +} + +button:disabled { + opacity: 0.5; + cursor: not-allowed; +} + +.status { + font-size: 0.85rem; + color: var(--muted); +} + +.status.ok { + color: var(--ok); +} + +.status.err { + color: var(--err); +} + +.err-banner { + background: rgba(255, 107, 107, 0.08); + border: 1px solid rgba(255, 107, 107, 0.4); + color: #ffb3b3; + padding: 0.75rem 1rem; + border-radius: 6px; + margin-bottom: 1rem; + font-size: 0.9rem; +} + +.results { + line-height: 1.9; + font-size: 1rem; + word-wrap: break-word; +} + +.entity { + border-radius: 4px; + padding: 0.05em 0.3em; + margin: 0 0.05em; + font-weight: 500; + color: #0b0d12; +} + +.entity .chip { + font-size: 0.7em; + font-weight: 700; + padding: 0 0.35em; + margin-left: 0.3em; + border-radius: 3px; + background: rgba(0, 0, 0, 0.25); + color: #fff; + text-transform: uppercase; + letter-spacing: 0.04em; +} + +.entity-list, +.relation-list { + margin-top: 1rem; + display: flex; + flex-direction: column; + gap: 0.35rem; +} + +.entity-list li, +.relation-list li { + list-style: none; + font-size: 0.9rem; + color: var(--muted); + font-family: ui-monospace, SFMono-Regular, "SF Mono", Menlo, monospace; +} + +.entity-list .label, +.relation-list .label { + color: var(--text); + font-weight: 600; +} + +.score { + color: var(--accent); +} + +.muted { + color: var(--muted); + font-size: 0.85rem; +} diff --git a/demos/rgliner-web/src/main.rs b/demos/rgliner-web/src/main.rs new file mode 100644 index 000000000..4d09cd942 --- /dev/null +++ b/demos/rgliner-web/src/main.rs @@ -0,0 +1,567 @@ +use dioxus::prelude::*; +use rgliner::{ + relation_decoding::Relation, + relex::{GlinerRelEx, GlinerRelExSource}, + DecodingMode, Entity, Gliner, GlinerSource, +}; + +fn main() { + #[cfg(target_arch = "wasm32")] + { + console_error_panic_hook::set_once(); + tracing_wasm::set_as_global_default(); + } + dioxus::launch(App); +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum Mode { + Ner, + Relex, +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum ModelChoice { + Edge, + Small, + Base, + Large, + RelexMulti, + RelexBase, + RelexLarge, +} + +impl ModelChoice { + fn label(self) -> &'static str { + match self { + ModelChoice::Edge => "edge · 60M · fastest", + ModelChoice::Small => "small · 108M", + ModelChoice::Base => "base · 194M", + ModelChoice::Large => "large · 530M · best NER", + ModelChoice::RelexMulti => "relex-multi · multilingual", + ModelChoice::RelexBase => "relex-base · English", + ModelChoice::RelexLarge => "relex-large · English · best", + } + } + + fn value(self) -> &'static str { + match self { + ModelChoice::Edge => "edge", + ModelChoice::Small => "small", + ModelChoice::Base => "base", + ModelChoice::Large => "large", + ModelChoice::RelexMulti => "relex-multi", + ModelChoice::RelexBase => "relex-base", + ModelChoice::RelexLarge => "relex-large", + } + } + + fn from_value(v: &str) -> Option { + Some(match v { + "edge" => ModelChoice::Edge, + "small" => ModelChoice::Small, + "base" => ModelChoice::Base, + "large" => ModelChoice::Large, + "relex-multi" => ModelChoice::RelexMulti, + "relex-base" => ModelChoice::RelexBase, + "relex-large" => ModelChoice::RelexLarge, + _ => return None, + }) + } + + fn default_for(mode: Mode) -> Self { + match mode { + Mode::Ner => ModelChoice::Edge, + Mode::Relex => ModelChoice::RelexMulti, + } + } + + fn for_mode(mode: Mode) -> &'static [ModelChoice] { + match mode { + Mode::Ner => &[ + ModelChoice::Edge, + ModelChoice::Small, + ModelChoice::Base, + ModelChoice::Large, + ], + Mode::Relex => &[ + ModelChoice::RelexMulti, + ModelChoice::RelexBase, + ModelChoice::RelexLarge, + ], + } + } +} + +enum LoadedModel { + Ner { + choice: ModelChoice, + inner: Gliner, + }, + Relex { + choice: ModelChoice, + inner: GlinerRelEx, + }, +} + +impl LoadedModel { + fn choice(&self) -> ModelChoice { + match self { + LoadedModel::Ner { choice, .. } | LoadedModel::Relex { choice, .. } => *choice, + } + } +} + +#[derive(Clone, Default)] +struct RelexResult { + entities: Vec, + relations: Vec, +} + +#[component] +fn App() -> Element { + let mut mode = use_signal(|| Mode::Ner); + let mut choice = use_signal(|| ModelChoice::Edge); + + let mut text = use_signal(|| { + "Apple Inc. was founded by Steve Jobs in California. Microsoft is headquartered in Redmond." + .to_string() + }); + let mut entity_labels = + use_signal(|| "person, organization, location".to_string()); + let mut relation_labels = use_signal(|| "founded by, located in".to_string()); + + let mut model = use_signal(|| None::); + let mut loading = use_signal(|| false); + let mut running = use_signal(|| false); + let mut error = use_signal(|| None::); + let mut ner_out = use_signal(Vec::::new); + let mut relex_out = use_signal(RelexResult::default); + let mut status = use_signal(|| "No model loaded".to_string()); + + let mut on_mode_change = move |new_mode: Mode| { + if mode() != new_mode { + mode.set(new_mode); + choice.set(ModelChoice::default_for(new_mode)); + ner_out.write().clear(); + *relex_out.write() = RelexResult::default(); + } + }; + + let on_load = move |_| { + if loading() { + return; + } + let selected = choice(); + loading.set(true); + error.set(None); + status.set(format!("Loading {}…", selected.label())); + spawn(async move { + match build_model(selected).await { + Ok(m) => { + let dev = match &m { + LoadedModel::Ner { inner, .. } => { + if inner.device().is_gpu() { "GPU" } else { "CPU" } + } + LoadedModel::Relex { inner, .. } => { + if inner.device().is_gpu() { "GPU" } else { "CPU" } + } + }; + model.set(Some(m)); + status.set(format!("{} ready on {dev}", selected.label())); + } + Err(e) => { + error.set(Some(format!("{e}"))); + status.set("Load failed".to_string()); + } + } + loading.set(false); + }); + }; + + let on_extract = move |_| { + if running() { + return; + } + let current_text = text(); + let ent_labels = parse_labels(&entity_labels()); + let rel_labels = parse_labels(&relation_labels()); + let current_mode = mode(); + + running.set(true); + error.set(None); + + let Some(mut taken) = model.write().take() else { + error.set(Some("Load a model first.".to_string())); + running.set(false); + return; + }; + + spawn(async move { + let outcome = run_extraction( + &mut taken, + current_mode, + ¤t_text, + &ent_labels, + &rel_labels, + ) + .await; + + match outcome { + Ok(ExtractionOutput::Ner(entities)) => { + status.set(format!("Extracted {} entities", entities.len())); + ner_out.set(entities); + } + Ok(ExtractionOutput::Relex(result)) => { + status.set(format!( + "Extracted {} entities, {} relations", + result.entities.len(), + result.relations.len() + )); + relex_out.set(result); + } + Err(e) => error.set(Some(e)), + } + + model.set(Some(taken)); + running.set(false); + }); + }; + + let has_model = model.read().is_some(); + let model_mismatch = model + .read() + .as_ref() + .map(|m| m.choice() != choice()) + .unwrap_or(true); + + rsx! { + document::Link { rel: "stylesheet", href: asset!("/assets/style.css") } + div { class: "app", + header { class: "site-header", + div { + h1 { "rgliner" } + div { class: "tag", + "GLiNER NER & relation extraction — running locally in your browser with WebGPU." + } + } + a { + href: "https://github.com/floneum/floneum", + target: "_blank", + "GitHub" + } + } + + div { class: "tabs", + button { + class: if mode() == Mode::Ner { "tab active" } else { "tab" }, + onclick: move |_| on_mode_change(Mode::Ner), + "NER" + } + button { + class: if mode() == Mode::Relex { "tab active" } else { "tab" }, + onclick: move |_| on_mode_change(Mode::Relex), + "NER + Relations" + } + } + + if let Some(e) = error() { + div { class: "err-banner", "{e}" } + } + + div { class: "panel", + label { "Model" } + div { class: "row", + select { + value: "{choice().value()}", + onchange: move |ev| { + if let Some(c) = ModelChoice::from_value(&ev.value()) { + choice.set(c); + } + }, + for c in ModelChoice::for_mode(mode()) { + option { value: "{c.value()}", "{c.label()}" } + } + } + button { + class: "primary", + disabled: loading() || running() || (has_model && !model_mismatch), + onclick: on_load, + if loading() { "Loading…" } + else if has_model && !model_mismatch { "Loaded" } + else if has_model { "Reload" } + else { "Load model" } + } + span { + class: if error().is_some() { "status err" } else if has_model { "status ok" } else { "status" }, + "{status()}" + } + } + p { class: "muted", + "First load fetches the GGUF weights (60 MB – 500 MB) from HuggingFace and caches them in the browser's Origin Private File System. Subsequent loads are instant." + } + } + + div { class: "panel", + label { "Text" } + textarea { + value: "{text}", + oninput: move |e| text.set(e.value()), + } + } + + div { class: "panel", + label { "Entity labels (comma-separated)" } + input { + r#type: "text", + value: "{entity_labels}", + oninput: move |e| entity_labels.set(e.value()), + } + if mode() == Mode::Relex { + div { style: "margin-top: 0.75rem;", + label { "Relation labels (comma-separated)" } + input { + r#type: "text", + value: "{relation_labels}", + oninput: move |e| relation_labels.set(e.value()), + } + } + } + } + + div { class: "panel", + div { class: "row", + button { + class: "primary", + disabled: !has_model || running() || loading() || model_mismatch, + onclick: on_extract, + if running() { "Extracting…" } else { "Extract" } + } + if model_mismatch && has_model { + span { class: "status", + "Model selection changed — reload to use it." + } + } + } + } + + { render_results(mode(), text(), ner_out(), relex_out()) } + } + } +} + +fn parse_labels(raw: &str) -> Vec { + raw.split(',') + .map(|s| s.trim()) + .filter(|s| !s.is_empty()) + .map(|s| s.to_string()) + .collect() +} + +async fn build_model( + choice: ModelChoice, +) -> Result { + match choice { + ModelChoice::Edge + | ModelChoice::Small + | ModelChoice::Base + | ModelChoice::Large => { + let source = match choice { + ModelChoice::Edge => GlinerSource::edge(), + ModelChoice::Small => GlinerSource::small(), + ModelChoice::Base => GlinerSource::base(), + ModelChoice::Large => GlinerSource::large(), + _ => unreachable!(), + }; + let inner = Gliner::builder() + .with_source(source) + .with_decoding_mode(DecodingMode::Flat) + .with_threshold(0.3) + .build_with_loading_handler(|_| {}) + .await + .map_err(|e| format!("{e}"))?; + Ok(LoadedModel::Ner { choice, inner }) + } + ModelChoice::RelexMulti | ModelChoice::RelexBase | ModelChoice::RelexLarge => { + let source = match choice { + ModelChoice::RelexMulti => GlinerRelExSource::relex_multi(), + ModelChoice::RelexBase => GlinerRelExSource::relex_base(), + ModelChoice::RelexLarge => GlinerRelExSource::relex_large(), + _ => unreachable!(), + }; + let inner = GlinerRelEx::builder() + .with_source(source) + .build_with_loading_handler(|_| {}) + .await + .map_err(|e| format!("{e}"))?; + Ok(LoadedModel::Relex { choice, inner }) + } + } +} + +enum ExtractionOutput { + Ner(Vec), + Relex(RelexResult), +} + +async fn run_extraction( + model: &mut LoadedModel, + mode: Mode, + text: &str, + entity_labels: &[String], + relation_labels: &[String], +) -> Result { + let ent_refs: Vec<&str> = entity_labels.iter().map(|s| s.as_str()).collect(); + + if ent_refs.is_empty() { + return Err("Add at least one entity label.".to_string()); + } + + match (mode, model) { + (Mode::Ner, LoadedModel::Ner { inner, .. }) => { + let entities = inner + .extract(text, &ent_refs) + .await + .map_err(|e| format!("{e}"))?; + Ok(ExtractionOutput::Ner(entities)) + } + (Mode::Relex, LoadedModel::Relex { inner, .. }) => { + let rel_refs: Vec<&str> = relation_labels.iter().map(|s| s.as_str()).collect(); + if rel_refs.is_empty() { + return Err("Add at least one relation label.".to_string()); + } + let (entities, relations) = inner + .extract(text, &ent_refs, &rel_refs) + .await + .map_err(|e| format!("{e}"))?; + Ok(ExtractionOutput::Relex(RelexResult { + entities, + relations, + })) + } + _ => Err("Loaded model doesn't match the current tab. Reload the model.".to_string()), + } +} + +fn render_results( + mode: Mode, + text: String, + ner_out: Vec, + relex_out: RelexResult, +) -> Element { + let (entities, relations): (Vec, Vec) = match mode { + Mode::Ner => (ner_out, Vec::new()), + Mode::Relex => (relex_out.entities, relex_out.relations), + }; + + if entities.is_empty() && relations.is_empty() { + return rsx! { + div { class: "panel", + p { class: "muted", "Run extraction to see results." } + } + }; + } + + rsx! { + div { class: "panel", + div { class: "results", + { highlighted_text(text.clone(), entities.clone()) } + } + + if !entities.is_empty() { + ul { class: "entity-list", + for (i, ent) in entities.iter().enumerate() { + li { key: "{i}", + span { class: "label", style: "color: {hsl_for(&ent.label)};", "{ent.label}" } + " · " + span { "{ent.text:?}" } + " " + span { class: "score", "{format_score(ent.score)}" } + } + } + } + } + + if !relations.is_empty() { + div { style: "margin-top: 1rem;", + label { "Relations" } + ul { class: "relation-list", + for (i, rel) in relations.iter().enumerate() { + li { key: "rel-{i}", + span { class: "label", "{rel.head.text}" } + " --[" + span { style: "color: var(--accent);", "{rel.relation}" } + "]--> " + span { class: "label", "{rel.tail.text}" } + " " + span { class: "score", "{format_score(rel.score)}" } + } + } + } + } + } + } + } +} + +fn highlighted_text(text: String, entities: Vec) -> Element { + let mut sorted = entities.clone(); + sorted.sort_by_key(|e| e.start_char); + + let mut segments: Vec<(bool, String, Option)> = Vec::new(); + let mut cursor = 0usize; + for ent in sorted.iter() { + if ent.start_char < cursor { + continue; + } + if ent.start_char > cursor { + segments.push((false, text[cursor..ent.start_char].to_string(), None)); + } + let end = ent.end_char.min(text.len()); + if end > ent.start_char { + segments.push(( + true, + text[ent.start_char..end].to_string(), + Some(ent.label.clone()), + )); + } + cursor = end; + } + if cursor < text.len() { + segments.push((false, text[cursor..].to_string(), None)); + } + + rsx! { + for (i, (is_entity, content, label)) in segments.into_iter().enumerate() { + if is_entity { + { + let label_text = label.clone().unwrap_or_default(); + let color = hsl_for(&label_text); + rsx! { + span { + key: "seg-{i}", + class: "entity", + style: "background-color: {color};", + "{content}" + span { class: "chip", "{label_text}" } + } + } + } + } else { + span { key: "seg-{i}", "{content}" } + } + } + } +} + +fn hsl_for(label: &str) -> String { + let hash: u32 = label + .bytes() + .fold(0u32, |acc, b| acc.wrapping_mul(31).wrapping_add(b as u32)); + let hue = hash % 360; + format!("hsl({hue}, 70%, 72%)") +} + +fn format_score(score: f32) -> String { + format!("{:.2}", score) +} diff --git a/fusor-ml/core/src/tensor.rs b/fusor-ml/core/src/tensor.rs index 41e6b2286..7457eda8f 100644 --- a/fusor-ml/core/src/tensor.rs +++ b/fusor-ml/core/src/tensor.rs @@ -904,7 +904,7 @@ impl Tensor { where D: FloatDataType, { - #[cfg(debug_assertions)] + #[cfg(all(debug_assertions, not(target_arch = "wasm32")))] { use pollster::FutureExt as _; let as_slice = self.as_slice().block_on().unwrap(); diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs index 4f3974ea3..303b7cea6 100644 --- a/models/rgliner/src/lib.rs +++ b/models/rgliner/src/lib.rs @@ -93,24 +93,14 @@ use fusor::{Device, Tensor, VarBuilder}; use kalosm_common::Cache; use kalosm_model_types::ModelLoadingProgress; use rbert::BertSource; -use std::sync::{Arc, Mutex, OnceLock}; +use std::sync::Arc; use tokenizers::Tokenizer; use crate::raw::{CachedLabels, LabelEncoder, Scorer, SpanLayer, TextEncoder}; use crate::tokenization::{first_subtoken_pooling, WordTokenizer}; -fn default_device() -> Device { - static PANIC_HOOK_LOCK: OnceLock> = OnceLock::new(); - - let lock = PANIC_HOOK_LOCK.get_or_init(|| Mutex::new(())); - let _guard = lock.lock().unwrap_or_else(|poisoned| poisoned.into_inner()); - - let hook = std::panic::take_hook(); - std::panic::set_hook(Box::new(|_| {})); - let result = std::panic::catch_unwind(Device::gpu_blocking); - std::panic::set_hook(hook); - - result.ok().and_then(Result::ok).unwrap_or_else(Device::cpu) +async fn default_device() -> Device { + Device::gpu().await.unwrap_or_else(|_| Device::cpu()) } /// Builder for constructing a [`Gliner`] model. @@ -266,7 +256,7 @@ impl Gliner { // Initialize device let device = match device { Some(device) => device, - None => default_device(), + None => default_device().await, }; // Load text encoder diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index 3ef546da7..46ffc2ada 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -82,9 +82,9 @@ impl GlinerRelExSource { pub fn relex_multi() -> Self { Self { model: FileSource::huggingface( - "knowledgator/gliner-relex-multi-v1.0-gguf".to_string(), + "Demonthos/gliner-gguf".to_string(), "main".to_string(), - "gliner-relex-multi-v1.0-Q8_0.gguf".to_string(), + "gliner-relex-multi-v1.0-Q4_K.gguf".to_string(), ), tokenizer: None, config: None, @@ -100,9 +100,9 @@ impl GlinerRelExSource { pub fn relex_base() -> Self { Self { model: FileSource::huggingface( - "knowledgator/gliner-relex-base-v1.0-gguf".to_string(), + "Demonthos/gliner-gguf".to_string(), "main".to_string(), - "gliner-relex-base-v1.0-Q8_0.gguf".to_string(), + "gliner-relex-base-v1.0-Q4_K.gguf".to_string(), ), tokenizer: None, config: None, @@ -119,9 +119,9 @@ impl GlinerRelExSource { pub fn relex_large() -> Self { Self { model: FileSource::huggingface( - "knowledgator/gliner-relex-large-v1.0-gguf".to_string(), + "Demonthos/gliner-gguf".to_string(), "main".to_string(), - "gliner-relex-large-v1.0-Q8_0.gguf".to_string(), + "gliner-relex-large-v1.0-Q4_K.gguf".to_string(), ), tokenizer: None, config: None, @@ -291,11 +291,8 @@ pub struct GlinerRelEx { config: GlinerRelExConfig, } -fn default_device() -> Device { - std::panic::catch_unwind(Device::gpu_blocking) - .ok() - .and_then(Result::ok) - .unwrap_or_else(Device::cpu) +async fn default_device() -> Device { + Device::gpu().await.unwrap_or_else(|_| Device::cpu()) } impl GlinerRelEx { @@ -331,7 +328,10 @@ impl GlinerRelEx { .await?; // Initialize device - let device = device.unwrap_or_else(default_device); + let device = match device { + Some(d) => d, + None => default_device().await, + }; // Load model components from GGUF let mut model_cursor = std::io::Cursor::new(&model_bytes); diff --git a/models/rgliner/src/source.rs b/models/rgliner/src/source.rs index 47ba2079d..e563dfeeb 100644 --- a/models/rgliner/src/source.rs +++ b/models/rgliner/src/source.rs @@ -84,23 +84,19 @@ impl GlinerSource { std::env::var_os("HOME").map(|home| PathBuf::from(home).join(".cache").join("huggingface")) } - /// GLiNER bi-encoder v2.0 Edge variant (60M parameters). + fn demonthos_gguf(file: &str) -> FileSource { + Self::huggingface_or_cached("Demonthos/gliner-gguf", "main", file) + } + + /// GLiNER bi-encoder v2.0 Edge variant (60M parameters, Q4_K). /// /// The smallest and fastest variant, using: /// - Text encoder: ettin-encoder-32m /// - Label encoder: all-MiniLM-L6-v2 pub fn edge() -> Self { Self { - model: FileSource::huggingface( - "knowledgator/gliner-bi-edge-v2.0-gguf".to_string(), - "main".to_string(), - "gliner-bi-edge-v2.0-Q8_0.gguf".to_string(), - ), - label_encoder: FileSource::huggingface( - "knowledgator/gliner-bi-edge-v2.0-gguf".to_string(), - "main".to_string(), - "label-encoder-Q8_0.gguf".to_string(), - ), + model: Self::demonthos_gguf("gliner-bi-edge-v2.0-Q4_K.gguf"), + label_encoder: Self::demonthos_gguf("gliner-bi-edge-v2.0-Q4_K-label-encoder.gguf"), label_encoder_config: FileSource::huggingface( "sentence-transformers/all-MiniLM-L6-v2".to_string(), "main".to_string(), @@ -124,171 +120,15 @@ impl GlinerSource { } } - /// Demonthos GLiNER GGUF edge upload. - /// - /// Uses the GGUF weights and sidecar tokenizer/config files from - /// `Demonthos/gliner-gguf`. - pub fn demonthos_edge() -> Self { - Self { - model: Self::huggingface_or_cached("Demonthos/gliner-gguf", "main", "gliner-edge.gguf"), - label_encoder: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "gliner-edge-label-encoder.gguf", - ), - label_encoder_config: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "edge-label-encoder-config.json", - ), - label_encoder_tokenizer: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "edge-label-encoder-tokenizer.json", - ), - tokenizer: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "edge-text-tokenizer.json", - ), - config: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "edge-text-gliner-config.json", - ), - } - } - - /// Demonthos GLiNER GGUF small upload. - /// - /// Uses the GGUF weights and sidecar tokenizer/config files from - /// `Demonthos/gliner-gguf`. - pub fn demonthos_small() -> Self { - Self { - model: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "gliner-small.gguf", - ), - label_encoder: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "gliner-small-label-encoder.gguf", - ), - label_encoder_config: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "small-label-encoder-config.json", - ), - label_encoder_tokenizer: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "small-label-encoder-tokenizer.json", - ), - tokenizer: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "small-text-tokenizer.json", - ), - config: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "small-text-gliner-config.json", - ), - } - } - - /// Demonthos GLiNER GGUF base upload. - /// - /// Uses the GGUF weights and sidecar tokenizer/config files from - /// `Demonthos/gliner-gguf`. - pub fn demonthos_base() -> Self { - Self { - model: Self::huggingface_or_cached("Demonthos/gliner-gguf", "main", "gliner-base.gguf"), - label_encoder: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "gliner-base-label-encoder.gguf", - ), - label_encoder_config: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "base-label-encoder-config.json", - ), - label_encoder_tokenizer: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "base-label-encoder-tokenizer.json", - ), - tokenizer: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "base-text-tokenizer.json", - ), - config: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "base-text-gliner-config.json", - ), - } - } - - /// Demonthos GLiNER GGUF large upload. - /// - /// Uses the GGUF weights and sidecar tokenizer/config files from - /// `Demonthos/gliner-gguf`. - pub fn demonthos_large() -> Self { - Self { - model: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "gliner-large.gguf", - ), - label_encoder: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "gliner-large-label-encoder.gguf", - ), - label_encoder_config: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "large-label-encoder-config.json", - ), - label_encoder_tokenizer: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "large-label-encoder-tokenizer.json", - ), - tokenizer: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "large-text-tokenizer.json", - ), - config: Self::huggingface_or_cached( - "Demonthos/gliner-gguf", - "main", - "large-text-gliner-config.json", - ), - } - } - - /// GLiNER bi-encoder v2.0 Small variant (108M parameters). + /// GLiNER bi-encoder v2.0 Small variant (108M parameters, Q4_K). /// /// Good balance of speed and accuracy, using: /// - Text encoder: ettin-encoder-68m /// - Label encoder: all-MiniLM-L12-v2 pub fn small() -> Self { Self { - model: FileSource::huggingface( - "knowledgator/gliner-bi-small-v2.0-gguf".to_string(), - "main".to_string(), - "gliner-bi-small-v2.0-Q8_0.gguf".to_string(), - ), - label_encoder: FileSource::huggingface( - "knowledgator/gliner-bi-small-v2.0-gguf".to_string(), - "main".to_string(), - "label-encoder-Q8_0.gguf".to_string(), - ), + model: Self::demonthos_gguf("gliner-bi-small-v2.0-Q4_K.gguf"), + label_encoder: Self::demonthos_gguf("gliner-bi-small-v2.0-Q4_K-label-encoder.gguf"), label_encoder_config: FileSource::huggingface( "sentence-transformers/all-MiniLM-L12-v2".to_string(), "main".to_string(), @@ -312,23 +152,15 @@ impl GlinerSource { } } - /// GLiNER bi-encoder v2.0 Base variant (194M parameters). + /// GLiNER bi-encoder v2.0 Base variant (194M parameters, Q4_K). /// /// Default variant with good accuracy, using: /// - Text encoder: ettin-encoder-150m /// - Label encoder: bge-small-en-v1.5 pub fn base() -> Self { Self { - model: FileSource::huggingface( - "knowledgator/gliner-bi-base-v2.0-gguf".to_string(), - "main".to_string(), - "gliner-bi-base-v2.0-Q8_0.gguf".to_string(), - ), - label_encoder: FileSource::huggingface( - "knowledgator/gliner-bi-base-v2.0-gguf".to_string(), - "main".to_string(), - "label-encoder-Q8_0.gguf".to_string(), - ), + model: Self::demonthos_gguf("gliner-bi-base-v2.0-Q4_K.gguf"), + label_encoder: Self::demonthos_gguf("gliner-bi-base-v2.0-Q4_K-label-encoder.gguf"), label_encoder_config: FileSource::huggingface( "BAAI/bge-small-en-v1.5".to_string(), "main".to_string(), @@ -352,23 +184,15 @@ impl GlinerSource { } } - /// GLiNER bi-encoder v2.0 Large variant (530M parameters). + /// GLiNER bi-encoder v2.0 Large variant (530M parameters, Q4_K). /// /// Highest accuracy variant, using: /// - Text encoder: ettin-encoder-400m /// - Label encoder: bge-base-en-v1.5 pub fn large() -> Self { Self { - model: FileSource::huggingface( - "knowledgator/gliner-bi-large-v2.0-gguf".to_string(), - "main".to_string(), - "gliner-bi-large-v2.0-Q8_0.gguf".to_string(), - ), - label_encoder: FileSource::huggingface( - "knowledgator/gliner-bi-large-v2.0-gguf".to_string(), - "main".to_string(), - "label-encoder-Q8_0.gguf".to_string(), - ), + model: Self::demonthos_gguf("gliner-bi-large-v2.0-Q4_K.gguf"), + label_encoder: Self::demonthos_gguf("gliner-bi-large-v2.0-Q4_K-label-encoder.gguf"), label_encoder_config: FileSource::huggingface( "BAAI/bge-base-en-v1.5".to_string(), "main".to_string(), From f1d419eb7a1d1431b49a7834ed58f7faf57e0e34 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Tue, 14 Apr 2026 19:07:45 -0500 Subject: [PATCH 14/34] chunking --- Cargo.lock | 9 + Cargo.toml | 2 + demos/rgliner-web/src/main.rs | 29 +- fusor-ml/core/src/matmul/sgemm.rs | 11 +- interfaces/kalosm-chunking/Cargo.toml | 12 + interfaces/kalosm-chunking/src/lib.rs | 154 +++++ .../src}/sentence/assets/segment.srx | 0 .../kalosm-chunking/src/sentence/mod.rs | 78 +++ interfaces/kalosm-language/Cargo.toml | 1 + .../src/search/preprocessing/chunking.rs | 197 +----- .../src/search/preprocessing/mod.rs | 3 +- .../src/search/preprocessing/sentence/mod.rs | 76 +-- models/rgliner/src/lib.rs | 68 ++ models/rgliner/src/raw/bilstm.rs | 81 ++- models/rgliner/src/raw/span_layer.rs | 74 ++- models/rgliner/src/relex.rs | 595 ++++++++++++++++++ models/rgliner/src/tokenization.rs | 53 +- 17 files changed, 1125 insertions(+), 318 deletions(-) create mode 100644 interfaces/kalosm-chunking/Cargo.toml create mode 100644 interfaces/kalosm-chunking/src/lib.rs rename interfaces/{kalosm-language/src/search/preprocessing => kalosm-chunking/src}/sentence/assets/segment.srx (100%) create mode 100644 interfaces/kalosm-chunking/src/sentence/mod.rs diff --git a/Cargo.lock b/Cargo.lock index 12b60aafa..6e4b0c281 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5623,6 +5623,14 @@ dependencies = [ "tracing-subscriber 0.2.25", ] +[[package]] +name = "kalosm-chunking" +version = "0.4.0" +dependencies = [ + "srx", + "whatlang", +] + [[package]] name = "kalosm-common" version = "0.4.0" @@ -5662,6 +5670,7 @@ dependencies = [ "heed", "image 0.24.9", "kalosm", + "kalosm-chunking", "kalosm-language-model", "kalosm-llama", "kalosm-sample", diff --git a/Cargo.toml b/Cargo.toml index 294a3863e..28252ec67 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -30,6 +30,7 @@ members = [ "interfaces/kalosm-learning-macro", "interfaces/kalosm-parse-macro", "interfaces/kalosm-common", + "interfaces/kalosm-chunking", "interfaces/kalosm-model-types", "fusor-ml/core", "fusor-ml/gguf", @@ -48,6 +49,7 @@ kalosm = { path = "./interfaces/kalosm", version = "0.4.0" } kalosm-sample = { path = "./interfaces/kalosm-sample", version = "0.4.0" } kalosm-parse-macro = { path = "./interfaces/kalosm-parse-macro", version = "0.4.0" } kalosm-common = { path = "./interfaces/kalosm-common", version = "0.4.0" } +kalosm-chunking = { path = "./interfaces/kalosm-chunking", version = "0.4.0" } kalosm-model-types = { path = "./interfaces/kalosm-model-types", version = "0.4.0" } kalosm-language-model = { path = "./interfaces/language-model", version = "0.4.0" } kalosm-streams = { path = "./interfaces/kalosm-streams", version = "0.4.0" } diff --git a/demos/rgliner-web/src/main.rs b/demos/rgliner-web/src/main.rs index 4d09cd942..fffc0eb0f 100644 --- a/demos/rgliner-web/src/main.rs +++ b/demos/rgliner-web/src/main.rs @@ -5,6 +5,8 @@ use rgliner::{ DecodingMode, Entity, Gliner, GlinerSource, }; +const TOKEN_BUDGET: usize = 128; + fn main() { #[cfg(target_arch = "wasm32")] { @@ -376,7 +378,7 @@ async fn build_model( let inner = Gliner::builder() .with_source(source) .with_decoding_mode(DecodingMode::Flat) - .with_threshold(0.3) + .with_threshold(0.05) .build_with_loading_handler(|_| {}) .await .map_err(|e| format!("{e}"))?; @@ -420,7 +422,7 @@ async fn run_extraction( match (mode, model) { (Mode::Ner, LoadedModel::Ner { inner, .. }) => { let entities = inner - .extract(text, &ent_refs) + .extract_auto(text, &ent_refs, Some(TOKEN_BUDGET)) .await .map_err(|e| format!("{e}"))?; Ok(ExtractionOutput::Ner(entities)) @@ -431,7 +433,7 @@ async fn run_extraction( return Err("Add at least one relation label.".to_string()); } let (entities, relations) = inner - .extract(text, &ent_refs, &rel_refs) + .extract_auto(text, &ent_refs, &rel_refs, Some(TOKEN_BUDGET)) .await .map_err(|e| format!("{e}"))?; Ok(ExtractionOutput::Relex(RelexResult { @@ -510,24 +512,23 @@ fn highlighted_text(text: String, entities: Vec) -> Element { let mut segments: Vec<(bool, String, Option)> = Vec::new(); let mut cursor = 0usize; + let len = text.len(); for ent in sorted.iter() { - if ent.start_char < cursor { + let start = ent.start_char.min(len); + let end = ent.end_char.min(len); + if start < cursor || end <= start { continue; } - if ent.start_char > cursor { - segments.push((false, text[cursor..ent.start_char].to_string(), None)); + if !text.is_char_boundary(start) || !text.is_char_boundary(end) { + continue; } - let end = ent.end_char.min(text.len()); - if end > ent.start_char { - segments.push(( - true, - text[ent.start_char..end].to_string(), - Some(ent.label.clone()), - )); + if start > cursor { + segments.push((false, text[cursor..start].to_string(), None)); } + segments.push((true, text[start..end].to_string(), Some(ent.label.clone()))); cursor = end; } - if cursor < text.len() { + if cursor < len { segments.push((false, text[cursor..].to_string(), None)); } diff --git a/fusor-ml/core/src/matmul/sgemm.rs b/fusor-ml/core/src/matmul/sgemm.rs index 3555287c2..a4e6e707a 100644 --- a/fusor-ml/core/src/matmul/sgemm.rs +++ b/fusor-ml/core/src/matmul/sgemm.rs @@ -65,7 +65,15 @@ pub(super) fn build_kernel( let block_k_size = parameters.block_k_size; let thread_m_size = parameters.thread_m_size; let thread_n_size = parameters.thread_n_size; - let double_buffer = parameters.double_buffer; + // WebGPU guarantees only 16KB of workgroup storage on many adapters and caps it + // at 32KB. Disable double-buffering when it would push us past that budget. + const PADDING: u32 = 1; + const MAX_WORKGROUP_STORAGE_BYTES: u32 = 32768; + let dtype_bytes = matmul.matmul_dtype().element_size() as u32; + let single_buffer_bytes = + (block_m_size + block_n_size) * (block_k_size + PADDING) * dtype_bytes; + let double_buffer = + parameters.double_buffer && single_buffer_bytes * 2 <= MAX_WORKGROUP_STORAGE_BYTES; let threads_per_workgroup: u32 = (block_m_size * block_n_size) / (thread_m_size * thread_n_size); let threads_per_k_a: u32 = threads_per_workgroup / block_k_size; @@ -115,7 +123,6 @@ pub(super) fn build_kernel( let output = kernel.add_tensor_input(matmul.rank(), true, matmul.post_element_wise.out_datatype()); - const PADDING: u32 = 1; // Add padding for bank conflict avoidance let cache_dtype = matmul.matmul_dtype(); let cache_a_size = if double_buffer { 2 } else { 1 } * block_m_size * (block_k_size + PADDING); let cache_a = kernel.add_global_array( diff --git a/interfaces/kalosm-chunking/Cargo.toml b/interfaces/kalosm-chunking/Cargo.toml new file mode 100644 index 000000000..1b6d253c0 --- /dev/null +++ b/interfaces/kalosm-chunking/Cargo.toml @@ -0,0 +1,12 @@ +[package] +name = "kalosm-chunking" +version = "0.4.0" +edition = "2021" +description = "Text chunking primitives (sentence, paragraph, word) with language-aware sentence splitting" +license = "MIT/Apache-2.0" +repository = "https://github.com/floneum/floneum" +authors = ["Evan Almloff "] + +[dependencies] +srx = { version = "0.1.4", features = ["from_xml"] } +whatlang = "0.16.3" diff --git a/interfaces/kalosm-chunking/src/lib.rs b/interfaces/kalosm-chunking/src/lib.rs new file mode 100644 index 000000000..161f280d8 --- /dev/null +++ b/interfaces/kalosm-chunking/src/lib.rs @@ -0,0 +1,154 @@ +//! Text chunking primitives used across kalosm crates. +//! +//! This crate provides [`ChunkStrategy`] (paragraph / sentence / word sliding windows) and +//! [`SentenceChunker`] (SRX-based language-aware sentence splitting). It intentionally has +//! only two dependencies (`srx`, `whatlang`) so it can be used from wasm targets and from +//! crates that don't want the full `kalosm-language` stack. + +use std::ops::Range; + +pub mod sentence; + +pub use sentence::{DefaultSentenceChunker, SentenceChunker}; + +/// A strategy for chunking a document into smaller pieces. +#[derive(Debug, Clone, Copy, PartialEq)] +pub enum ChunkStrategy { + /// Split the document into paragraphs. + Paragraph { + /// The number of paragraphs to include in each chunk. + paragraph_count: usize, + /// The number of paragraphs to overlap between chunks. + overlap: usize, + }, + /// Split the document into sentences. + Sentence { + /// The number of sentences to include in each chunk. + sentence_count: usize, + /// The number of sentences to overlap between chunks. + overlap: usize, + }, + /// Split the document into words. + Words { + /// The number of words to include in each chunk. + word_count: usize, + /// The number of words to overlap between chunks. + overlap: usize, + }, +} + +impl ChunkStrategy { + /// Chunk a string into smaller ranges. + pub fn chunk_str(&self, string: &str) -> Vec> { + match self { + Self::Paragraph { + paragraph_count, + overlap, + } => { + let mut chunks = Vec::new(); + let mut start = 0; + let mut newline_indexes = Vec::new(); + for (i, c) in string.char_indices() { + if c == '\n' { + newline_indexes.push(i + 1); + if newline_indexes.len() >= *paragraph_count { + if !string[start..i].trim().is_empty() { + chunks.push(start..i); + } + for _ in 0..(newline_indexes.len() - *overlap) { + start = newline_indexes.remove(0); + } + } + } + } + + if !string[start..].trim().is_empty() { + chunks.push(start..string.len()); + } + + chunks + } + Self::Sentence { + sentence_count, + overlap, + } => { + let mut chunks = Vec::new(); + + let splits = SentenceChunker::default().split_sentences(string); + + for window in splits + .windows(*sentence_count) + .step_by(*sentence_count - *overlap) + { + if window.len() < *sentence_count { + break; + } + let start = window.first().unwrap().start; + let end = window.last().unwrap().end; + chunks.push(start..end); + } + + chunks + } + Self::Words { + word_count, + overlap, + } => { + let mut chunks = Vec::new(); + let mut start = 0; + let mut word_start_indexes = Vec::new(); + for (i, c) in string.char_indices() { + if c == ' ' { + word_start_indexes.push(i + 1); + if word_start_indexes.len() >= *word_count { + if !string[start..i].trim().is_empty() { + chunks.push(start..i); + } + for _ in 0..(word_start_indexes.len() - *overlap) { + start = word_start_indexes.remove(0); + } + } + } + } + + if !string[start..].trim().is_empty() { + chunks.push(start..string.len()); + } + + chunks + } + } + } +} + +impl Default for ChunkStrategy { + fn default() -> Self { + Self::Paragraph { + paragraph_count: 3, + overlap: 1, + } + } +} + +#[test] +fn test_chunking() { + let string = "The quick brown fox jumps over the lazy dog."; + let chunks = ChunkStrategy::Words { + word_count: 3, + overlap: 1, + }; + let chunks = chunks.chunk_str(string); + assert_eq!(chunks.len(), 4); + assert_eq!(string[chunks[0].clone()].trim(), "The quick brown"); + assert_eq!(string[chunks[1].clone()].trim(), "brown fox jumps"); + assert_eq!(string[chunks[2].clone()].trim(), "jumps over the"); + assert_eq!(string[chunks[3].clone()].trim(), "the lazy dog."); + + let chunks = ChunkStrategy::Sentence { + sentence_count: 2, + overlap: 1, + }; + let string = "first sentence. second sentence. third sentence. fourth sentence."; + let chunks = chunks.chunk_str(string); + assert_eq!(chunks.len(), 3); +} diff --git a/interfaces/kalosm-language/src/search/preprocessing/sentence/assets/segment.srx b/interfaces/kalosm-chunking/src/sentence/assets/segment.srx similarity index 100% rename from interfaces/kalosm-language/src/search/preprocessing/sentence/assets/segment.srx rename to interfaces/kalosm-chunking/src/sentence/assets/segment.srx diff --git a/interfaces/kalosm-chunking/src/sentence/mod.rs b/interfaces/kalosm-chunking/src/sentence/mod.rs new file mode 100644 index 000000000..c2a9f09c9 --- /dev/null +++ b/interfaces/kalosm-chunking/src/sentence/mod.rs @@ -0,0 +1,78 @@ +use srx::SRX; +use std::cell::OnceCell; +use std::rc::Rc; +use std::str::FromStr; + +/// The default sentence chunker. Unlike [`SentenceChunker`], this is Send + Sync. +#[derive(Debug, Clone, Copy, Default)] +pub struct DefaultSentenceChunker; + +impl DefaultSentenceChunker { + /// Split a string into sentence byte ranges. + pub fn split_sentences(&self, string: &str) -> Vec> { + SentenceChunker::default().split_sentences(string) + } +} + +/// A sentence splitter backed by [SRX](https://www.unicode.org/uli/pas/srx/srx20.html) rules. +/// +/// Uses the [srx](https://crates.io/crates/srx) crate to parse and apply the rules. +#[derive(Debug, Clone)] +pub struct SentenceChunker { + srx: Rc, +} + +impl SentenceChunker { + /// Create a new sentence chunker from an xml rules string. + pub fn new(rules: &str) -> Self { + Self { + srx: SRX::from_str(rules) + .expect("the rules file is valid") + .into(), + } + } + + /// Create a new sentence chunker from anything that implements [`std::io::Read`] in the srx rules format. + pub fn load(reader: impl std::io::Read) -> Result { + Ok(Self { + srx: SRX::from_reader(reader)?.into(), + }) + } + + /// Split the body of a document into a list of sentence byte ranges. + pub fn split_sentences(&self, string: &str) -> Vec> { + let language = whatlang::detect_lang(string) + .map(|lang_code| lang_code.code()) + .unwrap_or("en"); + + let rules = self.srx.language_rules(language); + rules.split_ranges(string) + } + + /// Access the parsed SRX rules. Primarily useful to external `Chunker` trait implementations. + pub fn srx(&self) -> &SRX { + &self.srx + } +} + +impl Default for SentenceChunker { + fn default() -> Self { + // The rules are expensive to parse (~1 second), so cache them in a thread-local. + thread_local! { + static DEFAULT_RULES: OnceCell> = const { OnceCell::new() }; + } + + let rules = DEFAULT_RULES.with(|default| { + default + .get_or_init(|| { + // LanguageTool ruleset: https://github.com/languagetool-org/languagetool/blob/master/languagetool-core/src/main/resources/org/languagetool/resource/segment.srx + let rules = SRX::from_str(include_str!("./assets/segment.srx")) + .expect("the rules file is valid"); + Rc::new(rules) + }) + .clone() + }); + + Self { srx: rules } + } +} diff --git a/interfaces/kalosm-language/Cargo.toml b/interfaces/kalosm-language/Cargo.toml index 95e94796f..bc8948f22 100644 --- a/interfaces/kalosm-language/Cargo.toml +++ b/interfaces/kalosm-language/Cargo.toml @@ -37,6 +37,7 @@ docx-rs = "0.4.7" lopdf = { version = "0.35.0", features = ["async"] } convert_case = "0.6.0" kalosm-sample = { workspace = true } +kalosm-chunking = { workspace = true } ego-tree = "0.6.2" image = { version = "0.24.7", optional = true } whatlang = "0.16.3" diff --git a/interfaces/kalosm-language/src/search/preprocessing/chunking.rs b/interfaces/kalosm-language/src/search/preprocessing/chunking.rs index 779092abd..429830980 100644 --- a/interfaces/kalosm-language/src/search/preprocessing/chunking.rs +++ b/interfaces/kalosm-language/src/search/preprocessing/chunking.rs @@ -1,203 +1,8 @@ use kalosm_language_model::Embedder; -use std::ops::Range; -use super::{Chunker, SentenceChunker}; +use super::{ChunkStrategy, Chunker}; use crate::{prelude::Document, search::Chunk}; -/// A strategy for chunking a document into smaller pieces. -/// -/// This is used to split a document into smaller pieces to generate embeddings for each piece. -#[derive(Debug, Clone, Copy, PartialEq)] -pub enum ChunkStrategy { - /// Split the document into paragraphs. - Paragraph { - /// The number of paragraphs to include in each chunk. - paragraph_count: usize, - /// The number of paragraphs to overlap between chunks. - overlap: usize, - }, - /// Split the document into sentences. - Sentence { - /// The number of sentences to include in each chunk. - sentence_count: usize, - /// The number of sentences to overlap between chunks. - overlap: usize, - }, - /// Split the document into words. - Words { - /// The number of words to include in each chunk. - word_count: usize, - /// The number of words to overlap between chunks. - overlap: usize, - }, -} - -impl ChunkStrategy { - /// Chunk a string into smaller ranges. - pub fn chunk_str(&self, string: &str) -> Vec> { - match self { - Self::Paragraph { - paragraph_count, - overlap, - } => { - let mut chunks = Vec::new(); - let mut start = 0; - let mut newline_indexes = Vec::new(); - for (i, c) in string.char_indices() { - if c == '\n' { - newline_indexes.push(i + 1); - if newline_indexes.len() >= *paragraph_count { - if !string[start..i].trim().is_empty() { - chunks.push(start..i); - } - for _ in 0..(newline_indexes.len() - *overlap) { - start = newline_indexes.remove(0); - } - } - } - } - - if !string[start..].trim().is_empty() { - chunks.push(start..string.len()); - } - - chunks - } - Self::Sentence { - sentence_count, - overlap, - } => { - let mut chunks = Vec::new(); - - let splits = SentenceChunker::default().split_sentences(string); - - for window in splits - .windows(*sentence_count) - .step_by(*sentence_count - *overlap) - { - if window.len() < *sentence_count { - break; - } - let start = window.first().unwrap().start; - let end = window.last().unwrap().end; - chunks.push(start..end); - } - - chunks - } - Self::Words { - word_count, - overlap, - } => { - let mut chunks = Vec::new(); - let mut start = 0; - let mut word_start_indexes = Vec::new(); - for (i, c) in string.char_indices() { - if c == ' ' { - word_start_indexes.push(i + 1); - if word_start_indexes.len() >= *word_count { - if !string[start..i].trim().is_empty() { - chunks.push(start..i); - } - for _ in 0..(word_start_indexes.len() - *overlap) { - start = word_start_indexes.remove(0); - } - } - } - } - - if !string[start..].trim().is_empty() { - chunks.push(start..string.len()); - } - - chunks - } - } - } -} - -impl Default for ChunkStrategy { - fn default() -> Self { - Self::Paragraph { - paragraph_count: 3, - overlap: 1, - } - } -} - -#[test] -fn test_chunking() { - let string = "The quick brown fox jumps over the lazy dog."; - let chunks = ChunkStrategy::Words { - word_count: 3, - overlap: 1, - }; - let chunks = chunks.chunk_str(string); - assert_eq!(chunks.len(), 4); - assert_eq!(string[chunks[0].clone()].trim(), "The quick brown"); - assert_eq!(string[chunks[1].clone()].trim(), "brown fox jumps"); - assert_eq!(string[chunks[2].clone()].trim(), "jumps over the"); - assert_eq!(string[chunks[3].clone()].trim(), "the lazy dog."); - - let chunks = ChunkStrategy::Words { - word_count: 3, - overlap: 2, - }; - let chunks = chunks.chunk_str(string); - assert_eq!(chunks.len(), 7); - assert_eq!(string[chunks[0].clone()].trim(), "The quick brown"); - assert_eq!(string[chunks[1].clone()].trim(), "quick brown fox"); - assert_eq!(string[chunks[2].clone()].trim(), "brown fox jumps"); - assert_eq!(string[chunks[3].clone()].trim(), "fox jumps over"); - assert_eq!(string[chunks[4].clone()].trim(), "jumps over the"); - assert_eq!(string[chunks[5].clone()].trim(), "over the lazy"); - assert_eq!(string[chunks[6].clone()].trim(), "the lazy dog."); - - let chunks = ChunkStrategy::Sentence { - sentence_count: 2, - overlap: 1, - }; - - let string = "first sentence. second sentence. third sentence. fourth sentence."; - - let chunks = chunks.chunk_str(string); - assert_eq!(chunks.len(), 3); - assert_eq!( - string[chunks[0].clone()].trim(), - "first sentence. second sentence." - ); - assert_eq!( - string[chunks[1].clone()].trim(), - "second sentence. third sentence." - ); - assert_eq!( - string[chunks[2].clone()].trim(), - "third sentence. fourth sentence." - ); - - let chunks = ChunkStrategy::Paragraph { - paragraph_count: 3, - overlap: 1, - }; - - let string = "first paragraph\n\nsecond paragraph\n\nthird paragraph\n\nfourth paragraph"; - - let chunks = chunks.chunk_str(string); - assert_eq!(chunks.len(), 3); - assert_eq!( - string[chunks[0].clone()].trim(), - "first paragraph\n\nsecond paragraph" - ); - assert_eq!( - string[chunks[1].clone()].trim(), - "second paragraph\n\nthird paragraph" - ); - assert_eq!( - string[chunks[2].clone()].trim(), - "third paragraph\n\nfourth paragraph" - ); -} - impl Chunker for ChunkStrategy { type Error = E; diff --git a/interfaces/kalosm-language/src/search/preprocessing/mod.rs b/interfaces/kalosm-language/src/search/preprocessing/mod.rs index 53257c947..f03e0f746 100644 --- a/interfaces/kalosm-language/src/search/preprocessing/mod.rs +++ b/interfaces/kalosm-language/src/search/preprocessing/mod.rs @@ -18,12 +18,11 @@ use crate::context::Document; use super::Chunk; +pub use kalosm_chunking::{ChunkStrategy, DefaultSentenceChunker, SentenceChunker}; mod chunking; -pub use chunking::*; mod task; pub use task::*; mod sentence; -pub use sentence::*; mod semantic; pub use semantic::*; mod html; diff --git a/interfaces/kalosm-language/src/search/preprocessing/sentence/mod.rs b/interfaces/kalosm-language/src/search/preprocessing/sentence/mod.rs index fee6de195..95e0c1350 100644 --- a/interfaces/kalosm-language/src/search/preprocessing/sentence/mod.rs +++ b/interfaces/kalosm-language/src/search/preprocessing/sentence/mod.rs @@ -1,25 +1,17 @@ use crate::prelude::{Chunk, Chunker, Document, Embedder}; -use srx::SRX; -use std::cell::OnceCell; use std::future::Future; -use std::rc::Rc; -use std::str::FromStr; -/// The default sentence chunker. Unlike [`SentenceChunker`], this is Send + Sync -#[derive(Debug, Clone, Copy, Default)] -pub struct DefaultSentenceChunker; +use super::{DefaultSentenceChunker, SentenceChunker}; impl Chunker for DefaultSentenceChunker { type Error = E; - /// Chunk a document into embedded snippets. fn chunk( &self, document: &Document, embedder: &E, ) -> impl Future, Self::Error>> + Send { let default = SentenceChunker::default(); - // Split the document into sentences. We first just collect the sentences as strings and byte ranges let mut initial_chunks = Vec::new(); let body = document.body(); let ranges = default.split_sentences(document.body()); @@ -31,78 +23,14 @@ impl Chunker for DefaultSentenceChunker { } } -/// A [`Chunker`] that splits a string into sentences with a given [SRX](https://www.unicode.org/uli/pas/srx/srx20.html) rules. -/// -/// This uses the [srx](https://crates.io/crates/srx) crate to parse and apply the rules. -#[derive(Debug, Clone)] -pub struct SentenceChunker { - srx: Rc, -} - -impl SentenceChunker { - /// Create a new sentence chunker from a xml rules string - pub fn new(rules: &str) -> Self { - Self { - srx: SRX::from_str(rules) - .expect("the rules file is valid") - .into(), - } - } - - /// Create a new sentence chunker from anything that implements [`std::io::Read`](std::io::Read) in the srx rules format - pub fn load(reader: impl std::io::Read) -> Result { - Ok(Self { - srx: SRX::from_reader(reader)?.into(), - }) - } - - /// Split the body of a document into a list of ranges with sentences - pub fn split_sentences(&self, string: &str) -> Vec> { - // Try to autodetect the language of the document - let language = whatlang::detect_lang(string) - .map(|lang_code| lang_code.code()) - .unwrap_or("en"); - - // Then get the language specific rules to split the document into sentences - let rules = self.srx.language_rules(language); - - rules.split_ranges(string) - } -} - -impl Default for SentenceChunker { - fn default() -> Self { - // The rules are expensive to parse (~1 second), so we cache them in a static once cell - thread_local! { - static DEFAULT_RULES: OnceCell> = const { OnceCell::new() }; - } - - let rules = DEFAULT_RULES.with(|default| { - default - .get_or_init(|| { - // Defaults to the language tool ruleset: https://github.com/languagetool-org/languagetool/blob/master/languagetool-core/src/main/resources/org/languagetool/resource/segment.srx - let rules = SRX::from_str(include_str!("./assets/segment.srx")) - .expect("the rules file is valid"); - Rc::new(rules) - }) - .clone() - }); - - Self { srx: rules } - } -} - -/// A strategy for chunking a document into smaller pieces. impl Chunker for SentenceChunker { type Error = E; - /// Chunk a document into embedded snippets. fn chunk( &self, document: &Document, embedder: &E, ) -> impl Future, Self::Error>> + Send { - // Split the document into sentences. We first just collect the sentences as strings and byte ranges let mut initial_chunks = Vec::new(); let body = document.body(); let ranges = self.split_sentences(document.body()); @@ -119,10 +47,8 @@ async fn embed_chunk( initial_chunks: Vec, ranges: Vec>, ) -> Result, E::Error> { - // Next embed them all in one big batch let embeddings = embedder.embed_vec(initial_chunks).await?; - // Now merge the embeddings and ranges into chunks let mut chunks = Vec::new(); for (embedding, chunk) in embeddings.into_iter().zip(ranges) { let chunk = Chunk { diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs index 303b7cea6..e8bcce633 100644 --- a/models/rgliner/src/lib.rs +++ b/models/rgliner/src/lib.rs @@ -89,6 +89,28 @@ pub use error::{GlinerError, GlinerLoadingError}; pub use rbert::raw::{ModernBertConfig, ModernBertModel}; pub use source::GlinerSource; +/// Deduplicate entities appearing in more than one overlapping chunk, keeping the +/// highest-scoring occurrence and sorting by span position. +fn merge_entities(entities: &mut Vec) { + entities.sort_by(|a, b| { + a.start_char + .cmp(&b.start_char) + .then_with(|| a.end_char.cmp(&b.end_char)) + .then_with(|| a.label.cmp(&b.label)) + }); + entities.dedup_by(|b, a| { + if a.start_char == b.start_char && a.end_char == b.end_char && a.label == b.label { + if b.score > a.score { + a.score = b.score; + } + true + } else { + false + } + }); +} + + use fusor::{Device, Tensor, VarBuilder}; use kalosm_common::Cache; use kalosm_model_types::ModelLoadingProgress; @@ -333,6 +355,52 @@ impl Gliner { Ok(results.pop().unwrap_or_default()) } + /// Extract named entities from text, chunking the input first so long documents + /// that would otherwise be truncated by the text encoder's context window still + /// get full coverage. + /// + /// Uses the model's own tokenizer to pack whole words into chunks of at most + /// `token_budget` subtokens, with roughly 15% token overlap between adjacent + /// chunks. Each chunk is scored independently; entity offsets are remapped back + /// into the original text and deduped across overlapping windows (keeping the + /// highest score per span+label). + /// + /// `token_budget` defaults to 128 — empirically the sweet spot for the edge + /// variant's span-scoring quality. Larger budgets approach the context limit + /// but hurt F1; much smaller budgets hurt recall. + pub async fn extract_auto( + &mut self, + text: &str, + labels: &[&str], + token_budget: Option, + ) -> Result, GlinerError> { + let budget = token_budget.unwrap_or(128); + let ranges = crate::tokenization::token_packed_ranges( + &self.tokenizer.tokenizer, + text, + budget, + budget / 7, + )?; + if ranges.len() <= 1 { + return self.extract(text, labels).await; + } + + let chunk_texts: Vec<&str> = ranges.iter().map(|r| &text[r.clone()]).collect(); + let per_chunk = self.extract_batch(&chunk_texts, labels).await?; + + let mut all: Vec = Vec::new(); + for (range, entities) in ranges.iter().zip(per_chunk) { + let offset = range.start; + for mut ent in entities { + ent.start_char += offset; + ent.end_char += offset; + all.push(ent); + } + } + merge_entities(&mut all); + Ok(all) + } + /// Extract named entities using cached labels. /// /// Panics if no labels are cached. diff --git a/models/rgliner/src/raw/bilstm.rs b/models/rgliner/src/raw/bilstm.rs index 398b0442a..2d2f6a24d 100644 --- a/models/rgliner/src/raw/bilstm.rs +++ b/models/rgliner/src/raw/bilstm.rs @@ -76,11 +76,40 @@ impl BiLstm { /// # Returns /// Output tensor [batch, seq_len, 2*hidden_size] pub async fn forward(&self, input: &Tensor<3, f32>) -> Tensor<3, f32> { + let [_batch, seq_len, _input_size] = input.shape(); + let lengths = vec![seq_len; input.shape()[0]]; + self.forward_with_lengths(input, &lengths).await + } + + /// Forward pass through BiLSTM with explicit per-item sequence lengths. + /// + /// Padded timesteps are masked out so shorter sequences in a batch do not + /// corrupt the backward direction state. + pub async fn forward_with_lengths( + &self, + input: &Tensor<3, f32>, + lengths: &[usize], + ) -> Tensor<3, f32> { let [batch, seq_len, _input_size] = input.shape(); + assert_eq!(lengths.len(), batch, "lengths must match batch size"); let device = input.device(); - let fwd_out = run_direction(input, &self.forward, self.hidden_size, &device, false); - let bwd_out = run_direction(input, &self.backward, self.hidden_size, &device, true); + let fwd_out = run_direction( + input, + &self.forward, + self.hidden_size, + &device, + false, + lengths, + ); + let bwd_out = run_direction( + input, + &self.backward, + self.hidden_size, + &device, + true, + lengths, + ); // Concatenate forward and backward along the feature dim -> [batch, seq, 2*hidden] Tensor::cat([fwd_out, bwd_out], 2) @@ -97,6 +126,7 @@ fn run_direction( hidden_size: usize, device: &Device, reverse: bool, + lengths: &[usize], ) -> Tensor<3, f32> { let [batch, seq_len, _] = input.shape(); @@ -107,6 +137,7 @@ fn run_direction( // dim 1 so that a final cat along dim 1 yields [batch, seq_len, hidden]. let mut outputs: Vec> = Vec::with_capacity(seq_len); outputs.resize_with(seq_len, || Tensor::zeros(device, [batch, 1, hidden_size])); + let zero_output: Tensor<3, f32> = Tensor::zeros(device, [batch, 1, hidden_size]); let bias_broadcast: Tensor<2, f32> = dir .bias @@ -145,10 +176,16 @@ fn run_direction( let g_gate = g_raw.tanh(); let o_gate = sigmoid_2d(&o_raw); - c = (f_gate * c + i_gate * g_gate).to_concrete(); - h = (o_gate * c.clone().tanh()).to_concrete(); + let next_c = (f_gate * c.clone() + i_gate * g_gate).to_concrete(); + let next_h = (o_gate * next_c.clone().tanh()).to_concrete(); + + let active_mask_2d = timestep_mask_2d(device, batch, hidden_size, lengths, t); + c = active_mask_2d.where_cond(&next_c, &c).to_concrete(); + h = active_mask_2d.where_cond(&next_h, &h).to_concrete(); - outputs[t] = h.clone().unsqueeze(1).to_concrete(); + let active_mask_3d = timestep_mask_3d(device, batch, hidden_size, lengths, t); + let output_t = h.clone().unsqueeze(1).to_concrete(); + outputs[t] = active_mask_3d.where_cond(&output_t, &zero_output).to_concrete(); } Tensor::cat(outputs, 1) @@ -156,6 +193,40 @@ fn run_direction( .to_concrete() } +fn timestep_mask_2d( + device: &Device, + batch: usize, + hidden_size: usize, + lengths: &[usize], + timestep: usize, +) -> Tensor<2, f32> { + let mask_data: Vec = lengths + .iter() + .map(|&length| if timestep < length { 1.0 } else { 0.0 }) + .collect(); + Tensor::new(device, &mask_data) + .reshape([batch, 1]) + .broadcast_as([batch, hidden_size]) + .to_concrete() +} + +fn timestep_mask_3d( + device: &Device, + batch: usize, + hidden_size: usize, + lengths: &[usize], + timestep: usize, +) -> Tensor<3, f32> { + let mask_data: Vec = lengths + .iter() + .map(|&length| if timestep < length { 1.0 } else { 0.0 }) + .collect(); + Tensor::new(device, &mask_data) + .reshape([batch, 1, 1]) + .broadcast_as([batch, 1, hidden_size]) + .to_concrete() +} + /// sigmoid via `0.5 * (tanh(x / 2) + 1)` — avoids needing scalar-left division /// or a `recip` primitive, and keeps the computation on-device. fn sigmoid_2d(x: &Tensor<2, f32>) -> Tensor<2, f32> { diff --git a/models/rgliner/src/raw/span_layer.rs b/models/rgliner/src/raw/span_layer.rs index 31f9aee8f..82ce20b8e 100644 --- a/models/rgliner/src/raw/span_layer.rs +++ b/models/rgliner/src/raw/span_layer.rs @@ -173,11 +173,36 @@ impl SpanLayer { spans: &[(usize, usize)], device: &Device, ) -> Tensor<2, f32> { + let (batched, counts) = + self.forward_for_spans_batched(word_embeddings, &[spans.to_vec()], device); + let _count = counts.first().copied().unwrap_or(0); + batched + } + + /// Compute span representations for a batch of per-item span lists. + /// + /// Returns: + /// - flattened span embeddings in batch-major order + /// - one count per batch item so the caller can slice the flattened output + pub fn forward_for_spans_batched( + &self, + word_embeddings: &Tensor<3, f32>, + spans_per_batch: &[Vec<(usize, usize)>], + device: &Device, + ) -> (Tensor<2, f32>, Vec) { let [batch_size, num_words, hidden_dim] = word_embeddings.shape(); - assert_eq!(batch_size, 1, "only batch_size=1 supported"); - let num_spans = spans.len(); + assert_eq!( + batch_size, + spans_per_batch.len(), + "spans_per_batch must match batch size" + ); + + let span_counts: Vec = spans_per_batch.iter().map(Vec::len).collect(); + let total_spans: usize = span_counts.iter().sum(); + if total_spans == 0 { + return (Tensor::zeros(device, [1, hidden_dim]), span_counts); + } - // Apply project_start and project_end to the full word embeddings let start_rep = self .start_fc2 .forward(&self.start_fc1.forward(word_embeddings).relu()); @@ -185,32 +210,35 @@ impl SpanLayer { .end_fc2 .forward(&self.end_fc1.forward(word_embeddings).relu()); - // Gather at span positions - let start_rep_2d = start_rep.squeeze(0).to_concrete(); - let end_rep_2d = end_rep.squeeze(0).to_concrete(); - let _ = num_words; - - let start_indices: Vec = spans.iter().map(|(s, _)| *s as u32).collect(); - let end_indices: Vec = spans.iter().map(|(_, e)| *e as u32).collect(); - let start_idx_tensor = Tensor::new(device, &start_indices); - let end_idx_tensor = Tensor::new(device, &end_indices); + let start_rep_flat = start_rep + .to_concrete() + .reshape([batch_size * num_words, hidden_dim]) + .to_concrete(); + let end_rep_flat = end_rep + .to_concrete() + .reshape([batch_size * num_words, hidden_dim]) + .to_concrete(); - let start_gathered = start_rep_2d.index_select(0, &start_idx_tensor); - let end_gathered = end_rep_2d.index_select(0, &end_idx_tensor); + let mut start_offset_indices: Vec = Vec::with_capacity(total_spans); + let mut end_offset_indices: Vec = Vec::with_capacity(total_spans); + for (batch_idx, spans) in spans_per_batch.iter().enumerate() { + let offset = (batch_idx * num_words) as u32; + for &(start, end) in spans { + start_offset_indices.push(start as u32 + offset); + end_offset_indices.push(end as u32 + offset); + } + } - // Concat along last dim: [num_spans, hidden*2] - let start_3d: Tensor<3, f32> = start_gathered.unsqueeze(0).to_concrete(); - let end_3d: Tensor<3, f32> = end_gathered.unsqueeze(0).to_concrete(); - let combined = Tensor::cat([start_3d, end_3d], 2).relu(); + let start_idx_tensor = Tensor::new(device, &start_offset_indices); + let end_idx_tensor = Tensor::new(device, &end_offset_indices); - // Apply out_project: Linear -> ReLU -> Linear + let start_gathered = start_rep_flat.index_select(0, &start_idx_tensor); + let end_gathered = end_rep_flat.index_select(0, &end_idx_tensor); + let combined = Tensor::cat([start_gathered, end_gathered], 1).relu(); let hidden = self.out_fc1.forward(&combined).relu(); let out = self.out_fc2.forward(&hidden); - // [1, num_spans, hidden] -> [num_spans, hidden] - let _ = num_spans; - let _ = hidden_dim; - out.squeeze(0).to_concrete() + (out.to_concrete(), span_counts) } fn gather_span_embeddings( diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index 46ffc2ada..b9b7fe2b4 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -594,6 +594,101 @@ impl GlinerRelEx { Ok((entities, relations)) } + /// Extract entities and relations from text, chunking the input first so long + /// documents that would otherwise be truncated by the encoder's context window + /// still get full coverage. + /// + /// Uses the model's own tokenizer to pack whole words into chunks of at most + /// `token_budget` subtokens with ~15% overlap between adjacent chunks. Each + /// chunk is scored independently; entity and relation byte offsets are remapped + /// back into the original text and deduped across overlapping windows (keeping + /// the highest score per span+label / head+tail+label). + /// + /// `token_budget` defaults to 128. + pub async fn extract_auto( + &self, + text: &str, + entity_labels: &[&str], + relation_labels: &[&str], + token_budget: Option, + ) -> Result<(Vec, Vec), GlinerError> { + let budget = token_budget.unwrap_or(128); + let ranges = crate::tokenization::token_packed_ranges( + self.tokenizer.tokenizer(), + text, + budget, + budget / 7, + )?; + if ranges.len() <= 1 { + return self.extract(text, entity_labels, relation_labels).await; + } + + let shift = |ent: &mut Entity, offset: usize| { + ent.start_char += offset; + ent.end_char += offset; + }; + + let mut all_entities: Vec = Vec::new(); + let mut all_relations: Vec = Vec::new(); + for range in &ranges { + let chunk = &text[range.clone()]; + let (entities, relations) = + self.extract(chunk, entity_labels, relation_labels).await?; + let offset = range.start; + for mut ent in entities { + shift(&mut ent, offset); + all_entities.push(ent); + } + for mut rel in relations { + shift(&mut rel.head, offset); + shift(&mut rel.tail, offset); + all_relations.push(rel); + } + } + + all_entities.sort_by(|a, b| { + a.start_char + .cmp(&b.start_char) + .then_with(|| a.end_char.cmp(&b.end_char)) + .then_with(|| a.label.cmp(&b.label)) + }); + all_entities.dedup_by(|b, a| { + if a.start_char == b.start_char && a.end_char == b.end_char && a.label == b.label { + if b.score > a.score { + a.score = b.score; + } + true + } else { + false + } + }); + + all_relations.sort_by(|a, b| { + a.head + .start_char + .cmp(&b.head.start_char) + .then_with(|| a.tail.start_char.cmp(&b.tail.start_char)) + .then_with(|| a.relation.cmp(&b.relation)) + }); + all_relations.dedup_by(|b, a| { + if a.head.start_char == b.head.start_char + && a.head.end_char == b.head.end_char + && a.tail.start_char == b.tail.start_char + && a.tail.end_char == b.tail.end_char + && a.relation == b.relation + { + if b.score > a.score { + a.score = b.score; + } + true + } else { + false + } + }); + + Ok((all_entities, all_relations)) + } + /// Decode entities for `span_mode = markerV0` (used by the `large` variants). /// /// Enumerates every `(start, end)` pair up to `config.max_width` words, @@ -819,3 +914,503 @@ impl GlinerRelEx { &self.config } } + +#[cfg(test)] +mod tests { + use super::*; + use std::path::PathBuf; + use std::time::{Duration, Instant}; + + const PROFILE_TEXT: &str = "Apple Inc. was founded by Steve Jobs in Cupertino, California. \ +Microsoft was founded by Bill Gates in Albuquerque, New Mexico. \ +Google was founded by Larry Page and Sergey Brin in Menlo Park, California. \ +Amazon was founded by Jeff Bezos in Bellevue, Washington. \ +Meta Platforms was founded by Mark Zuckerberg in Cambridge, Massachusetts."; + const ENTITY_LABELS: &[&str] = &["organization", "person", "location"]; + const RELATION_LABELS: &[&str] = &["founded by", "located in"]; + + #[derive(Debug, Clone)] + struct ExtractProfile { + variant: &'static str, + device: &'static str, + span_mode: SpanMode, + seq_len: usize, + num_words: usize, + entity_labels: usize, + relation_labels: usize, + entity_count: usize, + relation_count: usize, + span_count: usize, + candidate_pairs: usize, + cold_total: Duration, + warm_total: Duration, + tokenize_cpu: Duration, + input_prep_cpu: Duration, + entity_span_prep_cpu: Duration, + entity_sync: Duration, + entity_decode_cpu: Duration, + relation_span_sync: Duration, + relation_pair_pack_cpu: Duration, + relation_score_sync: Duration, + relation_decode_cpu: Duration, + } + + impl ExtractProfile { + fn print(&self) { + println!( + "PROFILE variant={} device={} span_mode={:?} seq_len={} words={} ent_labels={} rel_labels={} entities={} relations={} spans={} pairs={}", + self.variant, + self.device, + self.span_mode, + self.seq_len, + self.num_words, + self.entity_labels, + self.relation_labels, + self.entity_count, + self.relation_count, + self.span_count, + self.candidate_pairs + ); + println!( + " cold_total_ms={:.2} warm_total_ms={:.2}", + self.cold_total.as_secs_f64() * 1000.0, + self.warm_total.as_secs_f64() * 1000.0 + ); + println!( + " tokenize_cpu_ms={:.2} input_prep_cpu_ms={:.2} entity_span_prep_cpu_ms={:.2}", + self.tokenize_cpu.as_secs_f64() * 1000.0, + self.input_prep_cpu.as_secs_f64() * 1000.0, + self.entity_span_prep_cpu.as_secs_f64() * 1000.0 + ); + println!( + " entity_sync_ms={:.2} entity_decode_cpu_ms={:.2}", + self.entity_sync.as_secs_f64() * 1000.0, + self.entity_decode_cpu.as_secs_f64() * 1000.0 + ); + println!( + " relation_span_sync_ms={:.2} relation_pair_pack_cpu_ms={:.2}", + self.relation_span_sync.as_secs_f64() * 1000.0, + self.relation_pair_pack_cpu.as_secs_f64() * 1000.0 + ); + println!( + " relation_score_sync_ms={:.2} relation_decode_cpu_ms={:.2}", + self.relation_score_sync.as_secs_f64() * 1000.0, + self.relation_decode_cpu.as_secs_f64() * 1000.0 + ); + } + } + + fn weights_path(file_name: &str) -> PathBuf { + PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("weights") + .join(file_name) + } + + async fn load_local_relex( + model_path: PathBuf, + device: Device, + ) -> Result { + GlinerRelEx::builder() + .with_source(GlinerRelExSource::local(model_path)) + .with_device(device) + .build_with_loading_handler(|_| {}) + .await + } + + async fn profile_extract( + model: &GlinerRelEx, + variant: &'static str, + cold_total: Duration, + ) -> Result { + let total_start = Instant::now(); + + let tokenize_start = Instant::now(); + let tokenized = model + .tokenizer + .tokenize(PROFILE_TEXT, ENTITY_LABELS, RELATION_LABELS)?; + let tokenize_cpu = tokenize_start.elapsed(); + + let seq_len = tokenized.token_ids.len(); + let num_words = tokenized.num_words; + + let input_start = Instant::now(); + let token_ids = Tensor::new(&model.device, &tokenized.token_ids); + let token_ids: Tensor<2, u32> = token_ids.unsqueeze(0).to_concrete(); + + let attention_mask = Tensor::new(&model.device, &tokenized.attention_mask); + let attention_mask: Tensor<2, u32> = attention_mask.unsqueeze(0).to_concrete(); + let input_prep_cpu = input_start.elapsed(); + + let entity_compute_start = Instant::now(); + let encoder_output = model.encoder.forward(&token_ids, Some(&attention_mask)); + let word_encoder_embs = + model.gather_at_positions(&encoder_output, &tokenized.text_positions); + let lstm_output = model.bilstm.forward(&word_encoder_embs).await; + let ent_embs_raw = model.gather_at_positions(&encoder_output, &tokenized.ent_positions); + let ent_embs = model.prompt_rep_layer.forward_3d(&ent_embs_raw); + let rel_embs = model.gather_at_positions(&encoder_output, &tokenized.rel_positions); + let text_embs = lstm_output.clone(); + let ent_embs_2d: Tensor<2, f32> = ent_embs.squeeze(0).to_concrete(); + + let (entities, span_count, entity_span_prep_cpu, entity_sync, entity_decode_cpu) = + match model.span_mode { + SpanMode::TokenLevel => { + let scorer = model + .scorer + .as_ref() + .expect("token_level requires scorer"); + let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs_2d); + let (entities, entity_sync, entity_decode_cpu) = + profile_decode_entities_from_tokens( + model, + &token_scores, + ENTITY_LABELS, + &tokenized.word_offsets, + PROFILE_TEXT, + entity_compute_start, + ) + .await?; + ( + entities, + 0, + Duration::ZERO, + entity_sync, + entity_decode_cpu, + ) + } + SpanMode::MarkerV0 => { + profile_decode_entities_marker_v0( + model, + &text_embs, + &ent_embs_2d, + ENTITY_LABELS, + &tokenized.word_offsets, + tokenized.num_words, + PROFILE_TEXT, + entity_compute_start, + ) + .await? + } + }; + + let mut relation_span_sync = Duration::ZERO; + let mut relation_pair_pack_cpu = Duration::ZERO; + let mut relation_score_sync = Duration::ZERO; + let mut relation_decode_cpu = Duration::ZERO; + let mut candidate_pairs = 0usize; + let mut relation_count = 0usize; + + if entities.len() >= 2 && !RELATION_LABELS.is_empty() { + let relation_span_start = Instant::now(); + let entity_spans: Vec<(usize, usize)> = entities + .iter() + .map(|e| (e.start_word, e.end_word)) + .collect(); + let span_reps = model + .span_layer + .forward_for_spans(&text_embs, &entity_spans, &model.device); + let span_reps_data = span_reps.clone().as_slice().await?; + relation_span_sync = relation_span_start.elapsed(); + + let pair_pack_start = Instant::now(); + let num_entities = entities.len(); + let hidden_size = model.config.hidden_size; + let span_reps_slice = span_reps_data.as_slice(); + + let mut pairs: Vec<(usize, usize)> = Vec::new(); + for head in 0..num_entities { + for tail in 0..num_entities { + if head != tail { + pairs.push((head, tail)); + } + } + } + candidate_pairs = pairs.len(); + + let mut head_embs = Vec::with_capacity(candidate_pairs * hidden_size); + let mut tail_embs = Vec::with_capacity(candidate_pairs * hidden_size); + for &(head_idx, tail_idx) in &pairs { + let h_start = head_idx * hidden_size; + let t_start = tail_idx * hidden_size; + head_embs.extend_from_slice(&span_reps_slice[h_start..h_start + hidden_size]); + tail_embs.extend_from_slice(&span_reps_slice[t_start..t_start + hidden_size]); + } + + let head_tensor = Tensor::new(&model.device, &head_embs) + .reshape([candidate_pairs, hidden_size]) + .to_concrete(); + let tail_tensor = Tensor::new(&model.device, &tail_embs) + .reshape([candidate_pairs, hidden_size]) + .to_concrete(); + relation_pair_pack_cpu = pair_pack_start.elapsed(); + + let relation_score_start = Instant::now(); + let pair_embs = model.pair_projector.forward(&head_tensor, &tail_tensor); + let rel_embs_squeezed: Tensor<2, f32> = rel_embs.squeeze(0).to_concrete(); + let rel_scores = pair_embs.mat_mul(&rel_embs_squeezed.transpose(0, 1)); + let rel_scores_slice = rel_scores.clone().as_slice().await?; + relation_score_sync = relation_score_start.elapsed(); + + let relation_decode_start = Instant::now(); + let n_rels = RELATION_LABELS.len(); + let threshold = model.config.relation_threshold; + let mut relations = Vec::new(); + for (pair_idx, &(head_idx, tail_idx)) in pairs.iter().enumerate() { + let base = pair_idx * n_rels; + for rel_idx in 0..n_rels { + let raw = rel_scores_slice.as_slice()[base + rel_idx]; + let prob = 1.0 / (1.0 + (-raw).exp()); + if prob > threshold { + relations.push(Relation { + head: entities[head_idx].clone(), + tail: entities[tail_idx].clone(), + relation: RELATION_LABELS[rel_idx].to_string(), + score: prob, + }); + } + } + } + relations.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); + relation_count = relations.len(); + relation_decode_cpu = relation_decode_start.elapsed(); + } + + let warm_total = total_start.elapsed(); + Ok(ExtractProfile { + variant, + device: if model.device.is_gpu() { "gpu" } else { "cpu" }, + span_mode: model.span_mode, + seq_len, + num_words, + entity_labels: ENTITY_LABELS.len(), + relation_labels: RELATION_LABELS.len(), + entity_count: entities.len(), + relation_count, + span_count, + candidate_pairs, + cold_total, + warm_total, + tokenize_cpu, + input_prep_cpu, + entity_span_prep_cpu, + entity_sync, + entity_decode_cpu, + relation_span_sync, + relation_pair_pack_cpu, + relation_score_sync, + relation_decode_cpu, + }) + } + + async fn profile_decode_entities_marker_v0( + model: &GlinerRelEx, + text_embs: &Tensor<3, f32>, + ent_embs_2d: &Tensor<2, f32>, + entity_labels: &[&str], + word_offsets: &[(usize, usize)], + num_words: usize, + text: &str, + entity_compute_start: Instant, + ) -> Result<(Vec, usize, Duration, Duration, Duration), GlinerError> { + let threshold = model.config.entity_threshold; + let max_width = model.config.max_width; + let n_labels = entity_labels.len(); + + if num_words == 0 || n_labels == 0 { + return Ok((Vec::new(), 0, Duration::ZERO, Duration::ZERO, Duration::ZERO)); + } + + let span_prep_start = Instant::now(); + let mut spans: Vec<(usize, usize)> = Vec::new(); + for start in 0..num_words { + for width in 1..=max_width.min(num_words - start) { + spans.push((start, start + width - 1)); + } + } + let entity_span_prep_cpu = span_prep_start.elapsed(); + + let span_reps = model + .span_layer + .forward_for_spans(text_embs, &spans, &model.device); + let label_rep_t: Tensor<2, f32> = ent_embs_2d.transpose(0, 1).to_concrete(); + let logits = span_reps.mat_mul(&label_rep_t); + let logits_data = logits.clone().as_slice().await?; + let entity_sync = entity_compute_start.elapsed(); + + let decode_start = Instant::now(); + let logits_slice = logits_data.as_slice(); + let mut candidates: Vec<(usize, usize, usize, f32)> = Vec::new(); + for (span_idx, &(s, e)) in spans.iter().enumerate() { + for l in 0..n_labels { + let raw = logits_slice[span_idx * n_labels + l]; + let prob = 1.0 / (1.0 + (-raw).exp()); + if prob >= threshold { + candidates.push((s, e, l, prob)); + } + } + } + + candidates.sort_by(|a, b| b.3.partial_cmp(&a.3).unwrap_or(std::cmp::Ordering::Equal)); + + let mut taken: Vec<(usize, usize)> = Vec::new(); + let mut entities = Vec::new(); + for (s, e, l, score) in candidates { + let overlap = taken.iter().any(|&(a, b)| !(e < a || s > b)); + if overlap { + continue; + } + taken.push((s, e)); + if s < word_offsets.len() && e < word_offsets.len() { + let (start_char, _) = word_offsets[s]; + let (_, end_char) = word_offsets[e]; + entities.push(Entity { + text: text[start_char..end_char].to_string(), + label: entity_labels[l].to_string(), + start_char, + end_char, + start_word: s, + end_word: e, + score, + }); + } + } + entities.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); + let entity_decode_cpu = decode_start.elapsed(); + + Ok(( + entities, + spans.len(), + entity_span_prep_cpu, + entity_sync, + entity_decode_cpu, + )) + } + + async fn profile_decode_entities_from_tokens( + model: &GlinerRelEx, + token_scores: &Tensor<4, f32>, + entity_labels: &[&str], + word_offsets: &[(usize, usize)], + text: &str, + entity_compute_start: Instant, + ) -> Result<(Vec, Duration, Duration), GlinerError> { + let [_batch_size, num_tokens, num_labels, num_channels] = token_scores.shape(); + assert_eq!(num_channels, 3, "expected [start, end, inside]"); + let scores_data = token_scores.clone().as_slice().await?; + let entity_sync = entity_compute_start.elapsed(); + + let decode_start = Instant::now(); + let scores = scores_data.as_slice(); + let threshold = model.config.entity_threshold; + let mut candidates: Vec<(usize, usize, usize, f32)> = Vec::new(); + + let score_at = |tok: usize, lab: usize, ch: usize| -> f32 { + scores[tok * num_labels * 3 + lab * 3 + ch] + }; + + for label_idx in 0..num_labels { + for start_tok in 0..num_tokens { + let start_score = score_at(start_tok, label_idx, 0); + if start_score < threshold { + continue; + } + + for end_tok in start_tok..num_tokens { + let end_score = score_at(end_tok, label_idx, 1); + if end_score < threshold { + continue; + } + + let mut min_score = start_score.min(end_score); + let mut valid = true; + for t in start_tok..=end_tok { + let inside = score_at(t, label_idx, 2); + if inside < threshold { + valid = false; + break; + } + if inside < min_score { + min_score = inside; + } + } + if valid { + candidates.push((start_tok, end_tok, label_idx, min_score)); + } + } + } + } + + candidates.sort_by(|a, b| b.3.partial_cmp(&a.3).unwrap_or(std::cmp::Ordering::Equal)); + + let mut taken: Vec<(usize, usize)> = Vec::new(); + let mut entities = Vec::new(); + for (start_tok, end_tok, label_idx, score) in candidates { + let overlap = taken.iter().any(|&(a, b)| !(end_tok < a || start_tok > b)); + if overlap { + continue; + } + taken.push((start_tok, end_tok)); + + if start_tok < word_offsets.len() && end_tok < word_offsets.len() { + let (start_char, _) = word_offsets[start_tok]; + let (_, end_char) = word_offsets[end_tok]; + entities.push(Entity { + text: text[start_char..end_char].to_string(), + label: entity_labels[label_idx].to_string(), + start_char, + end_char, + start_word: start_tok, + end_word: end_tok, + score, + }); + } + } + + entities.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); + let entity_decode_cpu = decode_start.elapsed(); + + Ok((entities, entity_sync, entity_decode_cpu)) + } + + #[tokio::test] + #[ignore = "diagnostic profile for rel-ex forward-pass stage breakdowns"] + async fn profile_relex_variants() -> Result<(), Box> { + let device = match std::panic::catch_unwind(Device::gpu_blocking) { + Ok(Ok(device)) => device, + _ => Device::cpu(), + }; + + let variants = [ + ("multi", "gliner-relex-multi-v1.0-Q4_K.gguf"), + ("base", "gliner-relex-base-v1.0-Q4_K.gguf"), + ("large", "gliner-relex-large-v1.0-Q4_K.gguf"), + ]; + + for (variant, file_name) in variants { + let model = load_local_relex(weights_path(file_name), device.clone()).await?; + + let cold_start = Instant::now(); + let _ = model + .extract(PROFILE_TEXT, ENTITY_LABELS, RELATION_LABELS) + .await?; + let cold_total = cold_start.elapsed(); + + let profile = profile_extract(&model, variant, cold_total).await?; + profile.print(); + } + + Ok(()) + } +} diff --git a/models/rgliner/src/tokenization.rs b/models/rgliner/src/tokenization.rs index 10f5ec3ae..6885e46f0 100644 --- a/models/rgliner/src/tokenization.rs +++ b/models/rgliner/src/tokenization.rs @@ -22,7 +22,7 @@ pub struct TokenizedText { /// Word-level tokenizer wrapper. pub struct WordTokenizer { - tokenizer: Tokenizer, + pub(crate) tokenizer: Tokenizer, add_special_tokens: bool, } @@ -84,6 +84,57 @@ impl WordTokenizer { } } +/// Pack `text` into token-budgeted byte ranges using the supplied tokenizer. +/// +/// Splits the input on GLiNER-style word boundaries, encodes each word to count +/// subtokens, and greedily fills windows of at most `token_budget` subtokens with +/// `overlap_tokens` of trailing-token overlap between adjacent windows. +pub(crate) fn token_packed_ranges( + tokenizer: &Tokenizer, + text: &str, + token_budget: usize, + overlap_tokens: usize, +) -> Result>, GlinerError> { + let words = split_words(text); + if words.is_empty() { + return Ok(Vec::new()); + } + + let mut word_token_counts = Vec::with_capacity(words.len()); + for (w, _) in &words { + let enc = tokenizer + .encode(w.clone(), false) + .map_err(GlinerError::Tokenizer)?; + word_token_counts.push(enc.get_ids().len().max(1)); + } + + let mut ranges = Vec::new(); + let mut word = 0usize; + while word < words.len() { + let mut end_word = word; + let mut tokens = 0usize; + while end_word < words.len() && tokens + word_token_counts[end_word] <= token_budget { + tokens += word_token_counts[end_word]; + end_word += 1; + } + if end_word == word { + end_word = word + 1; + } + ranges.push(words[word].1 .0..words[end_word - 1].1 .1); + if end_word == words.len() { + break; + } + let mut back_tokens = 0usize; + let mut next = end_word; + while next > word + 1 && back_tokens < overlap_tokens { + next -= 1; + back_tokens += word_token_counts[next]; + } + word = next.max(word + 1); + } + Ok(ranges) +} + fn split_words(text: &str) -> Vec<(String, (usize, usize))> { let mut words = Vec::new(); let mut chars = text.char_indices().peekable(); From 87e8b30231e65f40169e3498ddb16e76bc1f7178 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Tue, 14 Apr 2026 19:45:57 -0500 Subject: [PATCH 15/34] switch to dx components --- demos/rgliner-web/Cargo.toml | 1 + .../assets/dx-components-theme.css | 87 +++++ .../src/components/badge/component.rs | 82 +++++ demos/rgliner-web/src/components/badge/mod.rs | 2 + .../src/components/badge/style.css | 42 +++ .../src/components/button/component.rs | 68 ++++ .../rgliner-web/src/components/button/mod.rs | 2 + .../src/components/button/style.css | 60 ++++ .../src/components/card/component.rs | 107 ++++++ demos/rgliner-web/src/components/card/mod.rs | 3 + .../rgliner-web/src/components/card/style.css | 52 +++ .../src/components/input/component.rs | 54 +++ demos/rgliner-web/src/components/input/mod.rs | 2 + .../src/components/input/style.css | 39 +++ .../src/components/label/component.rs | 15 + demos/rgliner-web/src/components/label/mod.rs | 2 + .../src/components/label/style.css | 8 + demos/rgliner-web/src/components/mod.rs | 9 + .../src/components/select/component.rs | 116 +++++++ .../rgliner-web/src/components/select/mod.rs | 2 + .../src/components/select/style.css | 155 +++++++++ .../src/components/tabs/component.rs | 119 +++++++ demos/rgliner-web/src/components/tabs/mod.rs | 2 + .../rgliner-web/src/components/tabs/style.css | 72 ++++ .../src/components/textarea/component.rs | 78 +++++ .../src/components/textarea/mod.rs | 2 + .../src/components/textarea/style.css | 84 +++++ demos/rgliner-web/src/main.rs | 313 ++++++++++-------- 28 files changed, 1438 insertions(+), 140 deletions(-) create mode 100644 demos/rgliner-web/assets/dx-components-theme.css create mode 100644 demos/rgliner-web/src/components/badge/component.rs create mode 100644 demos/rgliner-web/src/components/badge/mod.rs create mode 100644 demos/rgliner-web/src/components/badge/style.css create mode 100644 demos/rgliner-web/src/components/button/component.rs create mode 100644 demos/rgliner-web/src/components/button/mod.rs create mode 100644 demos/rgliner-web/src/components/button/style.css create mode 100644 demos/rgliner-web/src/components/card/component.rs create mode 100644 demos/rgliner-web/src/components/card/mod.rs create mode 100644 demos/rgliner-web/src/components/card/style.css create mode 100644 demos/rgliner-web/src/components/input/component.rs create mode 100644 demos/rgliner-web/src/components/input/mod.rs create mode 100644 demos/rgliner-web/src/components/input/style.css create mode 100644 demos/rgliner-web/src/components/label/component.rs create mode 100644 demos/rgliner-web/src/components/label/mod.rs create mode 100644 demos/rgliner-web/src/components/label/style.css create mode 100644 demos/rgliner-web/src/components/mod.rs create mode 100644 demos/rgliner-web/src/components/select/component.rs create mode 100644 demos/rgliner-web/src/components/select/mod.rs create mode 100644 demos/rgliner-web/src/components/select/style.css create mode 100644 demos/rgliner-web/src/components/tabs/component.rs create mode 100644 demos/rgliner-web/src/components/tabs/mod.rs create mode 100644 demos/rgliner-web/src/components/tabs/style.css create mode 100644 demos/rgliner-web/src/components/textarea/component.rs create mode 100644 demos/rgliner-web/src/components/textarea/mod.rs create mode 100644 demos/rgliner-web/src/components/textarea/style.css diff --git a/demos/rgliner-web/Cargo.toml b/demos/rgliner-web/Cargo.toml index b9edaa640..f362cb1e1 100644 --- a/demos/rgliner-web/Cargo.toml +++ b/demos/rgliner-web/Cargo.toml @@ -9,6 +9,7 @@ dioxus = { version = "=0.7.2", features = ["web"] } rgliner = { path = "../../models/rgliner", default-features = false } getrandom = { version = "0.3", features = ["wasm_js"] } tracing = "0.1" +dioxus-primitives = { git = "https://github.com/DioxusLabs/components", version = "0.0.1", default-features = false } [target.'cfg(target_arch = "wasm32")'.dependencies] console_error_panic_hook = "0.1" diff --git a/demos/rgliner-web/assets/dx-components-theme.css b/demos/rgliner-web/assets/dx-components-theme.css new file mode 100644 index 000000000..c55ca4a63 --- /dev/null +++ b/demos/rgliner-web/assets/dx-components-theme.css @@ -0,0 +1,87 @@ +/* This file contains the global styles for the styled dioxus components. You only + * need to import this file once in your project root. + */ +@import url("https://fonts.googleapis.com/css2?family=Inter:ital,opsz,wght@0,14..32,100..900;1,14..32,100..900&display=swap"); + +body { + color: var(--secondary-color-4); + font-family: Inter, sans-serif; + font-optical-sizing: auto; + font-style: normal; + font-weight: 400; +} + +html[data-theme="dark"] { + --dark: initial; + --light: ; +} + +html[data-theme="light"] { + --dark: ; + --light: initial; +} + +@media (prefers-color-scheme: dark) { + :root { + --dark: initial; + --light: ; + } +} + +@media (prefers-color-scheme: light) { + :root { + --dark: ; + --light: initial; + } +} + +:root { + /* Primary colors */ + --primary-color: var(--dark, #000) var(--light, #fff); + --primary-color-1: var(--dark, #0e0e0e) var(--light, #fbfbfb); + --primary-color-2: var(--dark, #0a0a0a) var(--light, #fff); + --primary-color-3: var(--dark, #141313) var(--light, #f8f8f8); + --primary-color-4: var(--dark, #1a1a1a) var(--light, #f8f8f8); + --primary-color-5: var(--dark, #262626) var(--light, #f5f5f5); + --primary-color-6: var(--dark, #232323) var(--light, #e5e5e5); + --primary-color-7: var(--dark, #3e3e3e) var(--light, #b0b0b0); + + /* Secondary colors */ + --secondary-color: var(--dark, #fff) var(--light, #000); + --secondary-color-1: var(--dark, #fafafa) var(--light, #000); + --secondary-color-2: var(--dark, #e6e6e6) var(--light, #0d0d0d); + --secondary-color-3: var(--dark, #dcdcdc) var(--light, #2b2b2b); + --secondary-color-4: var(--dark, #d4d4d4) var(--light, #111); + --secondary-color-5: var(--dark, #a1a1a1) var(--light, #848484); + --secondary-color-6: var(--dark, #5d5d5d) var(--light, #d0d0d0); + + /* Highlight colors */ + --focused-border-color: var(--dark, #2b7fff) var(--light, #2b7fff); + --primary-success-color: var(--dark, #02271c) var(--light, #ecfdf5); + --secondary-success-color: var(--dark, #b6fae3) var(--light, #10b981); + --primary-warning-color: var(--dark, #342203) var(--light, #fffbeb); + --secondary-warning-color: var(--dark, #feeac7) var(--light, #f59e0b); + --primary-error-color: var(--dark, #a22e2e) var(--light, #dc2626); + --secondary-error-color: var(--dark, #9b1c1c) var(--light, #ef4444); + --contrast-error-color: var(--dark, var(--secondary-color-3)) var(--light, var(--primary-color)); + --primary-info-color: var(--dark, var(--primary-color-5)) var(--light, var(--primary-color)); + --secondary-info-color: var(--dark, var(--primary-color-7)) var(--light, var(--secondary-color-3)); +} + +/* Modern browsers with `scrollbar-*` support */ +@supports (scrollbar-width: auto) { + :not(:hover) { + scrollbar-color: rgb(0 0 0 / 0%) rgb(0 0 0 / 0%); + } + + :hover { + scrollbar-color: var(--secondary-color-2) rgb(0 0 0 / 0%); + } +} + +/* Legacy browsers with `::-webkit-scrollbar-*` support */ +@supports selector(::-webkit-scrollbar) { + :root::-webkit-scrollbar-track { + background: transparent; + } +} diff --git a/demos/rgliner-web/src/components/badge/component.rs b/demos/rgliner-web/src/components/badge/component.rs new file mode 100644 index 000000000..831150562 --- /dev/null +++ b/demos/rgliner-web/src/components/badge/component.rs @@ -0,0 +1,82 @@ +use dioxus::prelude::*; + +#[derive(Copy, Clone, PartialEq, Default)] +#[non_exhaustive] +pub enum BadgeVariant { + #[default] + Primary, + Secondary, + Destructive, + Outline, +} + +impl BadgeVariant { + pub fn class(&self) -> &'static str { + match self { + BadgeVariant::Primary => "primary", + BadgeVariant::Secondary => "secondary", + BadgeVariant::Destructive => "destructive", + BadgeVariant::Outline => "outline", + } + } +} + +/// The props for the [`Badge`] component. +#[derive(Props, Clone, PartialEq)] +pub struct BadgeProps { + #[props(default)] + pub variant: BadgeVariant, + + /// Additional attributes to extend the badge element + #[props(extends = GlobalAttributes)] + pub attributes: Vec, + + /// The children of the badge element + pub children: Element, +} + +#[component] +pub fn Badge(props: BadgeProps) -> Element { + rsx! { + document::Link { rel: "stylesheet", href: asset!("./style.css") } + + BadgeElement { + "padding": true, + variant: props.variant, + attributes: props.attributes, + {props.children} + } + } +} + +#[component] +fn BadgeElement(props: BadgeProps) -> Element { + rsx! { + span { + class: "badge", + "data-style": props.variant.class(), + ..props.attributes, + {props.children} + } + } +} + +#[component] +pub fn VerifiedIcon() -> Element { + rsx! { + // Badge icon from lucide https://lucide.dev/icons/badge + svg { + view_box: "0 0 24 24", + xmlns: "http://www.w3.org/2000/svg", + width: "12", + height: "12", + fill: "none", + stroke: "var(--secondary-color-4)", + stroke_linecap: "round", + stroke_linejoin: "round", + stroke_width: 2, + path { d: "M3.85 8.62a4 4 0 0 1 4.78-4.77 4 4 0 0 1 6.74 0 4 4 0 0 1 4.78 4.78 4 4 0 0 1 0 6.74 4 4 0 0 1-4.77 4.78 4 4 0 0 1-6.75 0 4 4 0 0 1-4.78-4.77 4 4 0 0 1 0-6.76Z" } + path { d: "m9 12 2 2 4-4" } + } + } +} diff --git a/demos/rgliner-web/src/components/badge/mod.rs b/demos/rgliner-web/src/components/badge/mod.rs new file mode 100644 index 000000000..9a8ae5565 --- /dev/null +++ b/demos/rgliner-web/src/components/badge/mod.rs @@ -0,0 +1,2 @@ +mod component; +pub use component::*; \ No newline at end of file diff --git a/demos/rgliner-web/src/components/badge/style.css b/demos/rgliner-web/src/components/badge/style.css new file mode 100644 index 000000000..e36df538e --- /dev/null +++ b/demos/rgliner-web/src/components/badge/style.css @@ -0,0 +1,42 @@ +.badge-example { + display: flex; + align-items: center; + gap: 1rem; +} + +.badge { + display: inline-flex; + min-width: 20px; + height: 20px; + align-items: center; + justify-content: center; + border-radius: 10px; + box-shadow: 0 0 0 1px var(--primary-color-2); + font-size: 12px; + gap: 4px +} + +.badge[padding="true"] { + padding: 0 8px; +} + +.badge[data-style="primary"] { + background-color: var(--secondary-color-2); + color: var(--primary-color); +} + +.badge[data-style="secondary"] { + background-color: var(--primary-color-5); + color: var(--secondary-color-1); +} + +.badge[data-style="outline"] { + border: 1px solid var(--primary-color-6); + background-color: var(--light, var(--primary-color)) var(--dark, var(--primary-color-3)); + color: var(--secondary-color-4); +} + +.badge[data-style="destructive"] { + background-color: var(--primary-error-color); + color: var(--contrast-error-color); +} \ No newline at end of file diff --git a/demos/rgliner-web/src/components/button/component.rs b/demos/rgliner-web/src/components/button/component.rs new file mode 100644 index 000000000..e057876ca --- /dev/null +++ b/demos/rgliner-web/src/components/button/component.rs @@ -0,0 +1,68 @@ +use dioxus::prelude::*; +use dioxus_primitives::dioxus_attributes::attributes; +use dioxus_primitives::merge_attributes; + +#[derive(Copy, Clone, PartialEq, Default)] +#[non_exhaustive] +pub enum ButtonVariant { + #[default] + Primary, + Secondary, + Destructive, + Outline, + Ghost, +} + +impl ButtonVariant { + pub fn class(&self) -> &'static str { + match self { + ButtonVariant::Primary => "primary", + ButtonVariant::Secondary => "secondary", + ButtonVariant::Destructive => "destructive", + ButtonVariant::Outline => "outline", + ButtonVariant::Ghost => "ghost", + } + } +} + +#[component] +pub fn Button( + #[props(default)] variant: ButtonVariant, + #[props(extends=GlobalAttributes)] + #[props(extends=button)] + attributes: Vec, + onclick: Option>, + onmousedown: Option>, + onmouseup: Option>, + children: Element, +) -> Element { + let base = attributes!(button { + class: "button", + "data-style": variant.class(), + }); + let merged = merge_attributes(vec![base, attributes]); + + rsx! { + document::Link { rel: "stylesheet", href: asset!("./style.css") } + + button { + onclick: move |event| { + if let Some(f) = &onclick { + f.call(event); + } + }, + onmousedown: move |event| { + if let Some(f) = &onmousedown { + f.call(event); + } + }, + onmouseup: move |event| { + if let Some(f) = &onmouseup { + f.call(event); + } + }, + ..merged, + {children} + } + } +} diff --git a/demos/rgliner-web/src/components/button/mod.rs b/demos/rgliner-web/src/components/button/mod.rs new file mode 100644 index 000000000..9a8ae5565 --- /dev/null +++ b/demos/rgliner-web/src/components/button/mod.rs @@ -0,0 +1,2 @@ +mod component; +pub use component::*; \ No newline at end of file diff --git a/demos/rgliner-web/src/components/button/style.css b/demos/rgliner-web/src/components/button/style.css new file mode 100644 index 000000000..2f0d9ed66 --- /dev/null +++ b/demos/rgliner-web/src/components/button/style.css @@ -0,0 +1,60 @@ +.button { + padding: 8px 18px; + border: none; + border-radius: 0.5rem; + cursor: pointer; + font-size: 1rem; + transition: background-color 0.2s ease, color 0.2s ease; +} + +.button:focus-visible { + box-shadow: 0 0 0 2px var(--focused-border-color); +} + +.button[data-style="primary"] { + background-color: var(--secondary-color-2); + color: var(--primary-color); +} + +.button[data-style="primary"]:hover { + background-color: var(--secondary-color-1); +} + +.button[data-style="secondary"] { + background-color: var(--primary-color-5); + color: var(--secondary-color-1); +} + +.button[data-style="secondary"]:hover { + background-color: var(--primary-color-4); +} + +.button[data-style="ghost"] { + background-color: transparent; + color: var(--secondary-color-4); +} + +.button[data-style="ghost"]:hover { + background-color: var(--primary-color-5); + color: var(--secondary-color-1); +} + +.button[data-style="outline"] { + border: 1px solid var(--primary-color-6); + background-color: var(--light, var(--primary-color)) + var(--dark, var(--primary-color-3)); + color: var(--secondary-color-4); +} + +.button[data-style="outline"]:hover { + background-color: var(--primary-color-4); +} + +.button[data-style="destructive"] { + background-color: var(--primary-error-color); + color: var(--contrast-error-color); +} + +.button[data-style="destructive"]:hover { + background-color: var(--secondary-error-color); +} diff --git a/demos/rgliner-web/src/components/card/component.rs b/demos/rgliner-web/src/components/card/component.rs new file mode 100644 index 000000000..036749a7a --- /dev/null +++ b/demos/rgliner-web/src/components/card/component.rs @@ -0,0 +1,107 @@ +use dioxus::prelude::*; + +#[component] +pub fn Card( + #[props(extends=GlobalAttributes)] attributes: Vec, + children: Element, +) -> Element { + rsx! { + document::Link { rel: "stylesheet", href: asset!("./style.css") } + div { + class: "card", + "data-slot": "card", + ..attributes, + {children} + } + } +} + +#[component] +pub fn CardHeader( + #[props(extends=GlobalAttributes)] attributes: Vec, + children: Element, +) -> Element { + rsx! { + div { + class: "card-header", + "data-slot": "card-header", + ..attributes, + {children} + } + } +} + +#[component] +pub fn CardTitle( + #[props(extends=GlobalAttributes)] attributes: Vec, + children: Element, +) -> Element { + rsx! { + div { + class: "card-title", + "data-slot": "card-title", + ..attributes, + {children} + } + } +} + +#[component] +pub fn CardDescription( + #[props(extends=GlobalAttributes)] attributes: Vec, + children: Element, +) -> Element { + rsx! { + div { + class: "card-description", + "data-slot": "card-description", + ..attributes, + {children} + } + } +} + +#[component] +pub fn CardAction( + #[props(extends=GlobalAttributes)] attributes: Vec, + children: Element, +) -> Element { + rsx! { + div { + class: "card-action", + "data-slot": "card-action", + ..attributes, + {children} + } + } +} + +#[component] +pub fn CardContent( + #[props(extends=GlobalAttributes)] attributes: Vec, + children: Element, +) -> Element { + rsx! { + div { + class: "card-content", + "data-slot": "card-content", + ..attributes, + {children} + } + } +} + +#[component] +pub fn CardFooter( + #[props(extends=GlobalAttributes)] attributes: Vec, + children: Element, +) -> Element { + rsx! { + div { + class: "card-footer", + "data-slot": "card-footer", + ..attributes, + {children} + } + } +} diff --git a/demos/rgliner-web/src/components/card/mod.rs b/demos/rgliner-web/src/components/card/mod.rs new file mode 100644 index 000000000..a3527a11b --- /dev/null +++ b/demos/rgliner-web/src/components/card/mod.rs @@ -0,0 +1,3 @@ +mod component; +pub use component::*; + diff --git a/demos/rgliner-web/src/components/card/style.css b/demos/rgliner-web/src/components/card/style.css new file mode 100644 index 000000000..2ad7e6e6a --- /dev/null +++ b/demos/rgliner-web/src/components/card/style.css @@ -0,0 +1,52 @@ +.card { + display: flex; + flex-direction: column; + padding: 1.5rem 0; + border: 1px solid var(--light, var(--primary-color-6)) var(--dark, var(--primary-color-5)); + border-radius: 1rem; + background-color: var(--light, var(--primary-color-2)) var(--dark, var(--primary-color-3)); + box-shadow: 0 2px 10px rgb(0 0 0 / 10%); + color: var(--secondary-color-4); + gap: 1.5rem; +} + +.card-header { + display: grid; + align-items: start; + padding: 0 1.5rem; + gap: 0.5rem; + grid-auto-rows: min-content; + grid-template-rows: auto auto; +} + +.card-header:has([data-slot="card-action"]) { + grid-template-columns: 1fr auto; +} + +.card-title { + font-size: 1rem; + font-weight: 600; + line-height: 1; +} + +.card-description { + color: var(--secondary-color-5); + font-size: 0.875rem; + line-height: 1.25rem; +} + +.card-action { + grid-column-start: 2; + grid-row: 1 / span 2; + place-self: start end; +} + +.card-content { + padding: 0 1.5rem; +} + +.card-footer { + display: flex; + align-items: center; + padding: 0 1.5rem; +} diff --git a/demos/rgliner-web/src/components/input/component.rs b/demos/rgliner-web/src/components/input/component.rs new file mode 100644 index 000000000..5a1807edc --- /dev/null +++ b/demos/rgliner-web/src/components/input/component.rs @@ -0,0 +1,54 @@ +use dioxus::prelude::*; + +#[component] +pub fn Input( + oninput: Option>, + onchange: Option>, + oninvalid: Option>, + onselect: Option>, + onselectionchange: Option>, + onfocus: Option>, + onblur: Option>, + onfocusin: Option>, + onfocusout: Option>, + onkeydown: Option>, + onkeypress: Option>, + onkeyup: Option>, + oncompositionstart: Option>, + oncompositionupdate: Option>, + oncompositionend: Option>, + oncopy: Option>, + oncut: Option>, + onpaste: Option>, + #[props(extends=GlobalAttributes)] + #[props(extends=input)] + attributes: Vec, + children: Element, +) -> Element { + rsx! { + document::Link { rel: "stylesheet", href: asset!("./style.css") } + input { + class: "input", + oninput: move |e| _ = oninput.map(|callback| callback(e)), + onchange: move |e| _ = onchange.map(|callback| callback(e)), + oninvalid: move |e| _ = oninvalid.map(|callback| callback(e)), + onselect: move |e| _ = onselect.map(|callback| callback(e)), + onselectionchange: move |e| _ = onselectionchange.map(|callback| callback(e)), + onfocus: move |e| _ = onfocus.map(|callback| callback(e)), + onblur: move |e| _ = onblur.map(|callback| callback(e)), + onfocusin: move |e| _ = onfocusin.map(|callback| callback(e)), + onfocusout: move |e| _ = onfocusout.map(|callback| callback(e)), + onkeydown: move |e| _ = onkeydown.map(|callback| callback(e)), + onkeypress: move |e| _ = onkeypress.map(|callback| callback(e)), + onkeyup: move |e| _ = onkeyup.map(|callback| callback(e)), + oncompositionstart: move |e| _ = oncompositionstart.map(|callback| callback(e)), + oncompositionupdate: move |e| _ = oncompositionupdate.map(|callback| callback(e)), + oncompositionend: move |e| _ = oncompositionend.map(|callback| callback(e)), + oncopy: move |e| _ = oncopy.map(|callback| callback(e)), + oncut: move |e| _ = oncut.map(|callback| callback(e)), + onpaste: move |e| _ = onpaste.map(|callback| callback(e)), + ..attributes, + {children} + } + } +} diff --git a/demos/rgliner-web/src/components/input/mod.rs b/demos/rgliner-web/src/components/input/mod.rs new file mode 100644 index 000000000..9a8ae5565 --- /dev/null +++ b/demos/rgliner-web/src/components/input/mod.rs @@ -0,0 +1,2 @@ +mod component; +pub use component::*; \ No newline at end of file diff --git a/demos/rgliner-web/src/components/input/style.css b/demos/rgliner-web/src/components/input/style.css new file mode 100644 index 000000000..9faf0c950 --- /dev/null +++ b/demos/rgliner-web/src/components/input/style.css @@ -0,0 +1,39 @@ +/* Input Styles */ +.input { + position: relative; + display: flex; + box-sizing: border-box; + flex-direction: row; + align-items: center; + justify-content: space-between; + padding: 0.25rem; + padding: 8px 12px; + border: none; + border-radius: 0.5rem; + border-radius: calc(0.5rem); + background: none; + background-color: var(--light, var(--primary-color)) var(--dark, color-mix(in oklab, #FFFFFF26 30%, transparent)); + box-shadow: inset 0 0 0 1px var(--light, var(--primary-color-6)) + var(--dark, var(--primary-color-7)); + color: var(--secondary-color-4); + cursor: pointer; + gap: 0.25rem; + transition: background-color 100ms ease-out; +} + +.input::placeholder { + color: var(--secondary-color-5); +} + +.input:disabled { + color: var(--secondary-color-5); + cursor: not-allowed; +} + +.input:hover:not(:disabled), +.input:focus-visible { + background: var(--light, var(--primary-color-4)) + var(--dark, color-mix(in oklab, #FFFFFF26 50%, transparent)); + color: var(--secondary-color-1); + outline: none; +} diff --git a/demos/rgliner-web/src/components/label/component.rs b/demos/rgliner-web/src/components/label/component.rs new file mode 100644 index 000000000..3874e5be5 --- /dev/null +++ b/demos/rgliner-web/src/components/label/component.rs @@ -0,0 +1,15 @@ +use dioxus::prelude::*; +use dioxus_primitives::label::{self, LabelProps}; + +#[component] +pub fn Label(props: LabelProps) -> Element { + rsx! { + document::Link { rel: "stylesheet", href: asset!("./style.css") } + label::Label { + class: "label", + html_for: props.html_for, + attributes: props.attributes, + {props.children} + } + } +} diff --git a/demos/rgliner-web/src/components/label/mod.rs b/demos/rgliner-web/src/components/label/mod.rs new file mode 100644 index 000000000..9a8ae5565 --- /dev/null +++ b/demos/rgliner-web/src/components/label/mod.rs @@ -0,0 +1,2 @@ +mod component; +pub use component::*; \ No newline at end of file diff --git a/demos/rgliner-web/src/components/label/style.css b/demos/rgliner-web/src/components/label/style.css new file mode 100644 index 000000000..3cc7f3a64 --- /dev/null +++ b/demos/rgliner-web/src/components/label/style.css @@ -0,0 +1,8 @@ +/* Label Styles */ +.label { + display: flex; + align-items: center; + color: var(--secondary-color-4); + font-size: 0.8rem; + line-height: 1; +} diff --git a/demos/rgliner-web/src/components/mod.rs b/demos/rgliner-web/src/components/mod.rs new file mode 100644 index 000000000..da7dda254 --- /dev/null +++ b/demos/rgliner-web/src/components/mod.rs @@ -0,0 +1,9 @@ +// AUTOGENERTED Components module +pub mod textarea; +pub mod tabs; +pub mod label; +pub mod button; +pub mod card; +pub mod input; +pub mod select; +pub mod badge; diff --git a/demos/rgliner-web/src/components/select/component.rs b/demos/rgliner-web/src/components/select/component.rs new file mode 100644 index 000000000..88e70f02d --- /dev/null +++ b/demos/rgliner-web/src/components/select/component.rs @@ -0,0 +1,116 @@ +use dioxus::prelude::*; +use dioxus_primitives::select::{ + self, SelectGroupLabelProps, SelectGroupProps, SelectListProps, SelectOptionProps, SelectProps, + SelectTriggerProps, SelectValueProps, +}; + +#[component] +pub fn Select(props: SelectProps) -> Element { + rsx! { + document::Link { rel: "stylesheet", href: asset!("./style.css") } + select::Select { + class: "select", + value: props.value, + default_value: props.default_value, + on_value_change: props.on_value_change, + disabled: props.disabled, + name: props.name, + placeholder: props.placeholder, + roving_loop: props.roving_loop, + typeahead_timeout: props.typeahead_timeout, + attributes: props.attributes, + {props.children} + } + } +} + +#[component] +pub fn SelectTrigger(props: SelectTriggerProps) -> Element { + rsx! { + select::SelectTrigger { class: "select-trigger", attributes: props.attributes, + {props.children} + svg { + class: "select-expand-icon", + view_box: "0 0 24 24", + xmlns: "http://www.w3.org/2000/svg", + polyline { points: "6 9 12 15 18 9" } + } + } + } +} + +#[component] +pub fn SelectValue(props: SelectValueProps) -> Element { + rsx! { + select::SelectValue { attributes: props.attributes } + } +} + +#[component] +pub fn SelectList(props: SelectListProps) -> Element { + rsx! { + select::SelectList { + class: "select-list", + id: props.id, + attributes: props.attributes, + {props.children} + } + } +} + +#[component] +pub fn SelectGroup(props: SelectGroupProps) -> Element { + rsx! { + select::SelectGroup { + class: "select-group", + disabled: props.disabled, + id: props.id, + attributes: props.attributes, + {props.children} + } + } +} + +#[component] +pub fn SelectGroupLabel(props: SelectGroupLabelProps) -> Element { + rsx! { + select::SelectGroupLabel { + class: "select-group-label", + id: props.id, + attributes: props.attributes, + {props.children} + } + } +} + +#[component] +pub fn SelectOption(props: SelectOptionProps) -> Element { + rsx! { + select::SelectOption:: { + class: "select-option", + value: props.value, + text_value: props.text_value, + disabled: props.disabled, + id: props.id, + index: props.index, + aria_label: props.aria_label, + aria_roledescription: props.aria_roledescription, + attributes: props.attributes, + {props.children} + } + } +} + +#[component] +pub fn SelectItemIndicator() -> Element { + rsx! { + select::SelectItemIndicator { + svg { + class: "select-check-icon", + view_box: "0 0 24 24", + xmlns: "http://www.w3.org/2000/svg", + path { d: "M5 13l4 4L19 7" } + } + } + } +} diff --git a/demos/rgliner-web/src/components/select/mod.rs b/demos/rgliner-web/src/components/select/mod.rs new file mode 100644 index 000000000..9a8ae5565 --- /dev/null +++ b/demos/rgliner-web/src/components/select/mod.rs @@ -0,0 +1,2 @@ +mod component; +pub use component::*; \ No newline at end of file diff --git a/demos/rgliner-web/src/components/select/style.css b/demos/rgliner-web/src/components/select/style.css new file mode 100644 index 000000000..5a98dd3ed --- /dev/null +++ b/demos/rgliner-web/src/components/select/style.css @@ -0,0 +1,155 @@ +.select { + position: relative; +} + +.select-trigger { + position: relative; + display: flex; + box-sizing: border-box; + flex-direction: row; + align-items: center; + justify-content: space-between; + padding: 0.25rem; + padding: 8px 12px; + border: none; + border-radius: 0.5rem; + border-radius: calc(0.5rem); + background: none; + background: var(--light, var(--primary-color)) + var(--dark, var(--primary-color-3)); + box-shadow: inset 0 0 0 1px var(--light, var(--primary-color-6)) + var(--dark, var(--primary-color-7)); + color: var(--secondary-color-4); + cursor: pointer; + gap: 0.25rem; + transition: background-color 100ms ease-out; +} + +.select-trigger span[data-placeholder="true"] { + color: var(--secondary-color-5); +} + +.select[data-state="open"] .select-trigger { + pointer-events: none; +} + +.select-expand-icon { + width: 20px; + height: 20px; + fill: none; + stroke: var(--primary-color-7); + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 2; +} + +.select-check-icon { + width: 1rem; + height: 1rem; + fill: none; + stroke: var(--secondary-color-5); + stroke-linecap: round; + stroke-linejoin: round; + stroke-width: 2; +} + +.select[data-disabled="true"] .select-trigger { + color: var(--secondary-color-5); + cursor: not-allowed; +} + +.select-trigger:hover:not([data-disabled="true"]), +.select-trigger:focus-visible { + background: var(--light, var(--primary-color-4)) + var(--dark, var(--primary-color-5)); + color: var(--secondary-color-1); + outline: none; +} + +.select-list { + position: absolute; + z-index: 1000; + top: 100%; + left: 0; + min-width: 100%; + box-sizing: border-box; + padding: 0.25rem; + border-radius: 0.5rem; + margin-top: 0.25rem; + background: var(--light, var(--primary-color)) + var(--dark, var(--primary-color-5)); + box-shadow: inset 0 0 0 1px var(--light, var(--primary-color-6)) + var(--dark, var(--primary-color-7)); + opacity: 0; + pointer-events: none; + transform-origin: top; + will-change: transform, opacity; +} + +.select-list[data-state="closed"] { + animation: select-list-animate-out 150ms ease-in forwards; + pointer-events: none; +} + +@keyframes select-list-animate-out { + 0% { + opacity: 1; + transform: scale(1) translateY(0); + } + + 100% { + opacity: 0; + transform: scale(0.95) translateY(-2px); + } +} + +.select-list[data-state="open"] { + animation: select-list-animate-in 150ms ease-out forwards; + pointer-events: auto; +} + +@keyframes select-list-animate-in { + 0% { + opacity: 0; + transform: scale(0.95) translateY(-2px); + } + + 100% { + opacity: 1; + transform: scale(1) translateY(0); + } +} + +.select-option { + display: flex; + align-items: center; + justify-content: space-between; + padding: 8px 12px; + border-radius: calc(0.5rem - 0.25rem); + cursor: pointer; + font-size: 14px; +} + +.select-option[data-disabled="true"] { + color: var(--secondary-color-5); + cursor: not-allowed; +} + +.select-option:hover:not([data-disabled="true"]), +.select-option:focus-visible { + background: var(--light, var(--primary-color-4)) + var(--dark, var(--primary-color-7)); + color: var(--secondary-color-1); + outline: none; +} + +.select-group-label { + padding: 4px 12px; + color: var(--secondary-color-5); + font-size: 0.75rem; +} + +[data-disabled="true"] { + cursor: not-allowed; + opacity: 0.5; +} diff --git a/demos/rgliner-web/src/components/tabs/component.rs b/demos/rgliner-web/src/components/tabs/component.rs new file mode 100644 index 000000000..93ade03ad --- /dev/null +++ b/demos/rgliner-web/src/components/tabs/component.rs @@ -0,0 +1,119 @@ +use dioxus::prelude::*; +use dioxus_primitives::tabs::{self, TabContentProps, TabListProps, TabTriggerProps}; + +/// The props for the [`Tabs`] component. +#[derive(Props, Clone, PartialEq)] +pub struct TabsProps { + /// The class of the tabs component. + #[props(default)] + pub class: String, + + /// The controlled value of the active tab. + pub value: ReadSignal>, + + /// The default active tab value when uncontrolled. + #[props(default)] + pub default_value: String, + + /// Callback fired when the active tab changes. + #[props(default)] + pub on_value_change: Callback, + + /// Whether the tabs are disabled. + #[props(default)] + pub disabled: ReadSignal, + + /// Whether the tabs are horizontal. + #[props(default)] + pub horizontal: ReadSignal, + + /// Whether focus should loop around when reaching the end. + #[props(default = ReadSignal::new(Signal::new(true)))] + pub roving_loop: ReadSignal, + + /// The variant of the tabs component. + #[props(default)] + pub variant: TabsVariant, + + /// Additional attributes to apply to the tabs element. + #[props(extends = GlobalAttributes)] + pub attributes: Vec, + + /// The children of the tabs component. + pub children: Element, +} + +/// The variant of the tabs component. +#[derive(Clone, Copy, PartialEq, Default)] +pub enum TabsVariant { + /// The default variant. + #[default] + Default, + /// The ghost variant. + Ghost, +} + +impl TabsVariant { + /// Convert the variant to a string for use in class names + fn to_class(self) -> &'static str { + match self { + TabsVariant::Default => "default", + TabsVariant::Ghost => "ghost", + } + } +} + +#[component] +pub fn Tabs(props: TabsProps) -> Element { + rsx! { + document::Link { rel: "stylesheet", href: asset!("./style.css") } + tabs::Tabs { + class: props.class + " tabs", + "data-variant": props.variant.to_class(), + value: props.value, + default_value: props.default_value, + on_value_change: props.on_value_change, + disabled: props.disabled, + horizontal: props.horizontal, + roving_loop: props.roving_loop, + attributes: props.attributes, + {props.children} + } + } +} + +#[component] +pub fn TabList(props: TabListProps) -> Element { + rsx! { + tabs::TabList { class: "tabs-list", attributes: props.attributes, {props.children} } + } +} + +#[component] +pub fn TabTrigger(props: TabTriggerProps) -> Element { + rsx! { + tabs::TabTrigger { + class: "tabs-trigger", + id: props.id, + value: props.value, + index: props.index, + disabled: props.disabled, + attributes: props.attributes, + {props.children} + } + } +} + +#[component] +pub fn TabContent(props: TabContentProps) -> Element { + rsx! { + tabs::TabContent { + class: props.class.unwrap_or_default() + " tabs-content tabs-content-themed", + value: props.value, + id: props.id, + index: props.index, + attributes: props.attributes, + {props.children} + } + } +} diff --git a/demos/rgliner-web/src/components/tabs/mod.rs b/demos/rgliner-web/src/components/tabs/mod.rs new file mode 100644 index 000000000..9a8ae5565 --- /dev/null +++ b/demos/rgliner-web/src/components/tabs/mod.rs @@ -0,0 +1,2 @@ +mod component; +pub use component::*; \ No newline at end of file diff --git a/demos/rgliner-web/src/components/tabs/style.css b/demos/rgliner-web/src/components/tabs/style.css new file mode 100644 index 000000000..5e04cd012 --- /dev/null +++ b/demos/rgliner-web/src/components/tabs/style.css @@ -0,0 +1,72 @@ +.tabs { + display: flex; + width: 100%; + flex-direction: column; + gap: 0.5rem; +} + +.tabs-list { + display: flex; + width: fit-content; + box-sizing: border-box; + flex: 1; + flex-direction: row; + padding: 0.25rem; + border: none; + border-radius: 0.5rem; + gap: 0.25rem; +} + +[data-variant="default"] .tabs-list { + background: var(--light, var(--primary-color-3)) + var(--dark, var(--primary-color-5)); +} + +.tabs-trigger { + padding: 4px 8px; + border: none; + border-radius: calc(0.5rem - 0.25rem); + background: none; + color: var(--secondary-color-5); + cursor: pointer; +} + +[data-variant="default"] .tabs-trigger[data-state="active"] { + background-color: var(--light, var(--primary-color)) + var(--dark, var(--primary-color-6)); + box-shadow: var(--dark, inset 0 0 0 1px var(--primary-color-7)) + var(--light, 0 1px 2px rgb(0 0 0 / 18%)); +} + +.tabs-trigger[data-state="active"] { + color: var(--secondary-color-1); +} + +.tabs-trigger[data-disabled="true"] { + color: var(--secondary-color-5); + cursor: not-allowed; +} + +.tabs-trigger:hover:not([data-disabled="true"]), +.tabs-trigger:focus-visible { + color: var(--secondary-color-3); +} + +.tabs-content { + width: 100%; + box-sizing: border-box; + padding: 0.25rem; +} + +[data-variant="default"] .tabs-content-themed { + border: 1px solid var(--light, var(--primary-color-6)) + var(--dark, var(--primary-color-7)); + border-radius: 0.5rem; + background: var(--light, var(--primary-color)) + var(--dark, var(--primary-color-3)); + box-shadow: var(--light, 0 1px 2px rgb(0 0 0 / 18%)) var(--dark, none); +} + +.tabs-content[data-state="inactive"] { + display: none; +} diff --git a/demos/rgliner-web/src/components/textarea/component.rs b/demos/rgliner-web/src/components/textarea/component.rs new file mode 100644 index 000000000..7f4db3ccb --- /dev/null +++ b/demos/rgliner-web/src/components/textarea/component.rs @@ -0,0 +1,78 @@ +use dioxus::prelude::*; + +#[derive(Copy, Clone, PartialEq, Default)] +#[non_exhaustive] +pub enum TextareaVariant { + #[default] + Default, + Fade, + Outline, + Ghost, +} + +impl TextareaVariant { + pub fn class(&self) -> &'static str { + match self { + TextareaVariant::Default => "default", + TextareaVariant::Fade => "fade", + TextareaVariant::Outline => "outline", + TextareaVariant::Ghost => "ghost", + } + } +} + +#[component] +pub fn Textarea( + oninput: Option>, + onchange: Option>, + oninvalid: Option>, + onselect: Option>, + onselectionchange: Option>, + onfocus: Option>, + onblur: Option>, + onfocusin: Option>, + onfocusout: Option>, + onkeydown: Option>, + onkeypress: Option>, + onkeyup: Option>, + oncompositionstart: Option>, + oncompositionupdate: Option>, + oncompositionend: Option>, + oncopy: Option>, + oncut: Option>, + onpaste: Option>, + #[props(default)] variant: TextareaVariant, + #[props(extends=GlobalAttributes)] + #[props(extends=textarea)] + attributes: Vec, + children: Element, +) -> Element { + rsx! { + document::Link { rel: "stylesheet", href: asset!("./style.css") } + textarea { + class: "textarea", + "data-slot": "textarea", + "data-style": variant.class(), + oninput: move |e| _ = oninput.map(|callback| callback(e)), + onchange: move |e| _ = onchange.map(|callback| callback(e)), + oninvalid: move |e| _ = oninvalid.map(|callback| callback(e)), + onselect: move |e| _ = onselect.map(|callback| callback(e)), + onselectionchange: move |e| _ = onselectionchange.map(|callback| callback(e)), + onfocus: move |e| _ = onfocus.map(|callback| callback(e)), + onblur: move |e| _ = onblur.map(|callback| callback(e)), + onfocusin: move |e| _ = onfocusin.map(|callback| callback(e)), + onfocusout: move |e| _ = onfocusout.map(|callback| callback(e)), + onkeydown: move |e| _ = onkeydown.map(|callback| callback(e)), + onkeypress: move |e| _ = onkeypress.map(|callback| callback(e)), + onkeyup: move |e| _ = onkeyup.map(|callback| callback(e)), + oncompositionstart: move |e| _ = oncompositionstart.map(|callback| callback(e)), + oncompositionupdate: move |e| _ = oncompositionupdate.map(|callback| callback(e)), + oncompositionend: move |e| _ = oncompositionend.map(|callback| callback(e)), + oncopy: move |e| _ = oncopy.map(|callback| callback(e)), + oncut: move |e| _ = oncut.map(|callback| callback(e)), + onpaste: move |e| _ = onpaste.map(|callback| callback(e)), + ..attributes, + {children} + } + } +} diff --git a/demos/rgliner-web/src/components/textarea/mod.rs b/demos/rgliner-web/src/components/textarea/mod.rs new file mode 100644 index 000000000..2590c0132 --- /dev/null +++ b/demos/rgliner-web/src/components/textarea/mod.rs @@ -0,0 +1,2 @@ +mod component; +pub use component::*; diff --git a/demos/rgliner-web/src/components/textarea/style.css b/demos/rgliner-web/src/components/textarea/style.css new file mode 100644 index 000000000..0a2c49986 --- /dev/null +++ b/demos/rgliner-web/src/components/textarea/style.css @@ -0,0 +1,84 @@ +/* Base */ +.textarea { + width: 100%; + min-height: 4rem; + box-sizing: border-box; + padding: 8px 12px; + border: none; + border-radius: 0.5rem; + margin: 0; + appearance: none; + background: none; + color: var(--secondary-color-4); + font-family: inherit; + line-height: 1.5; + outline: none; + resize: vertical; + transition: background-color 100ms ease-out, border-color 100ms ease-out, box-shadow 100ms ease-out; +} + +.textarea:disabled { + color: var(--secondary-color-5); + cursor: not-allowed; +} + +.textarea::placeholder { + color: var(--secondary-color-5); +} + +/* Default Variant */ +.textarea[data-style="default"] { + background: var(--light, var(--primary-color)) var(--dark, var(--primary-color-3)); + box-shadow: inset 0 0 0 1px var(--light, var(--primary-color-6)) var(--dark, var(--primary-color-7)); +} + +.textarea[data-style="default"]:hover:not(:disabled), +.textarea[data-style="default"]:focus { + background: var(--light, var(--primary-color-4)) var(--dark, var(--primary-color-5)); + color: var(--secondary-color-1); +} + +/* Fade Variant */ +.textarea[data-style="fade"] { + background: var(--light, var(--primary-color)) var(--dark, var(--primary-color-3)); +} + +.textarea[data-style="fade"]:hover:not(:disabled), +.textarea[data-style="fade"]:focus { + background: var(--light, var(--primary-color-4)) var(--dark, var(--primary-color-5)); + color: var(--secondary-color-1); +} + +/* Outline Variant */ +.textarea[data-style="outline"] { + border: 1px solid var(--primary-color-6); + background-color: var(--light, var(--primary-color)) + var(--dark, var(--primary-color-3)); +} + +.textarea[data-style="outline"]:hover:not(:disabled, :focus) { + border-color: var(--primary-color-7); +} + +.textarea[data-style="outline"]:focus { + border-color: var(--focused-border-color); +} + +.textarea[data-style="outline"]:invalid, +.textarea[data-style="outline"][aria-invalid="true"] { + border-color: var(--primary-error-color); +} + +/* Ghost Variant */ +.textarea[data-style="ghost"] { + background-color: transparent; +} + +.textarea[data-style="ghost"]:hover:not(:disabled) { + background-color: var(--primary-color-5); + color: var(--secondary-color-1); +} + +.textarea[data-style="ghost"]:focus { + border-color: var(--focused-border-color); +} diff --git a/demos/rgliner-web/src/main.rs b/demos/rgliner-web/src/main.rs index fffc0eb0f..dbae9fdcc 100644 --- a/demos/rgliner-web/src/main.rs +++ b/demos/rgliner-web/src/main.rs @@ -1,3 +1,15 @@ +mod components; + +use components::badge::{Badge, BadgeVariant}; +use components::button::Button; +use components::card::{Card, CardContent, CardDescription, CardHeader, CardTitle}; +use components::input::Input; +use components::label::Label; +use components::select::{ + Select, SelectItemIndicator, SelectList, SelectOption, SelectTrigger, SelectValue, +}; +use components::tabs::{TabContent, TabList, TabTrigger, Tabs}; +use components::textarea::Textarea; use dioxus::prelude::*; use rgliner::{ relation_decoding::Relation, @@ -22,6 +34,22 @@ enum Mode { Relex, } +impl Mode { + fn value(self) -> &'static str { + match self { + Mode::Ner => "ner", + Mode::Relex => "relex", + } + } + + fn from_value(v: &str) -> Mode { + match v { + "relex" => Mode::Relex, + _ => Mode::Ner, + } + } +} + #[derive(Clone, Copy, PartialEq, Eq)] enum ModelChoice { Edge, @@ -46,31 +74,6 @@ impl ModelChoice { } } - fn value(self) -> &'static str { - match self { - ModelChoice::Edge => "edge", - ModelChoice::Small => "small", - ModelChoice::Base => "base", - ModelChoice::Large => "large", - ModelChoice::RelexMulti => "relex-multi", - ModelChoice::RelexBase => "relex-base", - ModelChoice::RelexLarge => "relex-large", - } - } - - fn from_value(v: &str) -> Option { - Some(match v { - "edge" => ModelChoice::Edge, - "small" => ModelChoice::Small, - "base" => ModelChoice::Base, - "large" => ModelChoice::Large, - "relex-multi" => ModelChoice::RelexMulti, - "relex-base" => ModelChoice::RelexBase, - "relex-large" => ModelChoice::RelexLarge, - _ => return None, - }) - } - fn default_for(mode: Mode) -> Self { match mode { Mode::Ner => ModelChoice::Edge, @@ -96,14 +99,8 @@ impl ModelChoice { } enum LoadedModel { - Ner { - choice: ModelChoice, - inner: Gliner, - }, - Relex { - choice: ModelChoice, - inner: GlinerRelEx, - }, + Ner { choice: ModelChoice, inner: Gliner }, + Relex { choice: ModelChoice, inner: GlinerRelEx }, } impl LoadedModel { @@ -129,8 +126,7 @@ fn App() -> Element { "Apple Inc. was founded by Steve Jobs in California. Microsoft is headquartered in Redmond." .to_string() }); - let mut entity_labels = - use_signal(|| "person, organization, location".to_string()); + let mut entity_labels = use_signal(|| "person, organization, location".to_string()); let mut relation_labels = use_signal(|| "founded by, located in".to_string()); let mut model = use_signal(|| None::); @@ -141,7 +137,8 @@ fn App() -> Element { let mut relex_out = use_signal(RelexResult::default); let mut status = use_signal(|| "No model loaded".to_string()); - let mut on_mode_change = move |new_mode: Mode| { + let on_mode_change = move |v: String| { + let new_mode = Mode::from_value(&v); if mode() != new_mode { mode.set(new_mode); choice.set(ModelChoice::default_for(new_mode)); @@ -237,7 +234,11 @@ fn App() -> Element { .map(|m| m.choice() != choice()) .unwrap_or(true); + let current_mode = mode(); + let current_choice = choice(); + rsx! { + document::Link { rel: "stylesheet", href: asset!("/assets/dx-components-theme.css") } document::Link { rel: "stylesheet", href: asset!("/assets/style.css") } div { class: "app", header { class: "site-header", @@ -254,100 +255,129 @@ fn App() -> Element { } } - div { class: "tabs", - button { - class: if mode() == Mode::Ner { "tab active" } else { "tab" }, - onclick: move |_| on_mode_change(Mode::Ner), - "NER" - } - button { - class: if mode() == Mode::Relex { "tab active" } else { "tab" }, - onclick: move |_| on_mode_change(Mode::Relex), - "NER + Relations" + Tabs { + default_value: current_mode.value().to_string(), + horizontal: true, + on_value_change: on_mode_change, + TabList { + TabTrigger { value: "ner".to_string(), index: 0usize, "NER" } + TabTrigger { value: "relex".to_string(), index: 1usize, "NER + Relations" } } + TabContent { index: 0usize, value: "ner".to_string(), "" } + TabContent { index: 1usize, value: "relex".to_string(), "" } } if let Some(e) = error() { div { class: "err-banner", "{e}" } } - div { class: "panel", - label { "Model" } - div { class: "row", - select { - value: "{choice().value()}", - onchange: move |ev| { - if let Some(c) = ModelChoice::from_value(&ev.value()) { - choice.set(c); + Card { + CardHeader { + CardTitle { "Model" } + CardDescription { + "First load fetches GGUF weights (60 MB – 500 MB) from HuggingFace and caches them in the browser's Origin Private File System." + } + } + CardContent { + div { class: "row", + div { style: "min-width: 16rem;", + Select:: { + key: "{current_mode.value()}", + placeholder: "Select a model...", + default_value: current_choice, + on_value_change: move |v: Option| { + if let Some(c) = v { + choice.set(c); + } + }, + SelectTrigger { aria_label: "Model", SelectValue {} } + SelectList { aria_label: "Models", + for (i, c) in ModelChoice::for_mode(current_mode).iter().copied().enumerate() { + SelectOption:: { + index: i, + value: c, + text_value: c.label().to_string(), + "{c.label()}" + SelectItemIndicator {} + } + } + } } - }, - for c in ModelChoice::for_mode(mode()) { - option { value: "{c.value()}", "{c.label()}" } } - } - button { - class: "primary", - disabled: loading() || running() || (has_model && !model_mismatch), - onclick: on_load, - if loading() { "Loading…" } - else if has_model && !model_mismatch { "Loaded" } - else if has_model { "Reload" } - else { "Load model" } - } - span { - class: if error().is_some() { "status err" } else if has_model { "status ok" } else { "status" }, - "{status()}" + Button { + disabled: loading() || running() || (has_model && !model_mismatch), + onclick: on_load, + if loading() { + "Loading…" + } else if has_model && !model_mismatch { + "Loaded" + } else if has_model { + "Reload" + } else { + "Load model" + } + } + span { + class: if error().is_some() { "status err" } else if has_model { "status ok" } else { "status" }, + "{status()}" + } } } - p { class: "muted", - "First load fetches the GGUF weights (60 MB – 500 MB) from HuggingFace and caches them in the browser's Origin Private File System. Subsequent loads are instant." - } } - div { class: "panel", - label { "Text" } - textarea { - value: "{text}", - oninput: move |e| text.set(e.value()), + Card { + CardHeader { CardTitle { "Text" } } + CardContent { + Textarea { + value: "{text}", + rows: "4", + oninput: move |e: FormEvent| text.set(e.value()), + } } } - div { class: "panel", - label { "Entity labels (comma-separated)" } - input { - r#type: "text", - value: "{entity_labels}", - oninput: move |e| entity_labels.set(e.value()), - } - if mode() == Mode::Relex { - div { style: "margin-top: 0.75rem;", - label { "Relation labels (comma-separated)" } - input { - r#type: "text", - value: "{relation_labels}", - oninput: move |e| relation_labels.set(e.value()), + Card { + CardHeader { CardTitle { "Labels" } } + CardContent { + Label { html_for: "entity-labels", "Entity labels (comma-separated)" } + Input { + id: "entity-labels", + r#type: "text", + value: "{entity_labels}", + oninput: move |e: FormEvent| entity_labels.set(e.value()), + } + if current_mode == Mode::Relex { + div { style: "margin-top: 0.75rem;", + Label { html_for: "relation-labels", "Relation labels (comma-separated)" } + Input { + id: "relation-labels", + r#type: "text", + value: "{relation_labels}", + oninput: move |e: FormEvent| relation_labels.set(e.value()), + } } } } } - div { class: "panel", - div { class: "row", - button { - class: "primary", - disabled: !has_model || running() || loading() || model_mismatch, - onclick: on_extract, - if running() { "Extracting…" } else { "Extract" } - } - if model_mismatch && has_model { - span { class: "status", - "Model selection changed — reload to use it." + Card { + CardContent { + div { class: "row", + Button { + disabled: !has_model || running() || loading() || model_mismatch, + onclick: on_extract, + if running() { "Extracting…" } else { "Extract" } + } + if model_mismatch && has_model { + span { class: "status", + "Model selection changed — reload to use it." + } } } } } - { render_results(mode(), text(), ner_out(), relex_out()) } + { render_results(current_mode, text(), ner_out(), relex_out()) } } } } @@ -360,14 +390,9 @@ fn parse_labels(raw: &str) -> Vec { .collect() } -async fn build_model( - choice: ModelChoice, -) -> Result { +async fn build_model(choice: ModelChoice) -> Result { match choice { - ModelChoice::Edge - | ModelChoice::Small - | ModelChoice::Base - | ModelChoice::Large => { + ModelChoice::Edge | ModelChoice::Small | ModelChoice::Base | ModelChoice::Large => { let source = match choice { ModelChoice::Edge => GlinerSource::edge(), ModelChoice::Small => GlinerSource::small(), @@ -458,45 +483,53 @@ fn render_results( if entities.is_empty() && relations.is_empty() { return rsx! { - div { class: "panel", - p { class: "muted", "Run extraction to see results." } + Card { + CardContent { + p { class: "muted", "Run extraction to see results." } + } } }; } rsx! { - div { class: "panel", - div { class: "results", - { highlighted_text(text.clone(), entities.clone()) } - } + Card { + CardHeader { CardTitle { "Results" } } + CardContent { + div { class: "results", + { highlighted_text(text.clone(), entities.clone()) } + } - if !entities.is_empty() { - ul { class: "entity-list", - for (i, ent) in entities.iter().enumerate() { - li { key: "{i}", - span { class: "label", style: "color: {hsl_for(&ent.label)};", "{ent.label}" } - " · " - span { "{ent.text:?}" } - " " - span { class: "score", "{format_score(ent.score)}" } + if !entities.is_empty() { + ul { class: "entity-list", + for (i, ent) in entities.iter().enumerate() { + li { key: "{i}", + Badge { + variant: BadgeVariant::Outline, + span { style: "color: {hsl_for(&ent.label)};", "{ent.label}" } + } + " · " + span { "{ent.text:?}" } + " " + span { class: "score", "{format_score(ent.score)}" } + } } } } - } - if !relations.is_empty() { - div { style: "margin-top: 1rem;", - label { "Relations" } - ul { class: "relation-list", - for (i, rel) in relations.iter().enumerate() { - li { key: "rel-{i}", - span { class: "label", "{rel.head.text}" } - " --[" - span { style: "color: var(--accent);", "{rel.relation}" } - "]--> " - span { class: "label", "{rel.tail.text}" } - " " - span { class: "score", "{format_score(rel.score)}" } + if !relations.is_empty() { + div { style: "margin-top: 1rem;", + Label { html_for: "relations-list", "Relations" } + ul { class: "relation-list", + for (i, rel) in relations.iter().enumerate() { + li { key: "rel-{i}", + Badge { "{rel.head.text}" } + " --[" + span { style: "color: var(--accent);", "{rel.relation}" } + "]--> " + Badge { "{rel.tail.text}" } + " " + span { class: "score", "{format_score(rel.score)}" } + } } } } From c8e6d42a4c8d285c8935eb9c00b0914a210a1b74 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Tue, 14 Apr 2026 20:54:14 -0500 Subject: [PATCH 16/34] rgliner batching --- Cargo.lock | 310 +++++- demos/rgliner-web/Cargo.toml | 2 + demos/rgliner-web/assets/style.css | 563 +++++++---- .../src/components/badge/component.rs | 82 -- demos/rgliner-web/src/components/badge/mod.rs | 2 - .../src/components/badge/style.css | 42 - .../src/components/card/component.rs | 107 --- demos/rgliner-web/src/components/card/mod.rs | 3 - .../rgliner-web/src/components/card/style.css | 52 - demos/rgliner-web/src/components/mod.rs | 3 - .../src/components/tabs/component.rs | 119 --- demos/rgliner-web/src/components/tabs/mod.rs | 2 - .../rgliner-web/src/components/tabs/style.css | 72 -- demos/rgliner-web/src/main.rs | 611 +++++------- fusor-ml/core/src/compute_graph/mod.rs | 14 + fusor-ml/core/src/compute_graph/resolve.rs | 82 +- fusor-ml/core/src/nary_wise.rs | 8 +- fusor-ml/core/src/tensor.rs | 10 + fusor-ml/fusor/src/lib.rs | 33 + models/rbert/src/raw/mdeberta/attention.rs | 18 + models/rbert/src/raw/mdeberta/model.rs | 23 + models/rgliner/src/lib.rs | 140 ++- models/rgliner/src/raw/bilstm.rs | 165 +++- models/rgliner/src/raw/joint_scorer.rs | 28 +- models/rgliner/src/raw/mod.rs | 1 + models/rgliner/src/raw/span_layer.rs | 13 +- models/rgliner/src/relex.rs | 898 ++++++++++++++---- models/rgliner/src/relex_tokenization.rs | 12 + models/rgliner/src/tokenization.rs | 10 + models/rgliner/tests/example_regression.rs | 211 +++- 30 files changed, 2276 insertions(+), 1360 deletions(-) delete mode 100644 demos/rgliner-web/src/components/badge/component.rs delete mode 100644 demos/rgliner-web/src/components/badge/mod.rs delete mode 100644 demos/rgliner-web/src/components/badge/style.css delete mode 100644 demos/rgliner-web/src/components/card/component.rs delete mode 100644 demos/rgliner-web/src/components/card/mod.rs delete mode 100644 demos/rgliner-web/src/components/card/style.css delete mode 100644 demos/rgliner-web/src/components/tabs/component.rs delete mode 100644 demos/rgliner-web/src/components/tabs/mod.rs delete mode 100644 demos/rgliner-web/src/components/tabs/style.css diff --git a/Cargo.lock b/Cargo.lock index 6e4b0c281..f2c6180e0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -587,7 +587,7 @@ dependencies = [ "derive_builder 0.20.2", "diligent-date-parser", "never", - "quick-xml", + "quick-xml 0.37.5", ] [[package]] @@ -1001,6 +1001,31 @@ dependencies = [ "cipher", ] +[[package]] +name = "bon" +version = "3.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f47dbe92550676ee653353c310dfb9cf6ba17ee70396e1f7cf0a2020ad49b2fe" +dependencies = [ + "bon-macros", + "rustversion", +] + +[[package]] +name = "bon-macros" +version = "3.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "519bd3116aeeb42d5372c29d982d16d0170d3d4a5ed85fc7dd91642ffff3c67c" +dependencies = [ + "darling 0.20.11", + "ident_case", + "prettyplease", + "proc-macro2", + "quote", + "rustversion", + "syn 2.0.117", +] + [[package]] name = "borsh" version = "1.6.0" @@ -2634,6 +2659,17 @@ dependencies = [ "web-sys", ] +[[package]] +name = "dioxus-attributes" +version = "0.1.0" +source = "git+https://github.com/DioxusLabs/components#ccdb07f69383de008a0afadda0e5ab7ec14c1a9c" +dependencies = [ + "dioxus-rsx", + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "dioxus-cli-config" version = "0.7.4" @@ -2742,7 +2778,7 @@ dependencies = [ "futures-channel", "futures-util", "generational-box", - "lazy-js-bundle", + "lazy-js-bundle 0.7.4", "serde", "serde_json", "tracing", @@ -2892,7 +2928,7 @@ dependencies = [ "futures-util", "generational-box", "keyboard-types", - "lazy-js-bundle", + "lazy-js-bundle 0.7.4", "rustversion", "tracing", ] @@ -2916,7 +2952,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a8ce1cf487007f90d0ec4ec87dff111d74ac04fca0918f9dcc4e80dc3b0531b2" dependencies = [ "js-sys", - "lazy-js-bundle", + "lazy-js-bundle 0.7.4", "rustc-hash 2.1.1", "sledgehammer_bindgen", "sledgehammer_utils", @@ -2937,6 +2973,30 @@ dependencies = [ "tracing-wasm", ] +[[package]] +name = "dioxus-markdown" +version = "0.1.0" +source = "git+https://github.com/rambip/rust-web-markdown#22ab22566014a8bd5bac959dd2a4770d2eddf16b" +dependencies = [ + "dioxus", + "web-framework-markdown", +] + +[[package]] +name = "dioxus-primitives" +version = "0.0.1" +source = "git+https://github.com/DioxusLabs/components#ccdb07f69383de008a0afadda0e5ab7ec14c1a9c" +dependencies = [ + "dioxus", + "dioxus-attributes", + "dioxus-sdk-time", + "lazy-js-bundle 0.6.2", + "num-integer", + "serde", + "time", + "tracing", +] + [[package]] name = "dioxus-rsx" version = "0.7.4" @@ -2950,6 +3010,18 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "dioxus-sdk-time" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "80c25ae93a3f72e734873b97fbd09d9b1b6adff97205fb0ffd8543e3564fb78e" +dependencies = [ + "dioxus", + "futures", + "gloo-timers", + "tokio", +] + [[package]] name = "dioxus-signals" version = "0.7.4" @@ -3010,7 +3082,7 @@ dependencies = [ "generational-box", "gloo-timers", "js-sys", - "lazy-js-bundle", + "lazy-js-bundle 0.7.4", "rustc-hash 2.1.1", "send_wrapper", "serde", @@ -3492,6 +3564,17 @@ dependencies = [ "regex-syntax 0.8.10", ] +[[package]] +name = "fancy-regex" +version = "0.16.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "998b056554fbe42e03ae0e152895cd1a7e1002aec800fdc6635d20270260c46f" +dependencies = [ + "bit-set 0.8.0", + "regex-automata 0.4.14", + "regex-syntax 0.8.10", +] + [[package]] name = "fancy-regex" version = "0.17.0" @@ -5676,7 +5759,7 @@ dependencies = [ "kalosm-sample", "kalosm-streams", "lopdf", - "pulldown-cmark", + "pulldown-cmark 0.9.6", "rand 0.8.5", "rbert", "readability", @@ -5902,6 +5985,24 @@ dependencies = [ name = "kalosm-workspace" version = "0.4.0" +[[package]] +name = "katex-rs" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c5382fea1e8edf972c23050cdfec2d12beca242e6d15e32310e5d1543a51d103" +dependencies = [ + "bon", + "phf 0.13.1", + "phf_codegen 0.13.1", + "rapidhash", + "serde", + "serde_json", + "strum 0.28.0", + "strum_macros 0.28.0", + "thiserror 2.0.18", + "unicode-normalization", +] + [[package]] name = "keyboard-types" version = "0.7.0" @@ -5959,6 +6060,12 @@ dependencies = [ "regex-automata 0.4.14", ] +[[package]] +name = "lazy-js-bundle" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e49596223b9d9d4947a14a25c142a6e7d8ab3f27eb3ade269d238bb8b5c267e2" + [[package]] name = "lazy-js-bundle" version = "0.7.4" @@ -6069,6 +6176,12 @@ dependencies = [ "thiserror 1.0.69", ] +[[package]] +name = "linked-hash-map" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0717cef1bc8b636c6e1c1bbdefc09e6322da8a9321966e8928ef80d20f7f770f" + [[package]] name = "linux-raw-sys" version = "0.12.1" @@ -6970,6 +7083,15 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "num_threads" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c7398b9c8b70908f6371f47ed36737907c87c52af34c268fed0bf0ceb92ead9" +dependencies = [ + "libc", +] + [[package]] name = "number_prefix" version = "0.4.0" @@ -7465,10 +7587,21 @@ version = "0.11.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1fd6780a80ae0c52cc120a26a1a42c1ae51b247a253e4e06113d23d2c2edd078" dependencies = [ - "phf_macros", + "phf_macros 0.11.3", "phf_shared 0.11.3", ] +[[package]] +name = "phf" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c1562dc717473dbaa4c1f85a36410e03c047b2e7df7f45ee938fbef64ae7fadf" +dependencies = [ + "phf_macros 0.13.1", + "phf_shared 0.13.1", + "serde", +] + [[package]] name = "phf_codegen" version = "0.10.0" @@ -7489,6 +7622,16 @@ dependencies = [ "phf_shared 0.11.3", ] +[[package]] +name = "phf_codegen" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "49aa7f9d80421bca176ca8dbfebe668cc7a2684708594ec9f3c0db0805d5d6e1" +dependencies = [ + "phf_generator 0.13.1", + "phf_shared 0.13.1", +] + [[package]] name = "phf_generator" version = "0.10.0" @@ -7509,6 +7652,16 @@ dependencies = [ "rand 0.8.5", ] +[[package]] +name = "phf_generator" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "135ace3a761e564ec88c03a77317a7c6b80bb7f7135ef2544dbe054243b89737" +dependencies = [ + "fastrand", + "phf_shared 0.13.1", +] + [[package]] name = "phf_macros" version = "0.11.3" @@ -7523,6 +7676,19 @@ dependencies = [ "unicase", ] +[[package]] +name = "phf_macros" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "812f032b54b1e759ccd5f8b6677695d5268c588701effba24601f6932f8269ef" +dependencies = [ + "phf_generator 0.13.1", + "phf_shared 0.13.1", + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "phf_shared" version = "0.10.0" @@ -7542,6 +7708,15 @@ dependencies = [ "unicase", ] +[[package]] +name = "phf_shared" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e57fef6bc5981e38c2ce2d63bfa546861309f875b8a75f092d1d54ae2d64f266" +dependencies = [ + "siphasher 1.0.2", +] + [[package]] name = "pico-args" version = "0.5.0" @@ -7592,6 +7767,19 @@ version = "0.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b4596b6d070b27117e987119b4dac604f3c58cfb0b191112e24771b2faeac1a6" +[[package]] +name = "plist" +version = "1.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "740ebea15c5d1428f910cd1a5f52cebf8d25006245ed8ade92702f4943d91e07" +dependencies = [ + "base64 0.22.1", + "indexmap 2.13.0", + "quick-xml 0.38.4", + "serde", + "time", +] + [[package]] name = "plotters" version = "0.3.7" @@ -7888,6 +8076,25 @@ dependencies = [ "unicase", ] +[[package]] +name = "pulldown-cmark" +version = "0.13.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c3a14896dfa883796f1cb410461aef38810ea05f2b2c33c5aded3649095fdad" +dependencies = [ + "bitflags 2.11.0", + "getopts", + "memchr", + "pulldown-cmark-escape", + "unicase", +] + +[[package]] +name = "pulldown-cmark-escape" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "007d8adb5ddab6f8e3f491ac63566a7d5002cc7ed73901f72057943fa71ae1ae" + [[package]] name = "pulp" version = "0.18.22" @@ -7968,6 +8175,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "quick-xml" +version = "0.38.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b66c2058c55a409d601666cffe35f04333cf1013010882cec174a7467cd4e21c" +dependencies = [ + "memchr", +] + [[package]] name = "quick_cache" version = "0.5.2" @@ -8166,6 +8382,15 @@ version = "1.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "973443cf09a9c8656b574a866ab68dfa19f0867d0340648c7d2f6a71b8a8ea68" +[[package]] +name = "rapidhash" +version = "4.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b5e48930979c155e2f33aa36ab3119b5ee81332beb6482199a8ecd6029b80b59" +dependencies = [ + "rustversion", +] + [[package]] name = "rav1e" version = "0.8.1" @@ -8651,7 +8876,10 @@ version = "0.1.0" dependencies = [ "console_error_panic_hook", "dioxus", + "dioxus-markdown", + "dioxus-primitives", "getrandom 0.3.4", + "gloo-timers", "rgliner", "tracing", "tracing-wasm", @@ -8768,7 +8996,7 @@ dependencies = [ "atom_syndication", "derive_builder 0.20.2", "never", - "quick-xml", + "quick-xml 0.37.5", ] [[package]] @@ -9781,6 +10009,15 @@ dependencies = [ "strum_macros 0.27.2", ] +[[package]] +name = "strum" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9628de9b8791db39ceda2b119bbe13134770b56c138ec1d3af810d045c04f9bd" +dependencies = [ + "strum_macros 0.28.0", +] + [[package]] name = "strum_macros" version = "0.26.4" @@ -9806,6 +10043,18 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "strum_macros" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab85eea0270ee17587ed4156089e10b9e6880ee688791d45a905f5b1ca36f664" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "subsecond" version = "0.7.4" @@ -10099,6 +10348,27 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "syntect" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "656b45c05d95a5704399aeef6bd0ddec7b2b3531b7c9e900abbf7c4d2190c925" +dependencies = [ + "bincode", + "fancy-regex 0.16.2", + "flate2", + "fnv", + "once_cell", + "plist", + "regex-syntax 0.8.10", + "serde", + "serde_derive", + "serde_json", + "thiserror 2.0.18", + "walkdir", + "yaml-rust", +] + [[package]] name = "sysctl" version = "0.5.5" @@ -10368,7 +10638,9 @@ checksum = "743bd48c283afc0388f9b8827b976905fb217ad9e647fae3a379a9283c4def2c" dependencies = [ "deranged", "itoa", + "libc", "num-conv", + "num_threads", "powerfmt", "serde_core", "time-core", @@ -11416,6 +11688,19 @@ dependencies = [ "pkg-config", ] +[[package]] +name = "web-framework-markdown" +version = "0.1.0" +source = "git+https://github.com/rambip/rust-web-markdown#22ab22566014a8bd5bac959dd2a4770d2eddf16b" +dependencies = [ + "katex-rs", + "lazy_static", + "pulldown-cmark 0.13.3", + "regex", + "syntect", + "web-sys", +] + [[package]] name = "web-sys" version = "0.3.91" @@ -12361,6 +12646,15 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7a5a4b21e1a62b67a2970e6831bc091d7b87e119e7f9791aef9702e3bef04448" +[[package]] +name = "yaml-rust" +version = "0.4.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "56c1936c4cc7a1c9ab21a1ebb602eb942ba868cbd44a99cb7cdc5892335e1c85" +dependencies = [ + "linked-hash-map", +] + [[package]] name = "yansi" version = "1.0.1" diff --git a/demos/rgliner-web/Cargo.toml b/demos/rgliner-web/Cargo.toml index f362cb1e1..f8334b9f1 100644 --- a/demos/rgliner-web/Cargo.toml +++ b/demos/rgliner-web/Cargo.toml @@ -10,8 +10,10 @@ rgliner = { path = "../../models/rgliner", default-features = false } getrandom = { version = "0.3", features = ["wasm_js"] } tracing = "0.1" dioxus-primitives = { git = "https://github.com/DioxusLabs/components", version = "0.0.1", default-features = false } +dioxus-markdown = { git = "https://github.com/rambip/rust-web-markdown" } [target.'cfg(target_arch = "wasm32")'.dependencies] console_error_panic_hook = "0.1" tracing-wasm = "0.2" wasm-bindgen-futures = "0.4" +gloo-timers = { version = "0.3", features = ["futures"] } diff --git a/demos/rgliner-web/assets/style.css b/demos/rgliner-web/assets/style.css index 0cc1f6fe2..1a2527e9f 100644 --- a/demos/rgliner-web/assets/style.css +++ b/demos/rgliner-web/assets/style.css @@ -1,234 +1,479 @@ -:root { - --bg: #0b0d12; - --panel: #141821; - --panel-border: #232833; - --text: #e6e9ef; - --muted: #8b92a3; - --accent: #7aa7ff; - --accent-strong: #4a7fff; - --err: #ff6b6b; - --ok: #4ade80; - --radius: 8px; -} +/* rgliner — a minimal reader + typography · paper · quiet annotation */ -* { - box-sizing: border-box; -} +@import url("https://fonts.googleapis.com/css2?family=Fraunces:ital,opsz,wght@0,9..144,300..900;1,9..144,300..900&family=JetBrains+Mono:wght@400;500&display=swap"); -body { +:root { + --paper: #f6f1e4; + --paper-deep: #efe7d2; + --rule: #d8cdb0; + --ink: #14110d; + --ink-soft: #3b342a; + --ink-faint: #8a8071; + --accent: #b6391a; + --ok: #3d6b55; + + /* Overrides for dx-components */ + --primary-color: var(--paper); + --primary-color-1: var(--paper); + --primary-color-2: var(--paper); + --primary-color-3: var(--paper-deep); + --primary-color-4: var(--paper-deep); + --primary-color-5: var(--paper-deep); + --primary-color-6: var(--rule); + --primary-color-7: var(--ink-faint); + --secondary-color: var(--ink); + --secondary-color-1: var(--ink); + --secondary-color-2: var(--ink); + --secondary-color-3: var(--ink); + --secondary-color-4: var(--ink); + --secondary-color-5: var(--ink-soft); + --secondary-color-6: var(--rule); + --focused-border-color: var(--accent); + + --serif: "Fraunces", "Iowan Old Style", Georgia, serif; + --mono: "JetBrains Mono", ui-monospace, SFMono-Regular, Menlo, monospace; +} + +* { box-sizing: border-box; } + +html, body { margin: 0; - font-family: ui-sans-serif, system-ui, -apple-system, "Segoe UI", Roboto, sans-serif; - background: var(--bg); - color: var(--text); - line-height: 1.5; + background: var(--paper); + color: var(--ink); + font-family: var(--serif); + font-size: 17px; + line-height: 1.65; + -webkit-font-smoothing: antialiased; } -.app { - max-width: 960px; +::selection { background: var(--accent); color: var(--paper); } + +/* ── Layout ─────────────────────────────────────────────────── */ + +.reader { + max-width: 1200px; margin: 0 auto; - padding: 2rem 1.5rem 4rem; + padding: 3rem 2rem 6rem; } -header.site-header { +/* ── Masthead ──────────────────────────────────────────────── */ + +.masthead { display: flex; align-items: baseline; justify-content: space-between; - gap: 1rem; - margin-bottom: 1.5rem; + gap: 2rem; + padding-bottom: 1.25rem; + border-bottom: 1px solid var(--ink); + margin-bottom: 0.75rem; + flex-wrap: wrap; } -header.site-header h1 { - font-size: 1.5rem; - margin: 0; +.wordmark { + display: flex; + align-items: baseline; + gap: 1rem; + flex-wrap: wrap; } -header.site-header .tag { - color: var(--muted); - font-size: 0.9rem; +.wordmark .mark { + font-family: var(--serif); + font-style: italic; + font-weight: 400; + font-variation-settings: "opsz" 144, "SOFT" 100; + font-size: clamp(2.75rem, 5vw, 3.75rem); + line-height: 1; + letter-spacing: -0.03em; + color: var(--ink); } +.wordmark .mark::first-letter { color: var(--accent); } -header.site-header a { - color: var(--accent); - text-decoration: none; - font-size: 0.9rem; +.wordmark .byline { + font-family: var(--serif); + font-style: italic; + font-weight: 350; + font-size: 1rem; + color: var(--ink-faint); } -.tabs { +.picker { display: flex; - gap: 0.25rem; - border-bottom: 1px solid var(--panel-border); - margin-bottom: 1.5rem; + gap: 0.6rem; + align-items: center; } -.tab { - background: none; +button.load { + background: var(--ink); + color: var(--paper); border: none; - color: var(--muted); - padding: 0.6rem 1rem; + border-radius: 0; + font-family: var(--mono); + font-size: 0.72rem; + letter-spacing: 0.18em; + text-transform: uppercase; + padding: 0.55rem 0.9rem; cursor: pointer; - font-size: 0.95rem; - border-bottom: 2px solid transparent; + transition: background 0.15s; +} +button.load:hover:not(:disabled) { background: var(--accent); } +button.load:disabled { + background: transparent; + color: var(--ink-faint); + cursor: default; + border: 1px solid var(--rule); } -.tab:hover { - color: var(--text); +/* dx select, tuned to feel like an editorial dropdown */ +[class*="select-trigger"], [class*="selectTrigger"] { + background: transparent !important; + border: none !important; + border-bottom: 1px solid var(--ink) !important; + border-radius: 0 !important; + color: var(--ink) !important; + font-family: var(--serif) !important; + font-size: 1rem !important; + font-style: italic; + padding: 0.3rem 0.5rem 0.3rem 0 !important; + min-width: 18rem; } -.tab.active { - color: var(--text); - border-bottom-color: var(--accent); +[class*="select-list"], [class*="selectList"], [role="listbox"] { + background: var(--paper) !important; + border: 1px solid var(--ink) !important; + border-radius: 0 !important; + box-shadow: 4px 4px 0 var(--ink) !important; + padding: 0 !important; + font-family: var(--mono); } -.panel { - background: var(--panel); - border: 1px solid var(--panel-border); - border-radius: var(--radius); - padding: 1rem 1.25rem; - margin-bottom: 1rem; +[role="option"], [class*="select-option"] { + font-family: var(--mono) !important; + font-size: 0.8rem !important; + padding: 0.55rem 0.85rem !important; + border-bottom: 1px solid var(--rule); + color: var(--ink) !important; + border-radius: 0 !important; +} +[role="option"]:last-child { border-bottom: none; } +[role="option"]:hover, [role="option"][data-highlighted] { + background: var(--ink) !important; + color: var(--paper) !important; } -.row { - display: flex; - gap: 0.75rem; +/* ── Settings disclosure ───────────────────────────────────── */ + +details.settings { + border-bottom: 1px solid var(--rule); + padding: 0.5rem 0 0.9rem; + margin-bottom: 0.75rem; +} + +details.settings > summary { + list-style: none; + cursor: pointer; + font-family: var(--mono); + font-size: 0.72rem; + letter-spacing: 0.18em; + text-transform: uppercase; + color: var(--ink-faint); + user-select: none; + display: inline-flex; + gap: 0.5rem; align-items: center; - flex-wrap: wrap; } +details.settings > summary::-webkit-details-marker { display: none; } +details.settings > summary::before { + content: "▸"; + transition: transform 0.15s; + color: var(--ink-faint); +} +details.settings[open] > summary::before { transform: rotate(90deg); } +details.settings > summary:hover { color: var(--ink); } -label { - color: var(--muted); - font-size: 0.85rem; - display: block; - margin-bottom: 0.35rem; -} - -input[type="text"], -textarea, -select { - background: #0f1218; - color: var(--text); - border: 1px solid var(--panel-border); - border-radius: 6px; - padding: 0.55rem 0.7rem; - font-family: inherit; - font-size: 0.95rem; - width: 100%; -} - -textarea { - min-height: 110px; - resize: vertical; - font-family: ui-monospace, SFMono-Regular, "SF Mono", Menlo, monospace; - font-size: 0.9rem; +.settings-body { + margin-top: 0.9rem; + max-width: 560px; } -select { - width: auto; - min-width: 220px; +label, [class*="label"] { + font-family: var(--mono) !important; + font-size: 0.68rem !important; + letter-spacing: 0.18em !important; + text-transform: uppercase !important; + color: var(--ink-faint) !important; + margin-bottom: 0.35rem !important; + display: block !important; + font-weight: 500 !important; } -button.primary, -button.secondary { - background: var(--accent-strong); - color: white; - border: none; - border-radius: 6px; - padding: 0.55rem 1rem; - cursor: pointer; - font-size: 0.95rem; - font-weight: 500; +input[type="text"], [class*="input"] { + background: transparent !important; + color: var(--ink) !important; + border: none !important; + border-bottom: 1px solid var(--rule) !important; + border-radius: 0 !important; + padding: 0.5rem 0 !important; + font-family: var(--mono) !important; + font-size: 0.88rem !important; + width: 100% !important; +} +input:focus, [class*="input"]:focus, +input:focus-visible { + outline: none !important; + border-bottom-color: var(--accent) !important; } -button.primary:hover:not(:disabled) { +/* ── Status line ───────────────────────────────────────────── */ + +.status-line { + display: flex; + align-items: center; + gap: 0.5rem; + font-family: var(--mono); + font-size: 0.7rem; + letter-spacing: 0.14em; + text-transform: uppercase; + color: var(--ink-faint); + padding: 0.4rem 0; + margin-bottom: 1.5rem; +} +.status-line .dot { + width: 0.55rem; + height: 0.55rem; + border-radius: 50%; + background: var(--ink-faint); + display: inline-block; +} +.status-line .dot.ok { background: var(--ok); } +.status-line .dot.err { background: var(--accent); } +.status-line .dot.busy { background: var(--accent); + animation: pulse 1.1s ease-in-out infinite; +} +.status-line .err-text { color: var(--accent); letter-spacing: 0.04em; text-transform: none; } + +@keyframes pulse { + 0%, 100% { opacity: 0.35; transform: scale(0.85); } + 50% { opacity: 1; transform: scale(1.05); } } -button.secondary { +/* ── Split editor / reader ─────────────────────────────────── */ + +.split { + display: grid; + grid-template-columns: 1fr 1fr; + gap: 3rem; + align-items: stretch; + min-height: 70vh; +} + +textarea.editor { background: transparent; - color: var(--text); - border: 1px solid var(--panel-border); + color: var(--ink-soft); + border: none; + border-right: 1px dashed var(--rule); + padding: 0.5rem 2rem 0.5rem 0; + font-family: var(--mono); + font-size: 0.9rem; + line-height: 1.75; + resize: none; + outline: none; + min-height: 60vh; } +textarea.editor:focus { color: var(--ink); } -button:disabled { - opacity: 0.5; - cursor: not-allowed; +/* ── Article ───────────────────────────────────────────────── */ + +.article { + font-family: var(--serif); + font-size: 1.15rem; + line-height: 1.75; + color: var(--ink); + padding: 0.5rem 0; } -.status { - font-size: 0.85rem; - color: var(--muted); +.article .placeholder { + font-style: italic; + color: var(--ink-faint); + padding: 2rem 0; } -.status.ok { - color: var(--ok); +.article h1 { + font-family: var(--serif); + font-style: italic; + font-weight: 500; + font-size: 2.2rem; + line-height: 1.1; + letter-spacing: -0.02em; + margin: 0 0 1.25rem; + color: var(--ink); } +.article h1::first-letter { color: var(--accent); } -.status.err { - color: var(--err); +.article h2 { + font-family: var(--serif); + font-weight: 500; + font-size: 1.4rem; + line-height: 1.2; + margin: 2rem 0 0.75rem; + letter-spacing: -0.01em; } -.err-banner { - background: rgba(255, 107, 107, 0.08); - border: 1px solid rgba(255, 107, 107, 0.4); - color: #ffb3b3; - padding: 0.75rem 1rem; - border-radius: 6px; - margin-bottom: 1rem; - font-size: 0.9rem; +.article h3 { + font-family: var(--mono); + font-size: 0.72rem; + letter-spacing: 0.22em; + text-transform: uppercase; + color: var(--ink-faint); + margin: 1.75rem 0 0.5rem; } -.results { - line-height: 1.9; - font-size: 1rem; - word-wrap: break-word; +.article p { margin: 0 0 1.15rem; } + +.article em { font-style: italic; } +.article strong { font-weight: 600; } + +.article a { + color: var(--ink); + text-decoration: none; + background-image: linear-gradient(to top, var(--accent) 0.08em, transparent 0.08em); + background-repeat: no-repeat; + background-size: 100% 100%; } +.article a:hover { color: var(--accent); } -.entity { - border-radius: 4px; +.article ul, .article ol { padding-left: 1.25rem; margin: 0 0 1.15rem; } +.article li { margin-bottom: 0.25rem; } + +.article blockquote { + border-left: 2px solid var(--accent); + padding: 0.1rem 0 0.1rem 1rem; + margin: 1rem 0; + color: var(--ink-soft); + font-style: italic; +} + +.article code { + font-family: var(--mono); + font-size: 0.85em; + background: var(--paper-deep); padding: 0.05em 0.3em; - margin: 0 0.05em; - font-weight: 500; - color: #0b0d12; + border-radius: 2px; } -.entity .chip { - font-size: 0.7em; - font-weight: 700; - padding: 0 0.35em; - margin-left: 0.3em; - border-radius: 3px; - background: rgba(0, 0, 0, 0.25); - color: #fff; - text-transform: uppercase; - letter-spacing: 0.04em; +.article hr { + border: none; + text-align: center; + margin: 2rem 0; } +.article hr::after { + content: "❦"; + color: var(--ink-faint); + font-size: 1.2rem; + letter-spacing: 1em; +} + +/* ── Entity annotation ─────────────────────────────────────── */ -.entity-list, -.relation-list { - margin-top: 1rem; +.entity { + position: relative; + cursor: default; + padding: 0 0.05em; + text-decoration: underline; + text-decoration-color: var(--ec, var(--accent)); + text-decoration-thickness: 0.14em; + text-underline-offset: 0.18em; + transition: background 0.12s; +} +.entity:hover { + background: color-mix(in srgb, var(--ec, var(--accent)) 14%, transparent); +} + +.entity .pop { + position: absolute; + left: 50%; + top: calc(100% + 0.4rem); + transform: translateX(-50%) translateY(-4px); + min-width: 16rem; + max-width: 22rem; + background: var(--ink); + color: var(--paper); + padding: 0.65rem 0.8rem; + font-family: var(--mono); + font-size: 0.75rem; + line-height: 1.5; + letter-spacing: 0.02em; + text-transform: none; + text-decoration: none; + box-shadow: 5px 5px 0 var(--accent); + opacity: 0; + pointer-events: none; + transition: opacity 0.12s, transform 0.12s; + z-index: 10; display: flex; flex-direction: column; - gap: 0.35rem; + gap: 0.3rem; } - -.entity-list li, -.relation-list li { - list-style: none; - font-size: 0.9rem; - color: var(--muted); - font-family: ui-monospace, SFMono-Regular, "SF Mono", Menlo, monospace; +.entity:hover .pop { + opacity: 1; + transform: translateX(-50%) translateY(0); } -.entity-list .label, -.relation-list .label { - color: var(--text); - font-weight: 600; +.entity .pop::before { + content: ""; + position: absolute; + top: -5px; + left: 50%; + width: 10px; + height: 10px; + background: var(--ink); + transform: translateX(-50%) rotate(45deg); } -.score { - color: var(--accent); +.pop-label { + font-size: 0.6rem; + letter-spacing: 0.24em; + text-transform: uppercase; + color: var(--ec, var(--accent)); + font-weight: 500; + padding-bottom: 0.25rem; + border-bottom: 1px dashed rgba(255,255,255,0.15); + margin-bottom: 0.15rem; } -.muted { - color: var(--muted); +.pop-rel { + display: block; + font-size: 0.72rem; + color: var(--paper); + padding: 0.15rem 0; +} +.pop-rel .arrow { color: rgba(255,255,255,0.35); } +.pop-rel .rel-name { + color: var(--paper); + font-style: italic; + font-family: var(--serif); font-size: 0.85rem; } + +.pop-empty { + color: rgba(255,255,255,0.5); + font-style: italic; +} + +/* ── Responsive ────────────────────────────────────────────── */ + +@media (max-width: 820px) { + .reader { padding: 1.5rem 1.25rem 4rem; } + .split { + grid-template-columns: 1fr; + gap: 1.5rem; + } + textarea.editor { + border-right: none; + border-bottom: 1px dashed var(--rule); + padding: 0 0 1rem; + min-height: 30vh; + } + [class*="select-trigger"], [class*="selectTrigger"] { min-width: 0 !important; } +} diff --git a/demos/rgliner-web/src/components/badge/component.rs b/demos/rgliner-web/src/components/badge/component.rs deleted file mode 100644 index 831150562..000000000 --- a/demos/rgliner-web/src/components/badge/component.rs +++ /dev/null @@ -1,82 +0,0 @@ -use dioxus::prelude::*; - -#[derive(Copy, Clone, PartialEq, Default)] -#[non_exhaustive] -pub enum BadgeVariant { - #[default] - Primary, - Secondary, - Destructive, - Outline, -} - -impl BadgeVariant { - pub fn class(&self) -> &'static str { - match self { - BadgeVariant::Primary => "primary", - BadgeVariant::Secondary => "secondary", - BadgeVariant::Destructive => "destructive", - BadgeVariant::Outline => "outline", - } - } -} - -/// The props for the [`Badge`] component. -#[derive(Props, Clone, PartialEq)] -pub struct BadgeProps { - #[props(default)] - pub variant: BadgeVariant, - - /// Additional attributes to extend the badge element - #[props(extends = GlobalAttributes)] - pub attributes: Vec, - - /// The children of the badge element - pub children: Element, -} - -#[component] -pub fn Badge(props: BadgeProps) -> Element { - rsx! { - document::Link { rel: "stylesheet", href: asset!("./style.css") } - - BadgeElement { - "padding": true, - variant: props.variant, - attributes: props.attributes, - {props.children} - } - } -} - -#[component] -fn BadgeElement(props: BadgeProps) -> Element { - rsx! { - span { - class: "badge", - "data-style": props.variant.class(), - ..props.attributes, - {props.children} - } - } -} - -#[component] -pub fn VerifiedIcon() -> Element { - rsx! { - // Badge icon from lucide https://lucide.dev/icons/badge - svg { - view_box: "0 0 24 24", - xmlns: "http://www.w3.org/2000/svg", - width: "12", - height: "12", - fill: "none", - stroke: "var(--secondary-color-4)", - stroke_linecap: "round", - stroke_linejoin: "round", - stroke_width: 2, - path { d: "M3.85 8.62a4 4 0 0 1 4.78-4.77 4 4 0 0 1 6.74 0 4 4 0 0 1 4.78 4.78 4 4 0 0 1 0 6.74 4 4 0 0 1-4.77 4.78 4 4 0 0 1-6.75 0 4 4 0 0 1-4.78-4.77 4 4 0 0 1 0-6.76Z" } - path { d: "m9 12 2 2 4-4" } - } - } -} diff --git a/demos/rgliner-web/src/components/badge/mod.rs b/demos/rgliner-web/src/components/badge/mod.rs deleted file mode 100644 index 9a8ae5565..000000000 --- a/demos/rgliner-web/src/components/badge/mod.rs +++ /dev/null @@ -1,2 +0,0 @@ -mod component; -pub use component::*; \ No newline at end of file diff --git a/demos/rgliner-web/src/components/badge/style.css b/demos/rgliner-web/src/components/badge/style.css deleted file mode 100644 index e36df538e..000000000 --- a/demos/rgliner-web/src/components/badge/style.css +++ /dev/null @@ -1,42 +0,0 @@ -.badge-example { - display: flex; - align-items: center; - gap: 1rem; -} - -.badge { - display: inline-flex; - min-width: 20px; - height: 20px; - align-items: center; - justify-content: center; - border-radius: 10px; - box-shadow: 0 0 0 1px var(--primary-color-2); - font-size: 12px; - gap: 4px -} - -.badge[padding="true"] { - padding: 0 8px; -} - -.badge[data-style="primary"] { - background-color: var(--secondary-color-2); - color: var(--primary-color); -} - -.badge[data-style="secondary"] { - background-color: var(--primary-color-5); - color: var(--secondary-color-1); -} - -.badge[data-style="outline"] { - border: 1px solid var(--primary-color-6); - background-color: var(--light, var(--primary-color)) var(--dark, var(--primary-color-3)); - color: var(--secondary-color-4); -} - -.badge[data-style="destructive"] { - background-color: var(--primary-error-color); - color: var(--contrast-error-color); -} \ No newline at end of file diff --git a/demos/rgliner-web/src/components/card/component.rs b/demos/rgliner-web/src/components/card/component.rs deleted file mode 100644 index 036749a7a..000000000 --- a/demos/rgliner-web/src/components/card/component.rs +++ /dev/null @@ -1,107 +0,0 @@ -use dioxus::prelude::*; - -#[component] -pub fn Card( - #[props(extends=GlobalAttributes)] attributes: Vec, - children: Element, -) -> Element { - rsx! { - document::Link { rel: "stylesheet", href: asset!("./style.css") } - div { - class: "card", - "data-slot": "card", - ..attributes, - {children} - } - } -} - -#[component] -pub fn CardHeader( - #[props(extends=GlobalAttributes)] attributes: Vec, - children: Element, -) -> Element { - rsx! { - div { - class: "card-header", - "data-slot": "card-header", - ..attributes, - {children} - } - } -} - -#[component] -pub fn CardTitle( - #[props(extends=GlobalAttributes)] attributes: Vec, - children: Element, -) -> Element { - rsx! { - div { - class: "card-title", - "data-slot": "card-title", - ..attributes, - {children} - } - } -} - -#[component] -pub fn CardDescription( - #[props(extends=GlobalAttributes)] attributes: Vec, - children: Element, -) -> Element { - rsx! { - div { - class: "card-description", - "data-slot": "card-description", - ..attributes, - {children} - } - } -} - -#[component] -pub fn CardAction( - #[props(extends=GlobalAttributes)] attributes: Vec, - children: Element, -) -> Element { - rsx! { - div { - class: "card-action", - "data-slot": "card-action", - ..attributes, - {children} - } - } -} - -#[component] -pub fn CardContent( - #[props(extends=GlobalAttributes)] attributes: Vec, - children: Element, -) -> Element { - rsx! { - div { - class: "card-content", - "data-slot": "card-content", - ..attributes, - {children} - } - } -} - -#[component] -pub fn CardFooter( - #[props(extends=GlobalAttributes)] attributes: Vec, - children: Element, -) -> Element { - rsx! { - div { - class: "card-footer", - "data-slot": "card-footer", - ..attributes, - {children} - } - } -} diff --git a/demos/rgliner-web/src/components/card/mod.rs b/demos/rgliner-web/src/components/card/mod.rs deleted file mode 100644 index a3527a11b..000000000 --- a/demos/rgliner-web/src/components/card/mod.rs +++ /dev/null @@ -1,3 +0,0 @@ -mod component; -pub use component::*; - diff --git a/demos/rgliner-web/src/components/card/style.css b/demos/rgliner-web/src/components/card/style.css deleted file mode 100644 index 2ad7e6e6a..000000000 --- a/demos/rgliner-web/src/components/card/style.css +++ /dev/null @@ -1,52 +0,0 @@ -.card { - display: flex; - flex-direction: column; - padding: 1.5rem 0; - border: 1px solid var(--light, var(--primary-color-6)) var(--dark, var(--primary-color-5)); - border-radius: 1rem; - background-color: var(--light, var(--primary-color-2)) var(--dark, var(--primary-color-3)); - box-shadow: 0 2px 10px rgb(0 0 0 / 10%); - color: var(--secondary-color-4); - gap: 1.5rem; -} - -.card-header { - display: grid; - align-items: start; - padding: 0 1.5rem; - gap: 0.5rem; - grid-auto-rows: min-content; - grid-template-rows: auto auto; -} - -.card-header:has([data-slot="card-action"]) { - grid-template-columns: 1fr auto; -} - -.card-title { - font-size: 1rem; - font-weight: 600; - line-height: 1; -} - -.card-description { - color: var(--secondary-color-5); - font-size: 0.875rem; - line-height: 1.25rem; -} - -.card-action { - grid-column-start: 2; - grid-row: 1 / span 2; - place-self: start end; -} - -.card-content { - padding: 0 1.5rem; -} - -.card-footer { - display: flex; - align-items: center; - padding: 0 1.5rem; -} diff --git a/demos/rgliner-web/src/components/mod.rs b/demos/rgliner-web/src/components/mod.rs index da7dda254..d8297853f 100644 --- a/demos/rgliner-web/src/components/mod.rs +++ b/demos/rgliner-web/src/components/mod.rs @@ -1,9 +1,6 @@ // AUTOGENERTED Components module pub mod textarea; -pub mod tabs; pub mod label; pub mod button; -pub mod card; pub mod input; pub mod select; -pub mod badge; diff --git a/demos/rgliner-web/src/components/tabs/component.rs b/demos/rgliner-web/src/components/tabs/component.rs deleted file mode 100644 index 93ade03ad..000000000 --- a/demos/rgliner-web/src/components/tabs/component.rs +++ /dev/null @@ -1,119 +0,0 @@ -use dioxus::prelude::*; -use dioxus_primitives::tabs::{self, TabContentProps, TabListProps, TabTriggerProps}; - -/// The props for the [`Tabs`] component. -#[derive(Props, Clone, PartialEq)] -pub struct TabsProps { - /// The class of the tabs component. - #[props(default)] - pub class: String, - - /// The controlled value of the active tab. - pub value: ReadSignal>, - - /// The default active tab value when uncontrolled. - #[props(default)] - pub default_value: String, - - /// Callback fired when the active tab changes. - #[props(default)] - pub on_value_change: Callback, - - /// Whether the tabs are disabled. - #[props(default)] - pub disabled: ReadSignal, - - /// Whether the tabs are horizontal. - #[props(default)] - pub horizontal: ReadSignal, - - /// Whether focus should loop around when reaching the end. - #[props(default = ReadSignal::new(Signal::new(true)))] - pub roving_loop: ReadSignal, - - /// The variant of the tabs component. - #[props(default)] - pub variant: TabsVariant, - - /// Additional attributes to apply to the tabs element. - #[props(extends = GlobalAttributes)] - pub attributes: Vec, - - /// The children of the tabs component. - pub children: Element, -} - -/// The variant of the tabs component. -#[derive(Clone, Copy, PartialEq, Default)] -pub enum TabsVariant { - /// The default variant. - #[default] - Default, - /// The ghost variant. - Ghost, -} - -impl TabsVariant { - /// Convert the variant to a string for use in class names - fn to_class(self) -> &'static str { - match self { - TabsVariant::Default => "default", - TabsVariant::Ghost => "ghost", - } - } -} - -#[component] -pub fn Tabs(props: TabsProps) -> Element { - rsx! { - document::Link { rel: "stylesheet", href: asset!("./style.css") } - tabs::Tabs { - class: props.class + " tabs", - "data-variant": props.variant.to_class(), - value: props.value, - default_value: props.default_value, - on_value_change: props.on_value_change, - disabled: props.disabled, - horizontal: props.horizontal, - roving_loop: props.roving_loop, - attributes: props.attributes, - {props.children} - } - } -} - -#[component] -pub fn TabList(props: TabListProps) -> Element { - rsx! { - tabs::TabList { class: "tabs-list", attributes: props.attributes, {props.children} } - } -} - -#[component] -pub fn TabTrigger(props: TabTriggerProps) -> Element { - rsx! { - tabs::TabTrigger { - class: "tabs-trigger", - id: props.id, - value: props.value, - index: props.index, - disabled: props.disabled, - attributes: props.attributes, - {props.children} - } - } -} - -#[component] -pub fn TabContent(props: TabContentProps) -> Element { - rsx! { - tabs::TabContent { - class: props.class.unwrap_or_default() + " tabs-content tabs-content-themed", - value: props.value, - id: props.id, - index: props.index, - attributes: props.attributes, - {props.children} - } - } -} diff --git a/demos/rgliner-web/src/components/tabs/mod.rs b/demos/rgliner-web/src/components/tabs/mod.rs deleted file mode 100644 index 9a8ae5565..000000000 --- a/demos/rgliner-web/src/components/tabs/mod.rs +++ /dev/null @@ -1,2 +0,0 @@ -mod component; -pub use component::*; \ No newline at end of file diff --git a/demos/rgliner-web/src/components/tabs/style.css b/demos/rgliner-web/src/components/tabs/style.css deleted file mode 100644 index 5e04cd012..000000000 --- a/demos/rgliner-web/src/components/tabs/style.css +++ /dev/null @@ -1,72 +0,0 @@ -.tabs { - display: flex; - width: 100%; - flex-direction: column; - gap: 0.5rem; -} - -.tabs-list { - display: flex; - width: fit-content; - box-sizing: border-box; - flex: 1; - flex-direction: row; - padding: 0.25rem; - border: none; - border-radius: 0.5rem; - gap: 0.25rem; -} - -[data-variant="default"] .tabs-list { - background: var(--light, var(--primary-color-3)) - var(--dark, var(--primary-color-5)); -} - -.tabs-trigger { - padding: 4px 8px; - border: none; - border-radius: calc(0.5rem - 0.25rem); - background: none; - color: var(--secondary-color-5); - cursor: pointer; -} - -[data-variant="default"] .tabs-trigger[data-state="active"] { - background-color: var(--light, var(--primary-color)) - var(--dark, var(--primary-color-6)); - box-shadow: var(--dark, inset 0 0 0 1px var(--primary-color-7)) - var(--light, 0 1px 2px rgb(0 0 0 / 18%)); -} - -.tabs-trigger[data-state="active"] { - color: var(--secondary-color-1); -} - -.tabs-trigger[data-disabled="true"] { - color: var(--secondary-color-5); - cursor: not-allowed; -} - -.tabs-trigger:hover:not([data-disabled="true"]), -.tabs-trigger:focus-visible { - color: var(--secondary-color-3); -} - -.tabs-content { - width: 100%; - box-sizing: border-box; - padding: 0.25rem; -} - -[data-variant="default"] .tabs-content-themed { - border: 1px solid var(--light, var(--primary-color-6)) - var(--dark, var(--primary-color-7)); - border-radius: 0.5rem; - background: var(--light, var(--primary-color)) - var(--dark, var(--primary-color-3)); - box-shadow: var(--light, 0 1px 2px rgb(0 0 0 / 18%)) var(--dark, none); -} - -.tabs-content[data-state="inactive"] { - display: none; -} diff --git a/demos/rgliner-web/src/main.rs b/demos/rgliner-web/src/main.rs index dbae9fdcc..fd1a57c2f 100644 --- a/demos/rgliner-web/src/main.rs +++ b/demos/rgliner-web/src/main.rs @@ -1,16 +1,12 @@ mod components; -use components::badge::{Badge, BadgeVariant}; -use components::button::Button; -use components::card::{Card, CardContent, CardDescription, CardHeader, CardTitle}; use components::input::Input; use components::label::Label; use components::select::{ Select, SelectItemIndicator, SelectList, SelectOption, SelectTrigger, SelectValue, }; -use components::tabs::{TabContent, TabList, TabTrigger, Tabs}; -use components::textarea::Textarea; use dioxus::prelude::*; +use dioxus_markdown::Markdown; use rgliner::{ relation_decoding::Relation, relex::{GlinerRelEx, GlinerRelExSource}, @@ -18,6 +14,15 @@ use rgliner::{ }; const TOKEN_BUDGET: usize = 128; +const DEBOUNCE_MS: u32 = 500; + +#[cfg(target_arch = "wasm32")] +async fn sleep_ms(ms: u32) { + gloo_timers::future::TimeoutFuture::new(ms).await; +} + +#[cfg(not(target_arch = "wasm32"))] +async fn sleep_ms(_ms: u32) {} fn main() { #[cfg(target_arch = "wasm32")] @@ -34,22 +39,6 @@ enum Mode { Relex, } -impl Mode { - fn value(self) -> &'static str { - match self { - Mode::Ner => "ner", - Mode::Relex => "relex", - } - } - - fn from_value(v: &str) -> Mode { - match v { - "relex" => Mode::Relex, - _ => Mode::Ner, - } - } -} - #[derive(Clone, Copy, PartialEq, Eq)] enum ModelChoice { Edge, @@ -64,37 +53,35 @@ enum ModelChoice { impl ModelChoice { fn label(self) -> &'static str { match self { - ModelChoice::Edge => "edge · 60M · fastest", - ModelChoice::Small => "small · 108M", - ModelChoice::Base => "base · 194M", - ModelChoice::Large => "large · 530M · best NER", - ModelChoice::RelexMulti => "relex-multi · multilingual", - ModelChoice::RelexBase => "relex-base · English", - ModelChoice::RelexLarge => "relex-large · English · best", + ModelChoice::Edge => "edge · entities · 60M", + ModelChoice::Small => "small · entities · 108M", + ModelChoice::Base => "base · entities · 194M", + ModelChoice::Large => "large · entities · 530M", + ModelChoice::RelexMulti => "relex-multi · entities + relations", + ModelChoice::RelexBase => "relex-base · entities + relations · EN", + ModelChoice::RelexLarge => "relex-large · entities + relations · EN", } } - fn default_for(mode: Mode) -> Self { - match mode { - Mode::Ner => ModelChoice::Edge, - Mode::Relex => ModelChoice::RelexMulti, + fn mode(self) -> Mode { + match self { + ModelChoice::Edge | ModelChoice::Small | ModelChoice::Base | ModelChoice::Large => { + Mode::Ner + } + _ => Mode::Relex, } } - fn for_mode(mode: Mode) -> &'static [ModelChoice] { - match mode { - Mode::Ner => &[ - ModelChoice::Edge, - ModelChoice::Small, - ModelChoice::Base, - ModelChoice::Large, - ], - Mode::Relex => &[ - ModelChoice::RelexMulti, - ModelChoice::RelexBase, - ModelChoice::RelexLarge, - ], - } + fn all() -> &'static [ModelChoice] { + &[ + ModelChoice::Edge, + ModelChoice::Small, + ModelChoice::Base, + ModelChoice::Large, + ModelChoice::RelexMulti, + ModelChoice::RelexBase, + ModelChoice::RelexLarge, + ] } } @@ -112,39 +99,98 @@ impl LoadedModel { } #[derive(Clone, Default)] -struct RelexResult { +struct Extraction { entities: Vec, relations: Vec, } +const DEFAULT_TEXT: &str = "# Silicon Valley, Briefly\n\n*Apple Inc.* was founded by **Steve Jobs** in California. **Microsoft** is headquartered in Redmond, and was founded by Bill Gates.\n\nOpenAI operates out of San Francisco."; + #[component] fn App() -> Element { - let mut mode = use_signal(|| Mode::Ner); let mut choice = use_signal(|| ModelChoice::Edge); - - let mut text = use_signal(|| { - "Apple Inc. was founded by Steve Jobs in California. Microsoft is headquartered in Redmond." - .to_string() - }); + let mut text = use_signal(|| DEFAULT_TEXT.to_string()); let mut entity_labels = use_signal(|| "person, organization, location".to_string()); - let mut relation_labels = use_signal(|| "founded by, located in".to_string()); + let mut relation_labels = use_signal(|| "founded by, located in, headquartered in".to_string()); let mut model = use_signal(|| None::); let mut loading = use_signal(|| false); let mut running = use_signal(|| false); let mut error = use_signal(|| None::); - let mut ner_out = use_signal(Vec::::new); - let mut relex_out = use_signal(RelexResult::default); - let mut status = use_signal(|| "No model loaded".to_string()); - - let on_mode_change = move |v: String| { - let new_mode = Mode::from_value(&v); - if mode() != new_mode { - mode.set(new_mode); - choice.set(ModelChoice::default_for(new_mode)); - ner_out.write().clear(); - *relex_out.write() = RelexResult::default(); + let mut extraction = use_signal(Extraction::default); + let mut status = use_signal(|| "idle".to_string()); + let mut schedule = use_signal(|| 0u64); + + // Bump the schedule whenever any input that affects extraction changes. + // Do NOT read `model` here — the extractor writes it back on every run, + // which would re-trigger this effect in an endless loop. + use_effect(move || { + let _ = text(); + let _ = entity_labels(); + let _ = relation_labels(); + schedule.with_mut(|s| *s += 1); + }); + + // React to schedule changes: debounce, then extract. + use_effect(move || { + let current = schedule(); + if current == 0 { + return; + } + spawn(async move { + sleep_ms(DEBOUNCE_MS).await; + if schedule() != current { + return; + } + // Wait for any in-flight extraction to finish; bail if a newer change arrives. + while running() { + sleep_ms(80).await; + if schedule() != current { + return; + } + } + let Some(mut taken) = model.write().take() else { + return; + }; + running.set(true); + error.set(None); + status.set("extracting…".to_string()); + + let ent_labels = parse_labels(&entity_labels()); + let rel_labels = parse_labels(&relation_labels()); + let cur_text = text(); + let mode = taken.choice().mode(); + let outcome = + run_extraction(&mut taken, mode, &cur_text, &ent_labels, &rel_labels).await; + model.set(Some(taken)); + + match outcome { + Ok(e) => { + status.set(format!( + "{} entities · {} relations", + e.entities.len(), + e.relations.len() + )); + extraction.set(e); + } + Err(e) => { + error.set(Some(e)); + status.set("error".to_string()); + } + } + running.set(false); + }); + }); + + let on_choice_change = move |v: Option| { + let Some(c) = v else { return }; + if choice() == c { + return; } + choice.set(c); + // Unload the current model; user will see a "load" hint. + *model.write() = None; + extraction.set(Extraction::default()); }; let on_load = move |_| { @@ -154,201 +200,92 @@ fn App() -> Element { let selected = choice(); loading.set(true); error.set(None); - status.set(format!("Loading {}…", selected.label())); + status.set(format!("loading {}…", selected.label())); spawn(async move { match build_model(selected).await { Ok(m) => { - let dev = match &m { - LoadedModel::Ner { inner, .. } => { - if inner.device().is_gpu() { "GPU" } else { "CPU" } - } - LoadedModel::Relex { inner, .. } => { - if inner.device().is_gpu() { "GPU" } else { "CPU" } - } - }; model.set(Some(m)); - status.set(format!("{} ready on {dev}", selected.label())); + status.set("ready".to_string()); + // Kick an extraction now that a model is available. + schedule.with_mut(|s| *s += 1); } Err(e) => { error.set(Some(format!("{e}"))); - status.set("Load failed".to_string()); + status.set("load failed".to_string()); } } loading.set(false); }); }; - let on_extract = move |_| { - if running() { - return; - } - let current_text = text(); - let ent_labels = parse_labels(&entity_labels()); - let rel_labels = parse_labels(&relation_labels()); - let current_mode = mode(); - - running.set(true); - error.set(None); - - let Some(mut taken) = model.write().take() else { - error.set(Some("Load a model first.".to_string())); - running.set(false); - return; - }; - - spawn(async move { - let outcome = run_extraction( - &mut taken, - current_mode, - ¤t_text, - &ent_labels, - &rel_labels, - ) - .await; - - match outcome { - Ok(ExtractionOutput::Ner(entities)) => { - status.set(format!("Extracted {} entities", entities.len())); - ner_out.set(entities); - } - Ok(ExtractionOutput::Relex(result)) => { - status.set(format!( - "Extracted {} entities, {} relations", - result.entities.len(), - result.relations.len() - )); - relex_out.set(result); - } - Err(e) => error.set(Some(e)), - } - - model.set(Some(taken)); - running.set(false); - }); - }; - let has_model = model.read().is_some(); - let model_mismatch = model + let model_matches = model .read() .as_ref() - .map(|m| m.choice() != choice()) - .unwrap_or(true); + .map(|m| m.choice() == choice()) + .unwrap_or(false); - let current_mode = mode(); let current_choice = choice(); + let current_text = text(); + let cur_extraction = extraction(); rsx! { document::Link { rel: "stylesheet", href: asset!("/assets/dx-components-theme.css") } document::Link { rel: "stylesheet", href: asset!("/assets/style.css") } - div { class: "app", - header { class: "site-header", - div { - h1 { "rgliner" } - div { class: "tag", - "GLiNER NER & relation extraction — running locally in your browser with WebGPU." - } - } - a { - href: "https://github.com/floneum/floneum", - target: "_blank", - "GitHub" + main { class: "reader", + header { class: "masthead", + div { class: "wordmark", + span { class: "mark", "rgliner" } + span { class: "byline", "a reader for named entities & their relations" } } - } - - Tabs { - default_value: current_mode.value().to_string(), - horizontal: true, - on_value_change: on_mode_change, - TabList { - TabTrigger { value: "ner".to_string(), index: 0usize, "NER" } - TabTrigger { value: "relex".to_string(), index: 1usize, "NER + Relations" } - } - TabContent { index: 0usize, value: "ner".to_string(), "" } - TabContent { index: 1usize, value: "relex".to_string(), "" } - } - - if let Some(e) = error() { - div { class: "err-banner", "{e}" } - } - - Card { - CardHeader { - CardTitle { "Model" } - CardDescription { - "First load fetches GGUF weights (60 MB – 500 MB) from HuggingFace and caches them in the browser's Origin Private File System." - } - } - CardContent { - div { class: "row", - div { style: "min-width: 16rem;", - Select:: { - key: "{current_mode.value()}", - placeholder: "Select a model...", - default_value: current_choice, - on_value_change: move |v: Option| { - if let Some(c) = v { - choice.set(c); - } - }, - SelectTrigger { aria_label: "Model", SelectValue {} } - SelectList { aria_label: "Models", - for (i, c) in ModelChoice::for_mode(current_mode).iter().copied().enumerate() { - SelectOption:: { - index: i, - value: c, - text_value: c.label().to_string(), - "{c.label()}" - SelectItemIndicator {} - } - } + div { class: "picker", + Select:: { + placeholder: "choose a model", + default_value: current_choice, + on_value_change: on_choice_change, + SelectTrigger { aria_label: "Model", SelectValue {} } + SelectList { aria_label: "Models", + for (i, c) in ModelChoice::all().iter().copied().enumerate() { + SelectOption:: { + index: i, + value: c, + text_value: c.label().to_string(), + "{c.label()}" + SelectItemIndicator {} } } } - Button { - disabled: loading() || running() || (has_model && !model_mismatch), - onclick: on_load, - if loading() { - "Loading…" - } else if has_model && !model_mismatch { - "Loaded" - } else if has_model { - "Reload" - } else { - "Load model" - } - } - span { - class: if error().is_some() { "status err" } else if has_model { "status ok" } else { "status" }, - "{status()}" - } } - } - } - - Card { - CardHeader { CardTitle { "Text" } } - CardContent { - Textarea { - value: "{text}", - rows: "4", - oninput: move |e: FormEvent| text.set(e.value()), + button { + class: "load", + disabled: loading() || (has_model && model_matches), + onclick: on_load, + if loading() { + "loading…" + } else if has_model && model_matches { + "loaded" + } else if has_model { + "reload" + } else { + "load" + } } } } - Card { - CardHeader { CardTitle { "Labels" } } - CardContent { - Label { html_for: "entity-labels", "Entity labels (comma-separated)" } + details { class: "settings", + summary { "labels" } + div { class: "settings-body", + Label { html_for: "entity-labels", "entities" } Input { id: "entity-labels", r#type: "text", value: "{entity_labels}", oninput: move |e: FormEvent| entity_labels.set(e.value()), } - if current_mode == Mode::Relex { + if current_choice.mode() == Mode::Relex { div { style: "margin-top: 0.75rem;", - Label { html_for: "relation-labels", "Relation labels (comma-separated)" } + Label { html_for: "relation-labels", "relations" } Input { id: "relation-labels", r#type: "text", @@ -360,24 +297,31 @@ fn App() -> Element { } } - Card { - CardContent { - div { class: "row", - Button { - disabled: !has_model || running() || loading() || model_mismatch, - onclick: on_extract, - if running() { "Extracting…" } else { "Extract" } - } - if model_mismatch && has_model { - span { class: "status", - "Model selection changed — reload to use it." - } + div { class: "status-line", + span { class: "dot", class: if running() || loading() { "busy" } else if error().is_some() { "err" } else if has_model { "ok" } else { "idle" } } + span { class: "msg", "{status()}" } + if let Some(e) = error() { + span { class: "err-text", " · {e}" } + } + } + + div { class: "split", + textarea { + class: "editor", + spellcheck: "false", + value: "{current_text}", + oninput: move |e: FormEvent| text.set(e.value()), + } + article { class: "article", + if has_model { + { render_article(¤t_text, &cur_extraction) } + } else { + div { class: "placeholder", + "load a model to begin reading." } } } } - - { render_results(current_mode, text(), ner_out(), relex_out()) } } } } @@ -426,23 +370,17 @@ async fn build_model(choice: ModelChoice) -> Result { } } -enum ExtractionOutput { - Ner(Vec), - Relex(RelexResult), -} - async fn run_extraction( model: &mut LoadedModel, mode: Mode, text: &str, entity_labels: &[String], relation_labels: &[String], -) -> Result { - let ent_refs: Vec<&str> = entity_labels.iter().map(|s| s.as_str()).collect(); - - if ent_refs.is_empty() { - return Err("Add at least one entity label.".to_string()); +) -> Result { + if entity_labels.is_empty() { + return Err("add at least one entity label".to_string()); } + let ent_refs: Vec<&str> = entity_labels.iter().map(|s| s.as_str()).collect(); match (mode, model) { (Mode::Ner, LoadedModel::Ner { inner, .. }) => { @@ -450,103 +388,50 @@ async fn run_extraction( .extract_auto(text, &ent_refs, Some(TOKEN_BUDGET)) .await .map_err(|e| format!("{e}"))?; - Ok(ExtractionOutput::Ner(entities)) + Ok(Extraction { + entities, + relations: Vec::new(), + }) } (Mode::Relex, LoadedModel::Relex { inner, .. }) => { - let rel_refs: Vec<&str> = relation_labels.iter().map(|s| s.as_str()).collect(); - if rel_refs.is_empty() { - return Err("Add at least one relation label.".to_string()); + if relation_labels.is_empty() { + return Err("add at least one relation label".to_string()); } + let rel_refs: Vec<&str> = relation_labels.iter().map(|s| s.as_str()).collect(); let (entities, relations) = inner .extract_auto(text, &ent_refs, &rel_refs, Some(TOKEN_BUDGET)) .await .map_err(|e| format!("{e}"))?; - Ok(ExtractionOutput::Relex(RelexResult { + Ok(Extraction { entities, relations, - })) + }) } - _ => Err("Loaded model doesn't match the current tab. Reload the model.".to_string()), + _ => Err("loaded model doesn't match the selection".to_string()), } } -fn render_results( - mode: Mode, - text: String, - ner_out: Vec, - relex_out: RelexResult, -) -> Element { - let (entities, relations): (Vec, Vec) = match mode { - Mode::Ner => (ner_out, Vec::new()), - Mode::Relex => (relex_out.entities, relex_out.relations), - }; - - if entities.is_empty() && relations.is_empty() { - return rsx! { - Card { - CardContent { - p { class: "muted", "Run extraction to see results." } - } - } - }; - } - +fn render_article(text: &str, ex: &Extraction) -> Element { + let spliced = splice_entities(text, &ex.entities, &ex.relations); rsx! { - Card { - CardHeader { CardTitle { "Results" } } - CardContent { - div { class: "results", - { highlighted_text(text.clone(), entities.clone()) } - } - - if !entities.is_empty() { - ul { class: "entity-list", - for (i, ent) in entities.iter().enumerate() { - li { key: "{i}", - Badge { - variant: BadgeVariant::Outline, - span { style: "color: {hsl_for(&ent.label)};", "{ent.label}" } - } - " · " - span { "{ent.text:?}" } - " " - span { class: "score", "{format_score(ent.score)}" } - } - } - } - } - - if !relations.is_empty() { - div { style: "margin-top: 1rem;", - Label { html_for: "relations-list", "Relations" } - ul { class: "relation-list", - for (i, rel) in relations.iter().enumerate() { - li { key: "rel-{i}", - Badge { "{rel.head.text}" } - " --[" - span { style: "color: var(--accent);", "{rel.relation}" } - "]--> " - Badge { "{rel.tail.text}" } - " " - span { class: "score", "{format_score(rel.score)}" } - } - } - } - } - } - } - } + Markdown { src: spliced } } } -fn highlighted_text(text: String, entities: Vec) -> Element { - let mut sorted = entities.clone(); +/// Splice `` into the markdown source at each +/// entity boundary. The nested `.rels` span renders on hover. +fn splice_entities(text: &str, entities: &[Entity], relations: &[Relation]) -> String { + if entities.is_empty() { + return text.to_string(); + } + let mut sorted: Vec<&Entity> = entities.iter().collect(); sorted.sort_by_key(|e| e.start_char); - let mut segments: Vec<(bool, String, Option)> = Vec::new(); + let mut out = String::with_capacity(text.len() + entities.len() * 64); let mut cursor = 0usize; let len = text.len(); - for ent in sorted.iter() { + + for ent in sorted { let start = ent.start_char.min(len); let end = ent.end_char.min(len); if start < cursor || end <= start { @@ -555,47 +440,67 @@ fn highlighted_text(text: String, entities: Vec) -> Element { if !text.is_char_boundary(start) || !text.is_char_boundary(end) { continue; } - if start > cursor { - segments.push((false, text[cursor..start].to_string(), None)); + out.push_str(&text[cursor..start]); + let color = underline_color(&ent.label); + out.push_str(&format!( + r#""#, + color = color, + label = escape_attr(&ent.label) + )); + // The entity surface text. + out.push_str(&escape_html_content(&text[start..end])); + // Popover: label + relations involving this entity. + out.push_str(r#""#); + out.push_str(r#""#); + out.push_str(&escape_html_content(&ent.label)); + out.push_str(""); + + let ent_text = &text[start..end]; + let mut rel_lines = 0usize; + for rel in relations { + if rel.head.text == ent_text || rel.tail.text == ent_text { + out.push_str(r#""#); + out.push_str(&escape_html_content(&rel.head.text)); + out.push_str(r#""#); + out.push_str(r#""#); + out.push_str(&escape_html_content(&rel.relation)); + out.push_str(""); + out.push_str(r#""#); + out.push_str(&escape_html_content(&rel.tail.text)); + out.push_str(""); + rel_lines += 1; + } + } + if rel_lines == 0 && !relations.is_empty() { + out.push_str(r#"no relations"#); } - segments.push((true, text[start..end].to_string(), Some(ent.label.clone()))); + out.push_str(""); + out.push_str(""); cursor = end; } if cursor < len { - segments.push((false, text[cursor..].to_string(), None)); + out.push_str(&text[cursor..]); } + out +} - rsx! { - for (i, (is_entity, content, label)) in segments.into_iter().enumerate() { - if is_entity { - { - let label_text = label.clone().unwrap_or_default(); - let color = hsl_for(&label_text); - rsx! { - span { - key: "seg-{i}", - class: "entity", - style: "background-color: {color};", - "{content}" - span { class: "chip", "{label_text}" } - } - } - } - } else { - span { key: "seg-{i}", "{content}" } - } - } - } +fn escape_html_content(s: &str) -> String { + s.replace('&', "&") + .replace('<', "<") + .replace('>', ">") } -fn hsl_for(label: &str) -> String { +fn escape_attr(s: &str) -> String { + s.replace('&', "&") + .replace('"', """) + .replace('<', "<") + .replace('>', ">") +} + +fn underline_color(label: &str) -> String { let hash: u32 = label .bytes() .fold(0u32, |acc, b| acc.wrapping_mul(31).wrapping_add(b as u32)); let hue = hash % 360; - format!("hsl({hue}, 70%, 72%)") -} - -fn format_score(score: f32) -> String { - format!("{:.2}", score) + format!("hsl({hue}, 75%, 42%)") } diff --git a/fusor-ml/core/src/compute_graph/mod.rs b/fusor-ml/core/src/compute_graph/mod.rs index d89c6f76b..01ca5ca01 100644 --- a/fusor-ml/core/src/compute_graph/mod.rs +++ b/fusor-ml/core/src/compute_graph/mod.rs @@ -353,6 +353,20 @@ impl ComputeGraphInner { .and_then(|n| n.cached.as_ref()) } + pub(crate) fn debug_node_state(&self, key: NodeIndex) -> String { + self.nodes + .nodes + .node_weight(key) + .map(|n| { + format!( + "variant={:?} cached={}", + n.variant, + n.cached.is_some() + ) + }) + .unwrap_or_else(|| "missing".to_string()) + } + #[cfg(feature = "extra_assertions")] fn contains_key(&self, key: NodeIndex) -> bool { self.nodes.nodes.contains_node(key) diff --git a/fusor-ml/core/src/compute_graph/resolve.rs b/fusor-ml/core/src/compute_graph/resolve.rs index 6c5547d8c..a248d0819 100644 --- a/fusor-ml/core/src/compute_graph/resolve.rs +++ b/fusor-ml/core/src/compute_graph/resolve.rs @@ -81,6 +81,8 @@ impl<'a> Resolver<'a> { // Pass 2: Apply Rewrite Rules self.optimize(graph); + self.rebuild_execution_edges(graph); + self.prune_unreachable_from_target(); // Pass 3: Topological Sort let sorted_nodes = toposort(&self.execution_graph, None) @@ -396,26 +398,77 @@ impl<'a> Resolver<'a> { } fn remove_node_if_dead(&mut self, node_idx: ExecutionNodeIndex) { - if !self.execution_graph.contains_node(node_idx) { + let _ = node_idx; + // Defer dead-node pruning until rewrites are complete. During optimization, + // later rewrite passes may still need to reconnect nodes that temporarily + // have no outgoing edges. + } + + fn rebuild_execution_edges(&mut self, graph: &ComputeGraphInner) { + let edge_indices: Vec<_> = self.execution_graph.edge_indices().collect(); + for edge in edge_indices { + self.execution_graph.remove_edge(edge); + } + + let node_indices: Vec<_> = self.execution_graph.node_indices().collect(); + for node_idx in node_indices { + if !self.execution_graph.contains_node(node_idx) { + continue; + } + let variant = self.execution_graph[node_idx].variant.clone(); + let mut dependencies = Vec::new(); + variant.visit_dependencies(&mut |dependency| { + dependencies.push(dependency); + }); + + for dependency in dependencies { + if self.check_cached(graph, dependency) { + continue; + } + let Some(dep_exec_idx) = self.get_input_node_in_exec_graph(dependency) else { + continue; + }; + if !self.execution_graph.contains_node(dep_exec_idx) || dep_exec_idx == node_idx { + continue; + } + if self.execution_graph.find_edge(dep_exec_idx, node_idx).is_none() { + self.execution_graph.add_edge(dep_exec_idx, node_idx, ()); + } + } + } + } + + fn prune_unreachable_from_target(&mut self) { + let Some(&target_exec) = self.node_mapping.get(&self.target) else { + return; + }; + if !self.execution_graph.contains_node(target_exec) { return; } - if self - .execution_graph - .neighbors_directed(node_idx, petgraph::Direction::Outgoing) - .count() - == 0 - { - // Collect incoming neighbors before removing - let incoming: Vec<_> = self + + let mut reachable = FxHashSet::default(); + let mut stack = vec![target_exec]; + while let Some(node_idx) = stack.pop() { + if !reachable.insert(node_idx) { + continue; + } + for dependency in self .execution_graph .neighbors_directed(node_idx, petgraph::Direction::Incoming) - .collect(); - self.execution_graph.remove_node(node_idx); - // Recursively check if dependencies are now dead - for dep in incoming { - self.remove_node_if_dead(dep); + { + stack.push(dependency); } } + + let all_nodes: Vec<_> = self.execution_graph.node_indices().collect(); + for node_idx in all_nodes { + if !reachable.contains(&node_idx) { + self.execution_graph.remove_node(node_idx); + } + } + + self.node_mapping + .retain(|_, exec_idx| self.execution_graph.contains_node(*exec_idx)); } // Rules @@ -596,6 +649,7 @@ impl<'a> Resolver<'a> { } for &new_input in &new_inputs { if let Some(exec) = self.get_input_node_in_exec_graph(new_input) + && self.execution_graph.contains_node(exec) && self.execution_graph.find_edge(exec, node_idx).is_none() { self.execution_graph.add_edge(exec, node_idx, ()); diff --git a/fusor-ml/core/src/nary_wise.rs b/fusor-ml/core/src/nary_wise.rs index db0cd8e36..327d7566b 100644 --- a/fusor-ml/core/src/nary_wise.rs +++ b/fusor-ml/core/src/nary_wise.rs @@ -507,7 +507,13 @@ impl Operation for NaryOperation { } } // Otherwise use the normal path which may return QMatrix for Dequantize nodes - nodes.get_result_or_qmatrix(*idx).unwrap().into() + nodes.get_result_or_qmatrix(*idx).unwrap_or_else(|| { + let node_debug = nodes.debug_node_state(*idx); + panic!( + "nary input {i} missing for node {:?}: {node_debug}", + idx + ); + }).into() }) .collect(); diff --git a/fusor-ml/core/src/tensor.rs b/fusor-ml/core/src/tensor.rs index 7457eda8f..d68364595 100644 --- a/fusor-ml/core/src/tensor.rs +++ b/fusor-ml/core/src/tensor.rs @@ -873,6 +873,16 @@ impl Tensor { } } + /// Resolve this tensor's pending compute graph into a new concrete tensor. + /// + /// Unlike [`materialize`], which only waits for queued work, this returns a + /// tensor backed by the resolved output buffer so subsequent ops no longer + /// extend the original lazy graph. + pub fn materialized(&self) -> Self { + let (tensor, _) = self.data.materialize(); + Self::from(tensor) + } + /// How many kernel calls are needed to fully resolve this tensor pub fn count_kernels_to_resolve(&self) -> usize { let (_, count) = self.data.materialize(); diff --git a/fusor-ml/fusor/src/lib.rs b/fusor-ml/fusor/src/lib.rs index e58ae2cc8..90ceb07ff 100644 --- a/fusor-ml/fusor/src/lib.rs +++ b/fusor-ml/fusor/src/lib.rs @@ -417,6 +417,39 @@ where } } + /// Materialize pending work for this tensor. + /// + /// On CPU this evaluates lazy expressions; on GPU it waits for queued work + /// associated with the tensor to complete. + pub async fn materialize(&self) + where + B: TensorBacking, + D: SimdElement + DataType, + { + match self { + Tensor::Cpu(t) => { + let _ = t.to_concrete(); + } + Tensor::Gpu(t) => t.materialize().await, + } + } + + /// Resolve this tensor into a new concrete tensor. + /// + /// For CPU tensors this evaluates lazy expressions. For GPU tensors this + /// returns a tensor backed by the resolved output buffer, so later ops do + /// not keep extending the original lazy graph. + pub fn materialized(&self) -> Tensor + where + B: TensorBacking, + D: SimdElement + DataType, + { + match self { + Tensor::Cpu(t) => Tensor::Cpu(t.to_concrete()), + Tensor::Gpu(t) => Tensor::Gpu(t.materialized()), + } + } + /// Returns the shape of the tensor. pub fn shape(&self) -> [usize; R] where diff --git a/models/rbert/src/raw/mdeberta/attention.rs b/models/rbert/src/raw/mdeberta/attention.rs index dcb025ad2..9fda2f79a 100644 --- a/models/rbert/src/raw/mdeberta/attention.rs +++ b/models/rbert/src/raw/mdeberta/attention.rs @@ -247,6 +247,7 @@ impl MDebertaAttention { self.num_heads, self.head_dim, ); + let [batch_size, _, _, _] = query.shape(); // === Content-to-Content attention === let c2c_scores = query.mat_mul(&key.transpose(2, 3)); @@ -264,6 +265,23 @@ impl MDebertaAttention { self.num_heads, self.head_dim, ); + let num_relative_positions = pos_query.shape()[2]; + let pos_query = pos_query + .broadcast_as([ + batch_size, + self.num_heads, + num_relative_positions, + self.head_dim, + ]) + .to_concrete(); + let pos_key = pos_key + .broadcast_as([ + batch_size, + self.num_heads, + num_relative_positions, + self.head_dim, + ]) + .to_concrete(); // === Content-to-Position attention === // c2p = Q @ pos_key^T -> [batch, heads, seq, 2*max_pos] diff --git a/models/rbert/src/raw/mdeberta/model.rs b/models/rbert/src/raw/mdeberta/model.rs index fb6dc4486..52a6ccf98 100644 --- a/models/rbert/src/raw/mdeberta/model.rs +++ b/models/rbert/src/raw/mdeberta/model.rs @@ -117,6 +117,29 @@ impl MDebertaModel { hidden_states } + #[doc(hidden)] + pub fn debug_after_embedding_norm(&self, input_ids: &Tensor<2, u32>) -> Tensor<3, f32> { + let hidden_states = self.token_embeddings.forward(input_ids); + self.embedding_norm.forward(&hidden_states) + } + + #[doc(hidden)] + pub fn debug_first_layer_output( + &self, + hidden_states: &Tensor<3, f32>, + attention_mask: Option<&Tensor<2, u32>>, + ) -> Tensor<3, f32> { + let [b_sz, seq_len, _] = hidden_states.shape(); + let gather_idx = self.rel_pos_embedding.compute_gather_indices( + b_sz, + self.config.num_heads, + seq_len, + &self.device, + ); + let rel_pos_emb = self.rel_pos_embedding.get_embeddings(); + self.layers[0].forward_with_rel(hidden_states, &rel_pos_emb, &gather_idx, attention_mask) + } + /// Get the embedding dimension. pub fn embedding_dim(&self) -> usize { self.config.hidden_size diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs index e8bcce633..94f5d4eaf 100644 --- a/models/rgliner/src/lib.rs +++ b/models/rgliner/src/lib.rs @@ -118,8 +118,8 @@ use rbert::BertSource; use std::sync::Arc; use tokenizers::Tokenizer; -use crate::raw::{CachedLabels, LabelEncoder, Scorer, SpanLayer, TextEncoder}; -use crate::tokenization::{first_subtoken_pooling, WordTokenizer}; +use crate::raw::{CachedLabels, LabelEncoder, SpanLayer, TextEncoder}; +use crate::tokenization::{first_subtoken_pooling, TokenizedText, WordTokenizer}; async fn default_device() -> Device { Device::gpu().await.unwrap_or_else(|_| Device::cpu()) @@ -440,60 +440,48 @@ impl Gliner { self.label_encoder.encode_labels(labels).await? }; - let mut results = Vec::with_capacity(texts.len()); - for text in texts { - let entities = self - .extract_internal(text, labels, &label_embeddings) - .await?; - results.push(entities); - } - - Ok(results) + self.extract_internal_batch(texts, labels, &label_embeddings) + .await } - async fn extract_internal( + async fn extract_internal_batch( &self, - text: &str, + texts: &[&str], labels: &[&str], label_embeddings: &Tensor<2, f32>, - ) -> Result, GlinerError> { - // 1. Tokenize text - let tokenized = self.tokenizer.tokenize(text)?; - - if tokenized.num_words == 0 { - return Ok(Vec::new()); + ) -> Result>, GlinerError> { + let tokenized = self.tokenizer.tokenize_batch(texts)?; + if tokenized.iter().all(|tokenized| tokenized.num_words == 0) { + return Ok(vec![Vec::new(); texts.len()]); } - // 2. Prepare input tensors - let token_ids = Tensor::new(&self.device, &tokenized.token_ids); - let token_ids: Tensor<2, u32> = token_ids.unsqueeze(0).to_concrete(); - - let attention_mask = Tensor::new(&self.device, &tokenized.attention_mask); - let attention_mask: Tensor<2, u32> = attention_mask.unsqueeze(0).to_concrete(); + let (token_ids, attention_mask) = self.build_batched_inputs(&tokenized); - // 3. Encode text let token_embeddings = self.text_encoder.forward(&token_ids, Some(&attention_mask)); // Python's bi-encoder span model pools transformer token embeddings // directly to words; the checkpoint still contains LSTM weights, but // that path is not used in BaseBiEncoderModel.get_representations(). let (word_embeddings, _word_mask) = - first_subtoken_pooling(&token_embeddings, &[tokenized.clone()], &self.device); + first_subtoken_pooling(&token_embeddings, &tokenized, &self.device); - // 4. Generate span representations - let (span_embeddings, span_indices) = - self.span_layer.forward(&word_embeddings, &self.device); + let spans_per_batch: Vec> = tokenized + .iter() + .map(|tokenized| self.enumerate_spans(tokenized.num_words)) + .collect(); + let span_counts: Vec = spans_per_batch.iter().map(Vec::len).collect(); + let total_spans: usize = span_counts.iter().sum(); - // 5. Score spans against labels - let scores = Scorer::forward(&span_embeddings, label_embeddings); + if total_spans == 0 { + return Ok(vec![Vec::new(); texts.len()]); + } - // 6. Decode predictions - let shape = scores.shape(); - let num_spans = shape[1]; - let num_labels = shape[2]; + let (flat_span_embeddings, _) = + self.span_layer + .forward_for_spans_batched(&word_embeddings, &spans_per_batch, &self.device); - // Get scores for first batch item and apply sigmoid - let flat_scores: Tensor<2, f32> = scores.squeeze(0).to_concrete(); + let labels_t = label_embeddings.t(); + let flat_scores = flat_span_embeddings.mat_mul(&labels_t); let tensor_slice = flat_scores.as_slice().await?; let scores_data: Vec = tensor_slice .as_slice() @@ -501,17 +489,72 @@ impl Gliner { .map(|&x| 1.0 / (1.0 + (-x).exp())) // sigmoid .collect(); - let entities = self.decoder.decode( - &scores_data, - num_spans, - num_labels, - &span_indices, - &tokenized.word_offsets, - labels, - text, - ); + let num_labels = label_embeddings.shape()[0]; + let mut results = Vec::with_capacity(texts.len()); + let mut score_offset = 0usize; + + for (batch_idx, tokenized) in tokenized.iter().enumerate() { + let span_count = span_counts[batch_idx]; + if span_count == 0 { + results.push(Vec::new()); + continue; + } + + let next_offset = score_offset + span_count * num_labels; + let entities = self.decoder.decode( + &scores_data[score_offset..next_offset], + span_count, + num_labels, + &spans_per_batch[batch_idx], + &tokenized.word_offsets, + labels, + texts[batch_idx], + ); + results.push(entities); + score_offset = next_offset; + } + + Ok(results) + } - Ok(entities) + fn build_batched_inputs(&self, tokenized: &[TokenizedText]) -> (Tensor<2, u32>, Tensor<2, u32>) { + let batch_size = tokenized.len(); + let max_seq_len = tokenized + .iter() + .map(|tokenized| tokenized.token_ids.len()) + .max() + .unwrap_or(1); + let pad_id = self.tokenizer.pad_id(); + + let mut token_ids = vec![pad_id; batch_size * max_seq_len]; + let mut attention_mask = vec![0u32; batch_size * max_seq_len]; + + for (batch_idx, item) in tokenized.iter().enumerate() { + let offset = batch_idx * max_seq_len; + let len = item.token_ids.len(); + token_ids[offset..offset + len].copy_from_slice(&item.token_ids); + attention_mask[offset..offset + len].copy_from_slice(&item.attention_mask); + } + + ( + Tensor::new(&self.device, &token_ids) + .reshape([batch_size, max_seq_len]) + .to_concrete(), + Tensor::new(&self.device, &attention_mask) + .reshape([batch_size, max_seq_len]) + .to_concrete(), + ) + } + + fn enumerate_spans(&self, num_words: usize) -> Vec<(usize, usize)> { + let mut spans = Vec::new(); + for start in 0..num_words { + let max_width = self.max_width.min(num_words - start); + for width in 1..=max_width { + spans.push((start, start + width - 1)); + } + } + spans } /// Get the maximum span width. @@ -528,6 +571,7 @@ impl Gliner { #[cfg(test)] mod gpu_parity_tests { use super::*; + use crate::raw::Scorer; use fusor::layers::{Embedding, LayerNorm}; use std::path::Path; diff --git a/models/rgliner/src/raw/bilstm.rs b/models/rgliner/src/raw/bilstm.rs index 2d2f6a24d..3530f988f 100644 --- a/models/rgliner/src/raw/bilstm.rs +++ b/models/rgliner/src/raw/bilstm.rs @@ -116,6 +116,122 @@ impl BiLstm { .reshape([batch, seq_len, 2 * self.hidden_size]) .to_concrete() } + + #[cfg(test)] + #[doc(hidden)] + pub fn debug_forward_direction( + &self, + input: &Tensor<3, f32>, + lengths: &[usize], + reverse: bool, + ) -> Tensor<3, f32> { + run_direction( + input, + if reverse { + &self.backward + } else { + &self.forward + }, + self.hidden_size, + &input.device(), + reverse, + lengths, + ) + } + + #[cfg(test)] + #[doc(hidden)] + pub fn debug_first_step_gates( + &self, + input: &Tensor<3, f32>, + reverse: bool, + ) -> Tensor<2, f32> { + let [batch, seq_len, input_size] = input.shape(); + let hidden_size = self.hidden_size; + let dir = if reverse { + &self.backward + } else { + &self.forward + }; + let t = if reverse { seq_len - 1 } else { 0 }; + let h: Tensor<2, f32> = Tensor::zeros(&input.device(), [batch, hidden_size]); + let x_t: Tensor<2, f32> = input + .narrow(1, t, 1) + .reshape([batch, input_size]) + .to_concrete(); + let bias_broadcast: Tensor<2, f32> = dir + .bias + .unsqueeze(0) + .broadcast_as([batch, 4 * hidden_size]) + .to_concrete(); + (x_t.mat_mul(&dir.w_ih_t) + h.mat_mul(&dir.w_hh_t) + bias_broadcast).to_concrete() + } + + #[cfg(test)] + #[doc(hidden)] + pub fn debug_forward_direction_state_only( + &self, + input: &Tensor<3, f32>, + lengths: &[usize], + reverse: bool, + ) -> Tensor<2, f32> { + let [batch, seq_len, input_size] = input.shape(); + let hidden_size = self.hidden_size; + let device = input.device(); + let dir = if reverse { + &self.backward + } else { + &self.forward + }; + + let mut h: Tensor<2, f32> = Tensor::zeros(&device, [batch, hidden_size]); + let mut c: Tensor<2, f32> = Tensor::zeros(&device, [batch, hidden_size]); + let bias_broadcast: Tensor<2, f32> = dir + .bias + .unsqueeze(0) + .broadcast_as([batch, 4 * hidden_size]) + .to_concrete(); + + let iter: Box> = if reverse { + Box::new((0..seq_len).rev()) + } else { + Box::new(0..seq_len) + }; + + for t in iter { + let x_t: Tensor<2, f32> = input + .narrow(1, t, 1) + .reshape([batch, input_size]) + .to_concrete(); + let gates_pre: Tensor<2, f32> = + (x_t.mat_mul(&dir.w_ih_t) + h.mat_mul(&dir.w_hh_t) + bias_broadcast.clone()) + .to_concrete(); + + let i_raw: Tensor<2, f32> = gates_pre.narrow(1, 0, hidden_size).to_concrete(); + let f_raw: Tensor<2, f32> = + gates_pre.narrow(1, hidden_size, hidden_size).to_concrete(); + let g_raw: Tensor<2, f32> = gates_pre + .narrow(1, 2 * hidden_size, hidden_size) + .to_concrete(); + let o_raw: Tensor<2, f32> = gates_pre + .narrow(1, 3 * hidden_size, hidden_size) + .to_concrete(); + + let i_gate = sigmoid_2d(&i_raw); + let f_gate = sigmoid_2d(&f_raw); + let g_gate = g_raw.tanh(); + let o_gate = sigmoid_2d(&o_raw); + + let next_c = (f_gate * c.clone() + i_gate * g_gate).to_concrete(); + let next_h = (o_gate * next_c.clone().tanh()).to_concrete(); + + let active_mask_2d = timestep_mask_2d(&device, batch, hidden_size, lengths, t); + c = active_mask_2d.where_cond(&next_c, &c).to_concrete(); + h = active_mask_2d.where_cond(&next_h, &h).to_concrete(); + } + + h + } } /// Run one direction of the LSTM. Sequential over time; every timestep's gate @@ -132,12 +248,7 @@ fn run_direction( let mut h: Tensor<2, f32> = Tensor::zeros(device, [batch, hidden_size]); let mut c: Tensor<2, f32> = Tensor::zeros(device, [batch, hidden_size]); - - // outputs[t] holds the hidden state at timestep t, already unsqueezed on - // dim 1 so that a final cat along dim 1 yields [batch, seq_len, hidden]. - let mut outputs: Vec> = Vec::with_capacity(seq_len); - outputs.resize_with(seq_len, || Tensor::zeros(device, [batch, 1, hidden_size])); - let zero_output: Tensor<3, f32> = Tensor::zeros(device, [batch, 1, hidden_size]); + let mut outputs: Tensor<3, f32> = Tensor::zeros(device, [batch, seq_len, hidden_size]); let bias_broadcast: Tensor<2, f32> = dir .bias @@ -183,14 +294,27 @@ fn run_direction( c = active_mask_2d.where_cond(&next_c, &c).to_concrete(); h = active_mask_2d.where_cond(&next_h, &h).to_concrete(); - let active_mask_3d = timestep_mask_3d(device, batch, hidden_size, lengths, t); - let output_t = h.clone().unsqueeze(1).to_concrete(); - outputs[t] = active_mask_3d.where_cond(&output_t, &zero_output).to_concrete(); + let output_t = h.clone().unsqueeze(1).to_concrete().materialized(); + outputs = outputs + .slice_assign([0..batch, t..(t + 1), 0..hidden_size], &output_t) + .materialized(); } - Tensor::cat(outputs, 1) - .reshape([batch, seq_len, hidden_size]) - .to_concrete() + let all_active = lengths.iter().all(|&length| length >= seq_len); + if all_active { + outputs + } else { + let mask_data: Vec = lengths + .iter() + .flat_map(|&length| (0..seq_len).map(move |t| if t < length { 1.0 } else { 0.0 })) + .collect(); + let mask: Tensor<3, f32> = Tensor::new(device, &mask_data) + .reshape([batch, seq_len, 1]) + .broadcast_as([batch, seq_len, hidden_size]) + .to_concrete(); + let zeros: Tensor<3, f32> = Tensor::zeros(device, [batch, seq_len, hidden_size]); + mask.where_cond(&outputs, &zeros).to_concrete() + } } fn timestep_mask_2d( @@ -210,23 +334,6 @@ fn timestep_mask_2d( .to_concrete() } -fn timestep_mask_3d( - device: &Device, - batch: usize, - hidden_size: usize, - lengths: &[usize], - timestep: usize, -) -> Tensor<3, f32> { - let mask_data: Vec = lengths - .iter() - .map(|&length| if timestep < length { 1.0 } else { 0.0 }) - .collect(); - Tensor::new(device, &mask_data) - .reshape([batch, 1, 1]) - .broadcast_as([batch, 1, hidden_size]) - .to_concrete() -} - /// sigmoid via `0.5 * (tanh(x / 2) + 1)` — avoids needing scalar-left division /// or a `recip` primitive, and keeps the computation on-device. fn sigmoid_2d(x: &Tensor<2, f32>) -> Tensor<2, f32> { diff --git a/models/rgliner/src/raw/joint_scorer.rs b/models/rgliner/src/raw/joint_scorer.rs index 04e477859..9567312ad 100644 --- a/models/rgliner/src/raw/joint_scorer.rs +++ b/models/rgliner/src/raw/joint_scorer.rs @@ -43,7 +43,7 @@ impl JointScorer { /// /// # Arguments /// * `token_embs` - Token embeddings [batch, seq_len, hidden_dim] - /// * `label_embs` - Label embeddings [n_labels, hidden_dim] + /// * `label_embs` - Label embeddings [batch, n_labels, hidden_dim] /// /// # Returns /// Scores [batch, seq_len, n_labels, 3] (3 classes: O, B, I) @@ -57,26 +57,28 @@ impl JointScorer { pub fn forward( &self, token_embs: &Tensor<3, f32>, - label_embs: &Tensor<2, f32>, + label_embs: &Tensor<3, f32>, ) -> Tensor<4, f32> { let [batch_size, seq_len, _hidden_dim] = token_embs.shape(); - let [n_labels, _] = label_embs.shape(); + let [label_batch_size, n_labels, _] = label_embs.shape(); + assert_eq!( + batch_size, label_batch_size, + "label batch size must match token batch size" + ); // Project tokens: [batch, seq, hidden] -> [batch, seq, 2*half] let proj_tokens = self.proj_token.forward(token_embs); let [_, _, proj_dim] = proj_tokens.shape(); let half = proj_dim / 2; - // Project labels: [n_labels, hidden] -> [n_labels, 2*half] - let label_embs_3d: Tensor<3, f32> = label_embs.unsqueeze(0).to_concrete(); - let proj_labels = self.proj_label.forward(&label_embs_3d); - let proj_labels: Tensor<2, f32> = proj_labels.squeeze(0).to_concrete(); + // Project labels: [batch, n_labels, hidden] -> [batch, n_labels, 2*half] + let proj_labels = self.proj_label.forward(label_embs); // Split projections into first/second halves along the feature dim. let tokens_first = proj_tokens.narrow(2, 0, half).to_concrete(); // [b, s, half] let tokens_second = proj_tokens.narrow(2, half, half).to_concrete(); // [b, s, half] - let labels_first = proj_labels.narrow(1, 0, half).to_concrete(); // [n, half] - let labels_second = proj_labels.narrow(1, half, half).to_concrete(); // [n, half] + let labels_first = proj_labels.narrow(2, 0, half).to_concrete(); // [b, n, half] + let labels_second = proj_labels.narrow(2, half, half).to_concrete(); // [b, n, half] // Broadcast to [batch, seq, n_labels, half] and build the three concatenation parts: // [token_first, label_first, token_second * label_second] @@ -88,13 +90,11 @@ impl JointScorer { .broadcast_as(target) .to_concrete(); let lab_first_4d: Tensor<4, f32> = labels_first - .unsqueeze(0) - .unsqueeze(0) + .unsqueeze(1) .broadcast_as(target) .to_concrete(); let lab_second_4d: Tensor<4, f32> = labels_second - .unsqueeze(0) - .unsqueeze(0) + .unsqueeze(1) .broadcast_as(target) .to_concrete(); @@ -126,7 +126,7 @@ impl JointScorer { pub fn forward_entity_scores( &self, token_embs: &Tensor<3, f32>, - label_embs: &Tensor<2, f32>, + label_embs: &Tensor<3, f32>, ) -> Tensor<4, f32> { let logits = self.forward(token_embs, label_embs); // sigmoid(x) = 0.5 * (tanh(x / 2) + 1); stays on-device and avoids needing diff --git a/models/rgliner/src/raw/mod.rs b/models/rgliner/src/raw/mod.rs index b8557b656..c2565064e 100644 --- a/models/rgliner/src/raw/mod.rs +++ b/models/rgliner/src/raw/mod.rs @@ -12,6 +12,7 @@ pub use bilstm::BiLstm; pub use joint_scorer::{JointScorer, PromptRepLayer}; pub use label_encoder::{CachedLabels, LabelEncoder}; pub use pair_projector::PairProjector; +#[allow(unused_imports)] pub use scorer::Scorer; pub use span_layer::SpanLayer; pub use text_encoder::TextEncoder; diff --git a/models/rgliner/src/raw/span_layer.rs b/models/rgliner/src/raw/span_layer.rs index 82ce20b8e..1e71ba4b1 100644 --- a/models/rgliner/src/raw/span_layer.rs +++ b/models/rgliner/src/raw/span_layer.rs @@ -234,11 +234,18 @@ impl SpanLayer { let start_gathered = start_rep_flat.index_select(0, &start_idx_tensor); let end_gathered = end_rep_flat.index_select(0, &end_idx_tensor); - let combined = Tensor::cat([start_gathered, end_gathered], 1).relu(); + let combined = Tensor::cat([start_gathered, end_gathered], 1) + .reshape([1, total_spans, hidden_dim * 2]) + .to_concrete() + .relu(); let hidden = self.out_fc1.forward(&combined).relu(); - let out = self.out_fc2.forward(&hidden); + let out = self + .out_fc2 + .forward(&hidden) + .reshape([total_spans, hidden_dim]) + .to_concrete(); - (out.to_concrete(), span_counts) + (out, span_counts) } fn gather_span_embeddings( diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index b9b7fe2b4..f540a16e2 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -52,7 +52,7 @@ use crate::decoding::Entity; use crate::error::{GlinerError, GlinerLoadingError}; use crate::raw::{BiLstm, JointScorer, PairProjector, PromptRepLayer, SpanLayer}; use crate::relation_decoding::Relation; -use crate::relex_tokenization::{RelExTokenizer, SpecialTokenIds}; +use crate::relex_tokenization::{RelExTokenizedInput, RelExTokenizer, SpecialTokenIds}; use rbert::raw::MDebertaModel; /// Source configuration for GLiNER-RelEx models. @@ -442,156 +442,97 @@ impl GlinerRelEx { entity_labels: &[&str], relation_labels: &[&str], ) -> Result<(Vec, Vec), GlinerError> { - // 1. Tokenize with special tokens - let tokenized = self - .tokenizer - .tokenize(text, entity_labels, relation_labels)?; + let mut results = self.extract_batch(&[text], entity_labels, relation_labels).await?; + Ok(results.pop().unwrap_or_default()) + } - if tokenized.num_words == 0 { - return Ok((Vec::new(), Vec::new())); + /// Extract entities and relations from a batch of texts. + pub async fn extract_batch( + &self, + texts: &[&str], + entity_labels: &[&str], + relation_labels: &[&str], + ) -> Result, Vec)>, GlinerError> { + if texts.is_empty() { + return Ok(Vec::new()); } - // 2. Prepare input tensors - let token_ids = Tensor::new(&self.device, &tokenized.token_ids); - let token_ids: Tensor<2, u32> = token_ids.unsqueeze(0).to_concrete(); - - let attention_mask = Tensor::new(&self.device, &tokenized.attention_mask); - let attention_mask: Tensor<2, u32> = attention_mask.unsqueeze(0).to_concrete(); - - // 3. Forward pass through encoder + let tokenized = self + .tokenizer + .tokenize_batch(texts, entity_labels, relation_labels)?; + let (token_ids, attention_mask) = self.build_batched_inputs(&tokenized); let encoder_output = self.encoder.forward(&token_ids, Some(&attention_mask)); - // 4. Extract word-level embeddings from encoder output, THEN apply BiLSTM - // (Python applies BiLSTM to word-level embeddings, not the full token sequence.) + let text_positions: Vec> = tokenized + .iter() + .map(|item| item.text_positions.clone()) + .collect(); let word_encoder_embs = - self.gather_at_positions(&encoder_output, &tokenized.text_positions); - let lstm_output = self.bilstm.forward(&word_encoder_embs).await; - - // 5. Extract label embeddings at marker positions from ENCODER output and project them - // (Labels are extracted from encoder output, text tokens from BiLSTM output) - // Entity label embeddings: hidden states at <> positions - let ent_embs_raw = self.gather_at_positions(&encoder_output, &tokenized.ent_positions); + self.gather_at_positions_batched(&encoder_output, &text_positions); + let word_lengths: Vec = tokenized.iter().map(|item| item.num_words).collect(); + let text_embs = self + .bilstm + .forward_with_lengths(&word_encoder_embs, &word_lengths) + .await; + + let ent_positions: Vec> = tokenized + .iter() + .map(|item| item.ent_positions.clone()) + .collect(); + let ent_embs_raw = self.gather_at_positions_batched(&encoder_output, &ent_positions); let ent_embs = self.prompt_rep_layer.forward_3d(&ent_embs_raw); - // Relation label embeddings: raw hidden states at <> positions - // (unlike entity labels, relation labels are NOT projected through prompt_rep_layer) - let rel_embs = self.gather_at_positions(&encoder_output, &tokenized.rel_positions); - - // 6. Text embeddings = BiLSTM output (already at word level) - let text_embs = lstm_output.clone(); + let rel_positions: Vec> = tokenized + .iter() + .map(|item| item.rel_positions.clone()) + .collect(); + let rel_embs = self.gather_at_positions_batched(&encoder_output, &rel_positions); - // 7–8. Decode entities using the mode matching the trained head. - let ent_embs_2d: Tensor<2, f32> = ent_embs.squeeze(0).to_concrete(); - let entities = match self.span_mode { + let entities_per_item = match self.span_mode { SpanMode::TokenLevel => { let scorer = self.scorer.as_ref().expect("token_level requires scorer"); - let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs_2d); - self.decode_entities_from_tokens( + let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs); + self.decode_entities_from_tokens_batch( &token_scores, entity_labels, - &tokenized.word_offsets, - text, + &tokenized, + texts, ) .await? } SpanMode::MarkerV0 => { - self.decode_entities_marker_v0( + self.decode_entities_marker_v0_batch( &text_embs, - &ent_embs_2d, + &ent_embs, entity_labels, - &tokenized.word_offsets, - tokenized.num_words, - text, + &tokenized, + texts, ) .await? } }; - // If no entities or no relation labels, return early - if entities.len() < 2 || relation_labels.is_empty() { - return Ok((entities, Vec::new())); - } - - // 9. Compute span representations for each entity using span_layer - // (matches Python's TokenMarker: project_start/project_end MLPs + out_project MLP) - let entity_spans: Vec<(usize, usize)> = entities - .iter() - .map(|e| (e.start_word, e.end_word)) - .collect(); - let span_reps = self - .span_layer - .forward_for_spans(&text_embs, &entity_spans, &self.device); - // span_reps shape: [num_entities, hidden] - - let num_entities = entities.len(); - let hidden_size = self.config.hidden_size; - - // 10. Build all entity pairs (head, tail) with head != tail - let mut candidate_pairs: Vec<(usize, usize)> = Vec::new(); - for head in 0..num_entities { - for tail in 0..num_entities { - if head != tail { - candidate_pairs.push((head, tail)); - } - } - } - - // 11. Gather head and tail span reps using index_select - let span_reps_data = span_reps.clone().as_slice().await?; - let span_reps_slice = span_reps_data.as_slice(); - let mut head_embs = Vec::with_capacity(candidate_pairs.len() * hidden_size); - let mut tail_embs = Vec::with_capacity(candidate_pairs.len() * hidden_size); - for &(head_idx, tail_idx) in &candidate_pairs { - let h_start = head_idx * hidden_size; - let t_start = tail_idx * hidden_size; - head_embs.extend_from_slice(&span_reps_slice[h_start..h_start + hidden_size]); - tail_embs.extend_from_slice(&span_reps_slice[t_start..t_start + hidden_size]); - } - - let head_tensor = Tensor::new(&self.device, &head_embs) - .reshape([candidate_pairs.len(), hidden_size]) - .to_concrete(); - let tail_tensor = Tensor::new(&self.device, &tail_embs) - .reshape([candidate_pairs.len(), hidden_size]) - .to_concrete(); - - // 12. Apply pair_projector: concat(head, tail) -> MLP -> pair_rep - let pair_embs = self.pair_projector.forward(&head_tensor, &tail_tensor); - - // 13. Score pairs against relation labels via dot product (no sigmoid yet) - let rel_embs_squeezed: Tensor<2, f32> = rel_embs.squeeze(0).to_concrete(); - let rel_scores = pair_embs.mat_mul(&rel_embs_squeezed.transpose(0, 1)); - - // 14. Apply sigmoid and filter by relation_threshold - let rel_scores_slice = rel_scores.clone().as_slice().await?; - let n_rels = relation_labels.len(); - let mut relations = Vec::new(); - let threshold = self.config.relation_threshold; - - for (pair_idx, &(head_idx, tail_idx)) in candidate_pairs.iter().enumerate() { - let base = pair_idx * n_rels; - for rel_idx in 0..n_rels { - let raw = rel_scores_slice.as_slice()[base + rel_idx]; - let prob = 1.0 / (1.0 + (-raw).exp()); - if prob > threshold { - relations.push(Relation { - head: entities[head_idx].clone(), - tail: entities[tail_idx].clone(), - relation: relation_labels[rel_idx].to_string(), - score: prob, - }); - } - } + let mut results = Vec::with_capacity(texts.len()); + for (batch_idx, entities) in entities_per_item.into_iter().enumerate() { + let relations = if entities.len() < 2 || relation_labels.is_empty() { + Vec::new() + } else { + let text_embs_item: Tensor<3, f32> = + text_embs.narrow(0, batch_idx, 1).to_concrete(); + let rel_embs_item: Tensor<3, f32> = + rel_embs.narrow(0, batch_idx, 1).to_concrete(); + self.decode_relations( + &text_embs_item, + &rel_embs_item, + &entities, + relation_labels, + ) + .await? + }; + results.push((entities, relations)); } - // Sort by score descending - relations.sort_by(|a, b| { - b.score - .partial_cmp(&a.score) - .unwrap_or(std::cmp::Ordering::Equal) - }); - - Ok((entities, relations)) + Ok(results) } /// Extract entities and relations from text, chunking the input first so long @@ -628,12 +569,14 @@ impl GlinerRelEx { ent.end_char += offset; }; + let chunk_texts: Vec<&str> = ranges.iter().map(|range| &text[range.clone()]).collect(); + let per_chunk = self + .extract_batch(&chunk_texts, entity_labels, relation_labels) + .await?; + let mut all_entities: Vec = Vec::new(); let mut all_relations: Vec = Vec::new(); - for range in &ranges { - let chunk = &text[range.clone()]; - let (entities, relations) = - self.extract(chunk, entity_labels, relation_labels).await?; + for (range, (entities, relations)) in ranges.iter().zip(per_chunk) { let offset = range.start; for mut ent in entities { shift(&mut ent, offset); @@ -689,13 +632,170 @@ impl GlinerRelEx { Ok((all_entities, all_relations)) } - /// Decode entities for `span_mode = markerV0` (used by the `large` variants). - /// - /// Enumerates every `(start, end)` pair up to `config.max_width` words, - /// computes the span representation via `SpanLayer::forward_for_spans`, - /// scores each span against every projected entity prompt via a dot - /// product, applies sigmoid + `entity_threshold`, and greedy-filters - /// overlapping spans (keeping the highest-scoring one). + fn build_batched_inputs( + &self, + tokenized: &[RelExTokenizedInput], + ) -> (Tensor<2, u32>, Tensor<2, u32>) { + let batch_size = tokenized.len(); + let max_seq_len = tokenized + .iter() + .map(|item| item.token_ids.len()) + .max() + .unwrap_or(1) + .max(1); + let pad_id = self.tokenizer.special_tokens().pad_id; + + let mut token_ids = vec![pad_id; batch_size * max_seq_len]; + let mut attention_mask = vec![0u32; batch_size * max_seq_len]; + for (batch_idx, item) in tokenized.iter().enumerate() { + let offset = batch_idx * max_seq_len; + let len = item.token_ids.len(); + token_ids[offset..offset + len].copy_from_slice(&item.token_ids); + attention_mask[offset..offset + len].copy_from_slice(&item.attention_mask); + } + + let token_ids = Tensor::new(&self.device, &token_ids) + .reshape([batch_size, max_seq_len]) + .to_concrete(); + let attention_mask = Tensor::new(&self.device, &attention_mask) + .reshape([batch_size, max_seq_len]) + .to_concrete(); + (token_ids, attention_mask) + } + + async fn decode_relations( + &self, + text_embs: &Tensor<3, f32>, + rel_embs: &Tensor<3, f32>, + entities: &[Entity], + relation_labels: &[&str], + ) -> Result, GlinerError> { + if entities.len() < 2 || relation_labels.is_empty() { + return Ok(Vec::new()); + } + + let entity_spans: Vec<(usize, usize)> = entities + .iter() + .map(|e| (e.start_word, e.end_word)) + .collect(); + let span_reps = self + .span_layer + .forward_for_spans(text_embs, &entity_spans, &self.device); + + let num_entities = entities.len(); + let hidden_size = self.config.hidden_size; + let mut candidate_pairs: Vec<(usize, usize)> = Vec::new(); + for head in 0..num_entities { + for tail in 0..num_entities { + if head != tail { + candidate_pairs.push((head, tail)); + } + } + } + + let span_reps_data = span_reps.clone().as_slice().await?; + let span_reps_slice = span_reps_data.as_slice(); + let mut head_embs = Vec::with_capacity(candidate_pairs.len() * hidden_size); + let mut tail_embs = Vec::with_capacity(candidate_pairs.len() * hidden_size); + for &(head_idx, tail_idx) in &candidate_pairs { + let h_start = head_idx * hidden_size; + let t_start = tail_idx * hidden_size; + head_embs.extend_from_slice(&span_reps_slice[h_start..h_start + hidden_size]); + tail_embs.extend_from_slice(&span_reps_slice[t_start..t_start + hidden_size]); + } + + let head_tensor = Tensor::new(&self.device, &head_embs) + .reshape([candidate_pairs.len(), hidden_size]) + .to_concrete(); + let tail_tensor = Tensor::new(&self.device, &tail_embs) + .reshape([candidate_pairs.len(), hidden_size]) + .to_concrete(); + let pair_embs = self.pair_projector.forward(&head_tensor, &tail_tensor); + + let rel_embs_squeezed: Tensor<2, f32> = rel_embs.squeeze(0).to_concrete(); + let rel_scores = pair_embs.mat_mul(&rel_embs_squeezed.transpose(0, 1)); + let rel_scores_slice = rel_scores.clone().as_slice().await?; + let n_rels = relation_labels.len(); + let threshold = self.config.relation_threshold; + + let mut relations = Vec::new(); + for (pair_idx, &(head_idx, tail_idx)) in candidate_pairs.iter().enumerate() { + let base = pair_idx * n_rels; + for rel_idx in 0..n_rels { + let raw = rel_scores_slice.as_slice()[base + rel_idx]; + let prob = 1.0 / (1.0 + (-raw).exp()); + if prob > threshold { + relations.push(Relation { + head: entities[head_idx].clone(), + tail: entities[tail_idx].clone(), + relation: relation_labels[rel_idx].to_string(), + score: prob, + }); + } + } + } + + relations.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); + Ok(relations) + } + + async fn decode_entities_marker_v0_batch( + &self, + text_embs: &Tensor<3, f32>, + ent_embs: &Tensor<3, f32>, + entity_labels: &[&str], + tokenized: &[RelExTokenizedInput], + texts: &[&str], + ) -> Result>, GlinerError> { + let spans_per_batch: Vec> = tokenized + .iter() + .map(|item| { + let mut spans = Vec::new(); + for start in 0..item.num_words { + for width in 1..=self.config.max_width.min(item.num_words - start) { + spans.push((start, start + width - 1)); + } + } + spans + }) + .collect(); + + let (flat_span_reps, span_counts) = + self.span_layer + .forward_for_spans_batched(text_embs, &spans_per_batch, &self.device); + + let mut offset = 0usize; + let mut results = Vec::with_capacity(tokenized.len()); + for batch_idx in 0..tokenized.len() { + let span_count = span_counts[batch_idx]; + let entities = if span_count == 0 || entity_labels.is_empty() { + Vec::new() + } else { + let span_reps: Tensor<2, f32> = + flat_span_reps.narrow(0, offset, span_count).to_concrete(); + let ent_embs_2d: Tensor<2, f32> = + ent_embs.narrow(0, batch_idx, 1).squeeze(0).to_concrete(); + self.decode_entities_marker_v0_from_span_reps( + &span_reps, + &spans_per_batch[batch_idx], + &ent_embs_2d, + entity_labels, + &tokenized[batch_idx].word_offsets, + texts[batch_idx], + ) + .await? + }; + results.push(entities); + offset += span_count; + } + + Ok(results) + } + async fn decode_entities_marker_v0( &self, text_embs: &Tensor<3, f32>, @@ -705,36 +805,51 @@ impl GlinerRelEx { num_words: usize, text: &str, ) -> Result, GlinerError> { - let threshold = self.config.entity_threshold; - let max_width = self.config.max_width; - let hidden = self.config.hidden_size; - let n_labels = entity_labels.len(); - - if num_words == 0 || n_labels == 0 { + if num_words == 0 || entity_labels.is_empty() { return Ok(Vec::new()); } - // Enumerate spans: (start, end) with end-start+1 <= max_width. - let mut spans: Vec<(usize, usize)> = Vec::new(); + let mut spans = Vec::new(); for start in 0..num_words { - for width in 1..=max_width.min(num_words - start) { + for width in 1..=self.config.max_width.min(num_words - start) { spans.push((start, start + width - 1)); } } - // Compute span reps [num_spans, hidden]. let span_reps = self .span_layer .forward_for_spans(text_embs, &spans, &self.device); + self.decode_entities_marker_v0_from_span_reps( + &span_reps, + &spans, + ent_embs_2d, + entity_labels, + word_offsets, + text, + ) + .await + } + + async fn decode_entities_marker_v0_from_span_reps( + &self, + span_reps: &Tensor<2, f32>, + spans: &[(usize, usize)], + ent_embs_2d: &Tensor<2, f32>, + entity_labels: &[&str], + word_offsets: &[(usize, usize)], + text: &str, + ) -> Result, GlinerError> { + let threshold = self.config.entity_threshold; + let n_labels = entity_labels.len(); + if spans.is_empty() || n_labels == 0 { + return Ok(Vec::new()); + } - // Score: [num_spans, hidden] @ [hidden, n_labels] -> [num_spans, n_labels]. let label_rep_t: Tensor<2, f32> = ent_embs_2d.transpose(0, 1).to_concrete(); let logits = span_reps.mat_mul(&label_rep_t); - let logits_data = logits.clone().as_slice().await?; let logits_slice = logits_data.as_slice(); - // Candidate (start, end, label, score) above threshold. let mut candidates: Vec<(usize, usize, usize, f32)> = Vec::new(); for (span_idx, &(s, e)) in spans.iter().enumerate() { for l in 0..n_labels { @@ -746,7 +861,6 @@ impl GlinerRelEx { } } - // Sort by score descending and greedy non-overlapping filter. candidates.sort_by(|a, b| b.3.partial_cmp(&a.3).unwrap_or(std::cmp::Ordering::Equal)); let mut taken: Vec<(usize, usize)> = Vec::new(); @@ -772,20 +886,45 @@ impl GlinerRelEx { } } - // Ensure output is sorted by score descending for presentation. entities.sort_by(|a, b| { b.score .partial_cmp(&a.score) .unwrap_or(std::cmp::Ordering::Equal) }); - let _ = hidden; Ok(entities) } + async fn decode_entities_from_tokens_batch( + &self, + token_scores: &Tensor<4, f32>, + entity_labels: &[&str], + tokenized: &[RelExTokenizedInput], + texts: &[&str], + ) -> Result>, GlinerError> { + let [batch_size, padded_tokens, num_labels, num_channels] = token_scores.shape(); + assert_eq!(num_channels, 3, "expected [start, end, inside]"); + assert_eq!(batch_size, tokenized.len(), "tokenized batch size mismatch"); + let scores_data = token_scores.clone().as_slice().await?; + let scores = scores_data.as_slice(); + let batch_stride = padded_tokens * num_labels * 3; + + let mut results = Vec::with_capacity(batch_size); + for batch_idx in 0..batch_size { + let start = batch_idx * batch_stride; + let end = start + batch_stride; + results.push(self.decode_entities_from_tokens_slice( + &scores[start..end], + tokenized[batch_idx].num_words, + num_labels, + entity_labels, + &tokenized[batch_idx].word_offsets, + texts[batch_idx], + )); + } + Ok(results) + } + /// Decode entities using span-boundary detection with start/end/inside scores. - /// - /// `token_scores` has shape [batch, seq_len, n_labels, 3] where the last dim - /// is [start, end, inside] sigmoid probabilities. async fn decode_entities_from_tokens( &self, token_scores: &Tensor<4, f32>, @@ -796,11 +935,26 @@ impl GlinerRelEx { let [_batch_size, num_tokens, num_labels, num_channels] = token_scores.shape(); assert_eq!(num_channels, 3, "expected [start, end, inside]"); let scores_data = token_scores.clone().as_slice().await?; - let scores = scores_data.as_slice(); + Ok(self.decode_entities_from_tokens_slice( + scores_data.as_slice(), + num_tokens, + num_labels, + entity_labels, + word_offsets, + text, + )) + } + fn decode_entities_from_tokens_slice( + &self, + scores: &[f32], + num_tokens: usize, + num_labels: usize, + entity_labels: &[&str], + word_offsets: &[(usize, usize)], + text: &str, + ) -> Vec { let threshold = self.config.entity_threshold; - - // Candidate spans: (start, end, label, score) let mut candidates: Vec<(usize, usize, usize, f32)> = Vec::new(); let score_at = |tok: usize, lab: usize, ch: usize| -> f32 { @@ -820,7 +974,6 @@ impl GlinerRelEx { continue; } - // Check all inside scores from start_tok to end_tok let mut min_score = start_score.min(end_score); let mut valid = true; for t in start_tok..=end_tok { @@ -833,19 +986,15 @@ impl GlinerRelEx { min_score = inside; } } - if !valid { - continue; + if valid { + candidates.push((start_tok, end_tok, label_idx, min_score)); } - - candidates.push((start_tok, end_tok, label_idx, min_score)); } } } - // Sort candidates by score descending candidates.sort_by(|a, b| b.3.partial_cmp(&a.3).unwrap_or(std::cmp::Ordering::Equal)); - // Greedy filter non-overlapping spans (flat_ner equivalent) let mut taken: Vec<(usize, usize)> = Vec::new(); let mut entities = Vec::new(); for (start_tok, end_tok, label_idx, score) in candidates { @@ -870,38 +1019,60 @@ impl GlinerRelEx { } } - // Sort by score descending entities.sort_by(|a, b| { b.score .partial_cmp(&a.score) .unwrap_or(std::cmp::Ordering::Equal) }); - - Ok(entities) + entities } - /// Gather hidden states at specific positions. - fn gather_at_positions( + /// Gather hidden states at specific positions for a whole batch. + fn gather_at_positions_batched( &self, hidden_states: &Tensor<3, f32>, - positions: &[usize], + positions_per_batch: &[Vec], ) -> Tensor<3, f32> { - let [batch_size, _seq_len, hidden_size] = hidden_states.shape(); - let num_positions = positions.len(); - - if num_positions == 0 { + let [batch_size, seq_len, hidden_size] = hidden_states.shape(); + assert_eq!( + batch_size, + positions_per_batch.len(), + "positions_per_batch must match batch size" + ); + + let max_positions = positions_per_batch.iter().map(Vec::len).max().unwrap_or(0); + if max_positions == 0 { return Tensor::zeros(&self.device, [batch_size, 1, hidden_size]); } - // Build index tensor - let indices: Vec = positions.iter().map(|&p| p as u32).collect(); - let index_tensor = Tensor::new(&self.device, &indices); + let hidden_flat = hidden_states + .to_concrete() + .reshape([batch_size * seq_len, hidden_size]) + .to_concrete(); + + let mut offset_indices = Vec::with_capacity(batch_size * max_positions); + for (batch_idx, positions) in positions_per_batch.iter().enumerate() { + let offset = (batch_idx * seq_len) as u32; + for pos_idx in 0..max_positions { + let pos = positions.get(pos_idx).copied().unwrap_or(0) as u32; + offset_indices.push(pos + offset); + } + } - // For batch size 1, we can use index_select - let hidden_2d = hidden_states.squeeze(0).to_concrete(); - let gathered = hidden_2d.index_select(0, &index_tensor); + let offset_indices = Tensor::new(&self.device, &offset_indices); + hidden_flat + .index_select(0, &offset_indices) + .reshape([batch_size, max_positions, hidden_size]) + .to_concrete() + } - gathered.unsqueeze(0).to_concrete() + /// Gather hidden states at specific positions. + fn gather_at_positions( + &self, + hidden_states: &Tensor<3, f32>, + positions: &[usize], + ) -> Tensor<3, f32> { + self.gather_at_positions_batched(hidden_states, &[positions.to_vec()]) } /// Get the device. @@ -918,7 +1089,6 @@ impl GlinerRelEx { #[cfg(test)] mod tests { use super::*; - use std::path::PathBuf; use std::time::{Duration, Instant}; const PROFILE_TEXT: &str = "Apple Inc. was founded by Steve Jobs in Cupertino, California. \ @@ -1000,18 +1170,12 @@ Meta Platforms was founded by Mark Zuckerberg in Cambridge, Massachusetts."; } } - fn weights_path(file_name: &str) -> PathBuf { - PathBuf::from(env!("CARGO_MANIFEST_DIR")) - .join("weights") - .join(file_name) - } - - async fn load_local_relex( - model_path: PathBuf, + async fn load_relex( + source: GlinerRelExSource, device: Device, ) -> Result { GlinerRelEx::builder() - .with_source(GlinerRelExSource::local(model_path)) + .with_source(source) .with_device(device) .build_with_loading_handler(|_| {}) .await @@ -1050,7 +1214,6 @@ Meta Platforms was founded by Mark Zuckerberg in Cambridge, Massachusetts."; let ent_embs = model.prompt_rep_layer.forward_3d(&ent_embs_raw); let rel_embs = model.gather_at_positions(&encoder_output, &tokenized.rel_positions); let text_embs = lstm_output.clone(); - let ent_embs_2d: Tensor<2, f32> = ent_embs.squeeze(0).to_concrete(); let (entities, span_count, entity_span_prep_cpu, entity_sync, entity_decode_cpu) = match model.span_mode { @@ -1059,7 +1222,7 @@ Meta Platforms was founded by Mark Zuckerberg in Cambridge, Massachusetts."; .scorer .as_ref() .expect("token_level requires scorer"); - let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs_2d); + let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs); let (entities, entity_sync, entity_decode_cpu) = profile_decode_entities_from_tokens( model, @@ -1079,6 +1242,7 @@ Meta Platforms was founded by Mark Zuckerberg in Cambridge, Massachusetts."; ) } SpanMode::MarkerV0 => { + let ent_embs_2d: Tensor<2, f32> = ent_embs.squeeze(0).to_concrete(); profile_decode_entities_marker_v0( model, &text_embs, @@ -1393,13 +1557,13 @@ Meta Platforms was founded by Mark Zuckerberg in Cambridge, Massachusetts."; }; let variants = [ - ("multi", "gliner-relex-multi-v1.0-Q4_K.gguf"), - ("base", "gliner-relex-base-v1.0-Q4_K.gguf"), - ("large", "gliner-relex-large-v1.0-Q4_K.gguf"), + ("multi", GlinerRelExSource::relex_multi()), + ("base", GlinerRelExSource::relex_base()), + ("large", GlinerRelExSource::relex_large()), ]; - for (variant, file_name) in variants { - let model = load_local_relex(weights_path(file_name), device.clone()).await?; + for (variant, source) in variants { + let model = load_relex(source, device.clone()).await?; let cold_start = Instant::now(); let _ = model @@ -1413,4 +1577,340 @@ Meta Platforms was founded by Mark Zuckerberg in Cambridge, Massachusetts."; Ok(()) } + + fn entity_signature(entities: &[Entity]) -> Vec<(String, String, usize, usize, usize, usize)> { + entities + .iter() + .map(|entity| { + ( + entity.label.clone(), + entity.text.clone(), + entity.start_char, + entity.end_char, + entity.start_word, + entity.end_word, + ) + }) + .collect() + } + + fn relation_signature( + relations: &[Relation], + ) -> Vec<(String, String, String, usize, usize, usize, usize)> { + relations + .iter() + .map(|relation| { + ( + relation.head.text.clone(), + relation.tail.text.clone(), + relation.relation.clone(), + relation.head.start_char, + relation.head.end_char, + relation.tail.start_char, + relation.tail.end_char, + ) + }) + .collect() + } + + async fn assert_batch_matches_serial_extract( + variant: &'static str, + source: GlinerRelExSource, + ) -> Result<(), Box> { + let device = Device::gpu().await.unwrap_or_else(|_| Device::cpu()); + let texts = [ + "Apple was founded by Steve Jobs.", + "Google was founded by Larry Page in Mountain View.", + ]; + let model = load_relex(source, device).await?; + + let mut serial_results = Vec::with_capacity(texts.len()); + for text in texts.iter().copied() { + serial_results.push(model.extract(text, ENTITY_LABELS, RELATION_LABELS).await?); + } + let batched_results = model + .extract_batch(&texts, ENTITY_LABELS, RELATION_LABELS) + .await?; + + assert_eq!( + serial_results.len(), + batched_results.len(), + "batch size mismatch for {variant}" + ); + + for ((serial_entities, serial_relations), (batched_entities, batched_relations)) in + serial_results.iter().zip(&batched_results) + { + assert_eq!( + entity_signature(serial_entities), + entity_signature(batched_entities), + "entity mismatch for {variant}" + ); + assert_eq!( + relation_signature(serial_relations), + relation_signature(batched_relations), + "relation mismatch for {variant}" + ); + assert_eq!( + serial_entities.len(), + batched_entities.len(), + "entity count mismatch for {variant}" + ); + assert_eq!( + serial_relations.len(), + batched_relations.len(), + "relation count mismatch for {variant}" + ); + + for (serial_entity, batched_entity) in serial_entities.iter().zip(batched_entities.iter()) + { + assert!( + (serial_entity.score - batched_entity.score).abs() < 1e-5, + "entity score mismatch for {variant}: serial={:.6} batched={:.6}", + serial_entity.score, + batched_entity.score + ); + } + for (serial_relation, batched_relation) in + serial_relations.iter().zip(batched_relations.iter()) + { + assert!( + (serial_relation.score - batched_relation.score).abs() < 1e-5, + "relation score mismatch for {variant}: serial={:.6} batched={:.6}", + serial_relation.score, + batched_relation.score + ); + } + } + + Ok(()) + } + + #[tokio::test] + async fn extract_batch_matches_serial_extract_for_remote_multi( + ) -> Result<(), Box> { + assert_batch_matches_serial_extract("multi", GlinerRelExSource::relex_multi()).await + } + + #[tokio::test] + async fn extract_batch_matches_serial_extract_for_remote_large( + ) -> Result<(), Box> { + assert_batch_matches_serial_extract("large", GlinerRelExSource::relex_large()).await + } + + #[tokio::test] + #[ignore = "expensive remote coverage across all rel-ex variants"] + async fn extract_batch_matches_serial_extract_for_remote_variants( + ) -> Result<(), Box> { + for (variant, source) in [ + ("multi", GlinerRelExSource::relex_multi()), + ("base", GlinerRelExSource::relex_base()), + ("large", GlinerRelExSource::relex_large()), + ] { + assert_batch_matches_serial_extract(variant, source).await?; + } + Ok(()) + } + + #[tokio::test] + #[ignore = "cache the remote rel-ex checkpoints"] + async fn cache_remote_relex_variants() -> Result<(), Box> { + let device = Device::gpu().await.unwrap_or_else(|_| Device::cpu()); + for source in [ + GlinerRelExSource::relex_multi(), + GlinerRelExSource::relex_base(), + GlinerRelExSource::relex_large(), + ] { + let _model = load_relex(source, device.clone()).await?; + } + Ok(()) + } + + #[tokio::test] + #[ignore = "smoke-test batched rel-ex on GPU"] + async fn extract_batch_remote_multi_gpu_smoke() -> Result<(), Box> { + let device = Device::gpu().await?; + let model = load_relex(GlinerRelExSource::relex_multi(), device).await?; + let texts = [ + "Apple was founded by Steve Jobs.", + "Google was founded by Larry Page in Mountain View.", + ]; + + let results = model + .extract_batch(&texts, ENTITY_LABELS, RELATION_LABELS) + .await?; + assert_eq!(results.len(), texts.len()); + Ok(()) + } + + #[tokio::test] + #[ignore = "diagnose first failing GPU stage for remote rel-ex multi"] + async fn debug_remote_multi_gpu_stage_cutoff() -> Result<(), Box> { + let device = Device::gpu().await?; + let model = load_relex(GlinerRelExSource::relex_multi(), device.clone()).await?; + let tokenized = model + .tokenizer + .tokenize(PROFILE_TEXT, ENTITY_LABELS, RELATION_LABELS)?; + + println!( + "stage_cutoff: device={} seq_len={} words={} ents={} rels={}", + if device.is_gpu() { "gpu" } else { "cpu" }, + tokenized.token_ids.len(), + tokenized.num_words, + tokenized.num_entity_labels, + tokenized.num_relation_labels + ); + + let token_ids = Tensor::new(&device, &tokenized.token_ids); + let token_ids: Tensor<2, u32> = token_ids.unsqueeze(0).to_concrete(); + + let attention_mask = Tensor::new(&device, &tokenized.attention_mask); + let attention_mask: Tensor<2, u32> = attention_mask.unsqueeze(0).to_concrete(); + + println!("stage_cutoff: materializing post-embedding-norm"); + let post_embedding = model.encoder.debug_after_embedding_norm(&token_ids); + let _ = post_embedding.clone().as_slice().await?; + println!("stage_cutoff: post-embedding-norm ok"); + + println!("stage_cutoff: materializing first encoder layer"); + let first_layer = model + .encoder + .debug_first_layer_output(&post_embedding, Some(&attention_mask)); + let _ = first_layer.clone().as_slice().await?; + println!("stage_cutoff: first encoder layer ok"); + + println!("stage_cutoff: materializing full encoder"); + let full_encoder = model.encoder.forward(&token_ids, Some(&attention_mask)); + let _ = full_encoder.as_slice().await?; + println!("stage_cutoff: full encoder ok"); + + Ok(()) + } + + #[tokio::test] + #[ignore = "diagnose first failing GPU stage for batched remote rel-ex multi"] + async fn debug_remote_multi_gpu_batched_stage_cutoff( + ) -> Result<(), Box> { + let device = Device::gpu().await?; + let model = load_relex(GlinerRelExSource::relex_multi(), device.clone()).await?; + let texts = [ + "Apple was founded by Steve Jobs.", + "Google was founded by Larry Page in Mountain View.", + ]; + let tokenized = model + .tokenizer + .tokenize_batch(&texts, ENTITY_LABELS, RELATION_LABELS)?; + let seq_lens: Vec = tokenized.iter().map(|item| item.token_ids.len()).collect(); + let word_lengths: Vec = tokenized.iter().map(|item| item.num_words).collect(); + println!( + "batched_stage_cutoff: device={} batch={} seq_lens={seq_lens:?} word_lengths={word_lengths:?}", + if device.is_gpu() { "gpu" } else { "cpu" }, + texts.len(), + ); + + let (token_ids, attention_mask) = model.build_batched_inputs(&tokenized); + + println!("batched_stage_cutoff: materializing encoder output"); + let encoder_output = model.encoder.forward(&token_ids, Some(&attention_mask)); + let _ = encoder_output.clone().as_slice().await?; + println!("batched_stage_cutoff: encoder output ok"); + + let text_positions: Vec> = tokenized + .iter() + .map(|item| item.text_positions.clone()) + .collect(); + println!("batched_stage_cutoff: materializing word encoder embeddings"); + let word_encoder_embs = model.gather_at_positions_batched(&encoder_output, &text_positions); + let _ = word_encoder_embs.clone().as_slice().await?; + println!("batched_stage_cutoff: word encoder embeddings ok"); + + println!("batched_stage_cutoff: materializing BiLSTM output"); + let uniform_lengths = vec![word_lengths.iter().copied().max().unwrap_or(0); word_lengths.len()]; + println!( + "batched_stage_cutoff: materializing BiLSTM output with uniform lengths {uniform_lengths:?}" + ); + println!("batched_stage_cutoff: materializing first-step gates only"); + let first_step_gates = model + .bilstm + .debug_first_step_gates(&word_encoder_embs, false); + let _ = first_step_gates.as_slice().await?; + println!("batched_stage_cutoff: first-step gates ok"); + + println!("batched_stage_cutoff: materializing forward direction state-only"); + let fwd_state = model + .bilstm + .debug_forward_direction_state_only(&word_encoder_embs, &uniform_lengths, false); + let _ = fwd_state.clone().as_slice().await?; + println!("batched_stage_cutoff: forward direction state-only ok"); + + println!("batched_stage_cutoff: materializing repeated slice_assign only"); + let hidden = fwd_state.shape()[1]; + let mut stitched: Tensor<3, f32> = + Tensor::zeros(&device, [texts.len(), uniform_lengths[0], hidden]); + let zero_step: Tensor<3, f32> = Tensor::zeros(&device, [texts.len(), 1, hidden]); + for t in 0..uniform_lengths[0] { + stitched = stitched.slice_assign([0..texts.len(), t..(t + 1), 0..hidden], &zero_step); + } + let _ = stitched.as_slice().await?; + println!("batched_stage_cutoff: repeated slice_assign ok"); + + println!("batched_stage_cutoff: materializing forward direction only"); + let fwd_only = model + .bilstm + .debug_forward_direction(&word_encoder_embs, &uniform_lengths, false); + let _ = fwd_only.clone().as_slice().await?; + println!("batched_stage_cutoff: forward direction ok"); + + println!("batched_stage_cutoff: materializing backward direction only"); + let bwd_only = model + .bilstm + .debug_forward_direction(&word_encoder_embs, &uniform_lengths, true); + let _ = bwd_only.clone().as_slice().await?; + println!("batched_stage_cutoff: backward direction ok"); + + println!("batched_stage_cutoff: materializing full BiLSTM with uniform lengths"); + let uniform_text_embs = model + .bilstm + .forward_with_lengths(&word_encoder_embs, &uniform_lengths) + .await; + let _ = uniform_text_embs.clone().as_slice().await?; + println!("batched_stage_cutoff: full BiLSTM with uniform lengths ok"); + + println!( + "batched_stage_cutoff: materializing BiLSTM output with actual lengths {word_lengths:?}" + ); + let text_embs = model + .bilstm + .forward_with_lengths(&word_encoder_embs, &word_lengths) + .await; + let _ = text_embs.clone().as_slice().await?; + println!("batched_stage_cutoff: BiLSTM output ok"); + + let ent_positions: Vec> = tokenized + .iter() + .map(|item| item.ent_positions.clone()) + .collect(); + println!("batched_stage_cutoff: materializing entity prompt embeddings"); + let ent_embs_raw = model.gather_at_positions_batched(&encoder_output, &ent_positions); + let ent_embs = model.prompt_rep_layer.forward_3d(&ent_embs_raw); + let _ = ent_embs.clone().as_slice().await?; + println!("batched_stage_cutoff: entity prompt embeddings ok"); + + let rel_positions: Vec> = tokenized + .iter() + .map(|item| item.rel_positions.clone()) + .collect(); + println!("batched_stage_cutoff: materializing relation embeddings"); + let rel_embs = model.gather_at_positions_batched(&encoder_output, &rel_positions); + let _ = rel_embs.clone().as_slice().await?; + println!("batched_stage_cutoff: relation embeddings ok"); + + println!("batched_stage_cutoff: materializing token scores"); + let scorer = model.scorer.as_ref().expect("token_level requires scorer"); + let token_scores = scorer.forward_entity_scores(&text_embs, &ent_embs); + let _ = token_scores.clone().as_slice().await?; + println!("batched_stage_cutoff: token scores ok"); + + Ok(()) + } } diff --git a/models/rgliner/src/relex_tokenization.rs b/models/rgliner/src/relex_tokenization.rs index 2e429abfc..b3c553a3b 100644 --- a/models/rgliner/src/relex_tokenization.rs +++ b/models/rgliner/src/relex_tokenization.rs @@ -191,6 +191,18 @@ impl RelExTokenizer { }) } + /// Tokenize a batch of texts with a shared label prompt. + pub fn tokenize_batch( + &self, + texts: &[&str], + entity_labels: &[&str], + relation_labels: &[&str], + ) -> Result, GlinerError> { + texts.iter() + .map(|text| self.tokenize(text, entity_labels, relation_labels)) + .collect() + } + /// Split text into words with character offsets. /// /// Matches Python GLiNER's `WhitespaceTokenSplitter` regex: diff --git a/models/rgliner/src/tokenization.rs b/models/rgliner/src/tokenization.rs index 6885e46f0..5da595757 100644 --- a/models/rgliner/src/tokenization.rs +++ b/models/rgliner/src/tokenization.rs @@ -82,6 +82,16 @@ impl WordTokenizer { word_offsets, }) } + + /// Tokenize a batch of texts. + pub fn tokenize_batch(&self, texts: &[&str]) -> Result, GlinerError> { + texts.iter().map(|text| self.tokenize(text)).collect() + } + + /// Resolve the tokenizer's padding ID. + pub fn pad_id(&self) -> u32 { + self.tokenizer.token_to_id("[PAD]").unwrap_or(0) + } } /// Pack `text` into token-budgeted byte ranges using the supplied tokenizer. diff --git a/models/rgliner/tests/example_regression.rs b/models/rgliner/tests/example_regression.rs index 6e5182512..46669e55f 100644 --- a/models/rgliner/tests/example_regression.rs +++ b/models/rgliner/tests/example_regression.rs @@ -1,20 +1,44 @@ -use std::path::PathBuf; - use fusor::Device; +use kalosm_model_types::FileSource; use rgliner::{Gliner, GlinerSource}; -fn local_edge_source() -> GlinerSource { - let manifest_dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")); - let weights_dir = manifest_dir.join("weights"); - - GlinerSource::local( - weights_dir.join("gliner-edge.gguf"), - weights_dir.join("gliner-edge-label-encoder.gguf"), +fn remote_edge_source() -> GlinerSource { + GlinerSource::custom( + FileSource::huggingface( + "Demonthos/gliner-gguf".to_string(), + "main".to_string(), + "gliner-bi-edge-v2.0-Q4_K.gguf".to_string(), + ), + FileSource::huggingface( + "Demonthos/gliner-gguf".to_string(), + "main".to_string(), + "gliner-bi-edge-v2.0-Q4_K-label-encoder.gguf".to_string(), + ), + FileSource::huggingface( + "sentence-transformers/all-MiniLM-L6-v2".to_string(), + "main".to_string(), + "config.json".to_string(), + ), + FileSource::huggingface( + "sentence-transformers/all-MiniLM-L6-v2".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + FileSource::huggingface( + "knowledgator/gliner-bi-edge-v2.0".to_string(), + "main".to_string(), + "tokenizer.json".to_string(), + ), + FileSource::huggingface( + "knowledgator/gliner-bi-edge-v2.0".to_string(), + "main".to_string(), + "gliner_config.json".to_string(), + ), ) } #[test] -fn edge_example_sentences_regression() -> anyhow::Result<()> { +fn remote_edge_cached_labels_match_uncached_extract() -> anyhow::Result<()> { tokio::runtime::Builder::new_multi_thread() .enable_all() .build()? @@ -23,64 +47,155 @@ fn edge_example_sentences_regression() -> anyhow::Result<()> { // environment can interfere with the auto-device probe, while the plain // example covers the user-facing default path separately. let mut gliner = Gliner::builder() - .with_source(local_edge_source()) + .with_source(remote_edge_source()) .with_device(Device::cpu()) .build() .await?; let labels = ["person", "organization", "location"]; let cases = [ - ( - "Apple Inc. was founded by Steve Jobs in California.", - vec![ - ("organization", "Apple Inc."), - ("person", "Steve Jobs"), - ("location", "California"), - ], - ), - ( - "Microsoft Corporation is headquartered in Seattle.", - vec![ - ("organization", "Microsoft Corporation"), - ("location", "Seattle"), - ], - ), - ( - "Elon Musk is the CEO of Tesla.", - vec![("person", "Elon Musk"), ("organization", "Tesla")], - ), - ( - "Google was founded in Mountain View.", - vec![("organization", "Google"), ("location", "Mountain View")], - ), + "Apple Inc. was founded by Steve Jobs in California.", + "Microsoft Corporation is headquartered in Seattle.", + "Elon Musk is the CEO of Tesla.", + "Google was founded in Mountain View.", ]; - for (text, expected) in cases { + for text in cases { let uncached_entities = gliner.extract(text, &labels).await?; - let uncached: Vec<(&str, &str)> = uncached_entities + let uncached: Vec<(String, String, usize, usize, f32)> = uncached_entities .iter() - .map(|entity| (entity.label.as_str(), entity.text.as_str())) + .map(|entity| { + ( + entity.label.clone(), + entity.text.clone(), + entity.start_char, + entity.end_char, + entity.score, + ) + }) .collect(); - assert_eq!( - uncached, expected, - "unexpected uncached entities for input: {text}" - ); - gliner.cache_labels(&labels).await?; let entities = gliner.extract_with_cached_labels(text).await?; - let actual: Vec<(&str, &str)> = entities + let cached: Vec<(String, String, usize, usize, f32)> = entities .iter() - .map(|entity| (entity.label.as_str(), entity.text.as_str())) + .map(|entity| { + ( + entity.label.clone(), + entity.text.clone(), + entity.start_char, + entity.end_char, + entity.score, + ) + }) .collect(); - assert_eq!(actual, expected, "unexpected entities for input: {text}"); - assert!( - entities.iter().all(|entity| entity.score >= 0.5), - "all expected entities should remain above the default threshold for input: {text}" + assert_eq!(uncached.len(), cached.len(), "entity count mismatch for input: {text}"); + for (uncached_entity, cached_entity) in uncached.iter().zip(&cached) { + assert_eq!(uncached_entity.0, cached_entity.0, "label mismatch for input: {text}"); + assert_eq!(uncached_entity.1, cached_entity.1, "text mismatch for input: {text}"); + assert_eq!(uncached_entity.2, cached_entity.2, "start mismatch for input: {text}"); + assert_eq!(uncached_entity.3, cached_entity.3, "end mismatch for input: {text}"); + assert!( + (uncached_entity.4 - cached_entity.4).abs() < 1e-5, + "score mismatch for input: {text}: uncached={:.6} cached={:.6}", + uncached_entity.4, + cached_entity.4 + ); + } + } + + Ok(()) + }) +} + +#[test] +fn edge_extract_batch_matches_serial_extract() -> anyhow::Result<()> { + tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build()? + .block_on(async { + let mut gliner = Gliner::builder() + .with_source(remote_edge_source()) + .with_device(Device::cpu()) + .build() + .await?; + + let labels = ["person", "organization", "location"]; + let texts = [ + "Apple Inc. was founded by Steve Jobs in California.", + "Microsoft Corporation is headquartered in Seattle.", + "", + "Google was founded in Mountain View.", + ]; + + let mut serial = Vec::with_capacity(texts.len()); + for text in texts.iter().copied() { + let entities = gliner.extract(text, &labels).await?; + serial.push( + entities + .into_iter() + .map(|entity| { + ( + entity.label, + entity.text, + entity.start_char, + entity.end_char, + entity.score, + ) + }) + .collect::>(), ); } + let batched = gliner.extract_batch(&texts, &labels).await?; + assert_eq!(batched.len(), texts.len()); + + for (serial_entities, batched_entities) in serial.iter().zip(&batched) { + let batched: Vec<(String, String, usize, usize, f32)> = batched_entities + .iter() + .map(|entity| { + ( + entity.label.clone(), + entity.text.clone(), + entity.start_char, + entity.end_char, + entity.score, + ) + }) + .collect(); + + assert_eq!(serial_entities.len(), batched.len()); + for (serial_entity, batched_entity) in serial_entities.iter().zip(&batched) { + assert_eq!(serial_entity.0, batched_entity.0); + assert_eq!(serial_entity.1, batched_entity.1); + assert_eq!(serial_entity.2, batched_entity.2); + assert_eq!(serial_entity.3, batched_entity.3); + assert!( + (serial_entity.4 - batched_entity.4).abs() < 1e-5, + "score mismatch: serial={:.6} batched={:.6}", + serial_entity.4, + batched_entity.4 + ); + } + } + + Ok(()) + }) +} + +#[test] +#[ignore = "cache the remote edge checkpoint and sidecars"] +fn cache_remote_edge_checkpoint() -> anyhow::Result<()> { + tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build()? + .block_on(async { + let _gliner = Gliner::builder() + .with_source(remote_edge_source()) + .with_device(Device::cpu()) + .build() + .await?; Ok(()) }) } From c956136422d290cce5616738b0727141cd5cabf2 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Tue, 14 Apr 2026 21:15:15 -0500 Subject: [PATCH 17/34] closer ui --- Cargo.lock | 233 +----------------- demos/rgliner-web/.claude/settings.local.json | 7 + demos/rgliner-web/Cargo.toml | 2 +- demos/rgliner-web/src/main.rs | 180 +++++++------- models/rgliner/src/relex.rs | 87 +++++++ 5 files changed, 188 insertions(+), 321 deletions(-) create mode 100644 demos/rgliner-web/.claude/settings.local.json diff --git a/Cargo.lock b/Cargo.lock index f2c6180e0..43b256b1d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -587,7 +587,7 @@ dependencies = [ "derive_builder 0.20.2", "diligent-date-parser", "never", - "quick-xml 0.37.5", + "quick-xml", ] [[package]] @@ -1001,31 +1001,6 @@ dependencies = [ "cipher", ] -[[package]] -name = "bon" -version = "3.9.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f47dbe92550676ee653353c310dfb9cf6ba17ee70396e1f7cf0a2020ad49b2fe" -dependencies = [ - "bon-macros", - "rustversion", -] - -[[package]] -name = "bon-macros" -version = "3.9.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "519bd3116aeeb42d5372c29d982d16d0170d3d4a5ed85fc7dd91642ffff3c67c" -dependencies = [ - "darling 0.20.11", - "ident_case", - "prettyplease", - "proc-macro2", - "quote", - "rustversion", - "syn 2.0.117", -] - [[package]] name = "borsh" version = "1.6.0" @@ -2973,15 +2948,6 @@ dependencies = [ "tracing-wasm", ] -[[package]] -name = "dioxus-markdown" -version = "0.1.0" -source = "git+https://github.com/rambip/rust-web-markdown#22ab22566014a8bd5bac959dd2a4770d2eddf16b" -dependencies = [ - "dioxus", - "web-framework-markdown", -] - [[package]] name = "dioxus-primitives" version = "0.0.1" @@ -3564,17 +3530,6 @@ dependencies = [ "regex-syntax 0.8.10", ] -[[package]] -name = "fancy-regex" -version = "0.16.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "998b056554fbe42e03ae0e152895cd1a7e1002aec800fdc6635d20270260c46f" -dependencies = [ - "bit-set 0.8.0", - "regex-automata 0.4.14", - "regex-syntax 0.8.10", -] - [[package]] name = "fancy-regex" version = "0.17.0" @@ -5985,24 +5940,6 @@ dependencies = [ name = "kalosm-workspace" version = "0.4.0" -[[package]] -name = "katex-rs" -version = "0.2.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c5382fea1e8edf972c23050cdfec2d12beca242e6d15e32310e5d1543a51d103" -dependencies = [ - "bon", - "phf 0.13.1", - "phf_codegen 0.13.1", - "rapidhash", - "serde", - "serde_json", - "strum 0.28.0", - "strum_macros 0.28.0", - "thiserror 2.0.18", - "unicode-normalization", -] - [[package]] name = "keyboard-types" version = "0.7.0" @@ -6176,12 +6113,6 @@ dependencies = [ "thiserror 1.0.69", ] -[[package]] -name = "linked-hash-map" -version = "0.5.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0717cef1bc8b636c6e1c1bbdefc09e6322da8a9321966e8928ef80d20f7f770f" - [[package]] name = "linux-raw-sys" version = "0.12.1" @@ -7587,21 +7518,10 @@ version = "0.11.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1fd6780a80ae0c52cc120a26a1a42c1ae51b247a253e4e06113d23d2c2edd078" dependencies = [ - "phf_macros 0.11.3", + "phf_macros", "phf_shared 0.11.3", ] -[[package]] -name = "phf" -version = "0.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c1562dc717473dbaa4c1f85a36410e03c047b2e7df7f45ee938fbef64ae7fadf" -dependencies = [ - "phf_macros 0.13.1", - "phf_shared 0.13.1", - "serde", -] - [[package]] name = "phf_codegen" version = "0.10.0" @@ -7622,16 +7542,6 @@ dependencies = [ "phf_shared 0.11.3", ] -[[package]] -name = "phf_codegen" -version = "0.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "49aa7f9d80421bca176ca8dbfebe668cc7a2684708594ec9f3c0db0805d5d6e1" -dependencies = [ - "phf_generator 0.13.1", - "phf_shared 0.13.1", -] - [[package]] name = "phf_generator" version = "0.10.0" @@ -7652,16 +7562,6 @@ dependencies = [ "rand 0.8.5", ] -[[package]] -name = "phf_generator" -version = "0.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "135ace3a761e564ec88c03a77317a7c6b80bb7f7135ef2544dbe054243b89737" -dependencies = [ - "fastrand", - "phf_shared 0.13.1", -] - [[package]] name = "phf_macros" version = "0.11.3" @@ -7676,19 +7576,6 @@ dependencies = [ "unicase", ] -[[package]] -name = "phf_macros" -version = "0.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "812f032b54b1e759ccd5f8b6677695d5268c588701effba24601f6932f8269ef" -dependencies = [ - "phf_generator 0.13.1", - "phf_shared 0.13.1", - "proc-macro2", - "quote", - "syn 2.0.117", -] - [[package]] name = "phf_shared" version = "0.10.0" @@ -7708,15 +7595,6 @@ dependencies = [ "unicase", ] -[[package]] -name = "phf_shared" -version = "0.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e57fef6bc5981e38c2ce2d63bfa546861309f875b8a75f092d1d54ae2d64f266" -dependencies = [ - "siphasher 1.0.2", -] - [[package]] name = "pico-args" version = "0.5.0" @@ -7767,19 +7645,6 @@ version = "0.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b4596b6d070b27117e987119b4dac604f3c58cfb0b191112e24771b2faeac1a6" -[[package]] -name = "plist" -version = "1.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "740ebea15c5d1428f910cd1a5f52cebf8d25006245ed8ade92702f4943d91e07" -dependencies = [ - "base64 0.22.1", - "indexmap 2.13.0", - "quick-xml 0.38.4", - "serde", - "time", -] - [[package]] name = "plotters" version = "0.3.7" @@ -8083,18 +7948,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c3a14896dfa883796f1cb410461aef38810ea05f2b2c33c5aded3649095fdad" dependencies = [ "bitflags 2.11.0", - "getopts", "memchr", - "pulldown-cmark-escape", "unicase", ] -[[package]] -name = "pulldown-cmark-escape" -version = "0.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "007d8adb5ddab6f8e3f491ac63566a7d5002cc7ed73901f72057943fa71ae1ae" - [[package]] name = "pulp" version = "0.18.22" @@ -8175,15 +8032,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "quick-xml" -version = "0.38.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b66c2058c55a409d601666cffe35f04333cf1013010882cec174a7467cd4e21c" -dependencies = [ - "memchr", -] - [[package]] name = "quick_cache" version = "0.5.2" @@ -8382,15 +8230,6 @@ version = "1.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "973443cf09a9c8656b574a866ab68dfa19f0867d0340648c7d2f6a71b8a8ea68" -[[package]] -name = "rapidhash" -version = "4.4.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b5e48930979c155e2f33aa36ab3119b5ee81332beb6482199a8ecd6029b80b59" -dependencies = [ - "rustversion", -] - [[package]] name = "rav1e" version = "0.8.1" @@ -8876,10 +8715,10 @@ version = "0.1.0" dependencies = [ "console_error_panic_hook", "dioxus", - "dioxus-markdown", "dioxus-primitives", "getrandom 0.3.4", "gloo-timers", + "pulldown-cmark 0.13.3", "rgliner", "tracing", "tracing-wasm", @@ -8996,7 +8835,7 @@ dependencies = [ "atom_syndication", "derive_builder 0.20.2", "never", - "quick-xml 0.37.5", + "quick-xml", ] [[package]] @@ -10009,15 +9848,6 @@ dependencies = [ "strum_macros 0.27.2", ] -[[package]] -name = "strum" -version = "0.28.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9628de9b8791db39ceda2b119bbe13134770b56c138ec1d3af810d045c04f9bd" -dependencies = [ - "strum_macros 0.28.0", -] - [[package]] name = "strum_macros" version = "0.26.4" @@ -10043,18 +9873,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "strum_macros" -version = "0.28.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ab85eea0270ee17587ed4156089e10b9e6880ee688791d45a905f5b1ca36f664" -dependencies = [ - "heck", - "proc-macro2", - "quote", - "syn 2.0.117", -] - [[package]] name = "subsecond" version = "0.7.4" @@ -10348,27 +10166,6 @@ dependencies = [ "syn 2.0.117", ] -[[package]] -name = "syntect" -version = "5.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "656b45c05d95a5704399aeef6bd0ddec7b2b3531b7c9e900abbf7c4d2190c925" -dependencies = [ - "bincode", - "fancy-regex 0.16.2", - "flate2", - "fnv", - "once_cell", - "plist", - "regex-syntax 0.8.10", - "serde", - "serde_derive", - "serde_json", - "thiserror 2.0.18", - "walkdir", - "yaml-rust", -] - [[package]] name = "sysctl" version = "0.5.5" @@ -11688,19 +11485,6 @@ dependencies = [ "pkg-config", ] -[[package]] -name = "web-framework-markdown" -version = "0.1.0" -source = "git+https://github.com/rambip/rust-web-markdown#22ab22566014a8bd5bac959dd2a4770d2eddf16b" -dependencies = [ - "katex-rs", - "lazy_static", - "pulldown-cmark 0.13.3", - "regex", - "syntect", - "web-sys", -] - [[package]] name = "web-sys" version = "0.3.91" @@ -12646,15 +12430,6 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7a5a4b21e1a62b67a2970e6831bc091d7b87e119e7f9791aef9702e3bef04448" -[[package]] -name = "yaml-rust" -version = "0.4.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "56c1936c4cc7a1c9ab21a1ebb602eb942ba868cbd44a99cb7cdc5892335e1c85" -dependencies = [ - "linked-hash-map", -] - [[package]] name = "yansi" version = "1.0.1" diff --git a/demos/rgliner-web/.claude/settings.local.json b/demos/rgliner-web/.claude/settings.local.json new file mode 100644 index 000000000..1078b8c08 --- /dev/null +++ b/demos/rgliner-web/.claude/settings.local.json @@ -0,0 +1,7 @@ +{ + "permissions": { + "allow": [ + "Bash(cargo tree:*)" + ] + } +} diff --git a/demos/rgliner-web/Cargo.toml b/demos/rgliner-web/Cargo.toml index f8334b9f1..f67fafcc8 100644 --- a/demos/rgliner-web/Cargo.toml +++ b/demos/rgliner-web/Cargo.toml @@ -10,7 +10,7 @@ rgliner = { path = "../../models/rgliner", default-features = false } getrandom = { version = "0.3", features = ["wasm_js"] } tracing = "0.1" dioxus-primitives = { git = "https://github.com/DioxusLabs/components", version = "0.0.1", default-features = false } -dioxus-markdown = { git = "https://github.com/rambip/rust-web-markdown" } +pulldown-cmark = { version = "0.13", default-features = false } [target.'cfg(target_arch = "wasm32")'.dependencies] console_error_panic_hook = "0.1" diff --git a/demos/rgliner-web/src/main.rs b/demos/rgliner-web/src/main.rs index fd1a57c2f..6f8f43a29 100644 --- a/demos/rgliner-web/src/main.rs +++ b/demos/rgliner-web/src/main.rs @@ -6,7 +6,9 @@ use components::select::{ Select, SelectItemIndicator, SelectList, SelectOption, SelectTrigger, SelectValue, }; use dioxus::prelude::*; -use dioxus_markdown::Markdown; +use pulldown_cmark::{Event, HeadingLevel, Options, Parser, Tag, TagEnd}; +use std::iter::Peekable; +use std::ops::Range; use rgliner::{ relation_decoding::Relation, relex::{GlinerRelEx, GlinerRelExSource}, @@ -119,36 +121,31 @@ fn App() -> Element { let mut error = use_signal(|| None::); let mut extraction = use_signal(Extraction::default); let mut status = use_signal(|| "idle".to_string()); - let mut schedule = use_signal(|| 0u64); - - // Bump the schedule whenever any input that affects extraction changes. - // Do NOT read `model` here — the extractor writes it back on every run, - // which would re-trigger this effect in an endless loop. - use_effect(move || { - let _ = text(); - let _ = entity_labels(); - let _ = relation_labels(); - schedule.with_mut(|s| *s += 1); - }); - // React to schedule changes: debounce, then extract. - use_effect(move || { - let current = schedule(); - if current == 0 { + // Memoised so the effect below only re-runs when a model appears or + // disappears, not every time `run_extraction` writes the model back. + let model_ready = use_memo(move || model.read().is_some()); + + // One extraction pipeline: watch the inputs, debounce, run. + // `use_future` cancels and restarts whenever any tracked signal changes, + // which gives us the debounce for free. + use_future(move || async move { + let cur_text = text(); + let ent_raw = entity_labels(); + let rel_raw = relation_labels(); + if !model_ready() { return; } + + // Cancelled if any of the above change during the wait. + sleep_ms(DEBOUNCE_MS).await; + + // Detach the extraction so cancellation can't drop it mid-run + // (which would lose the model we've taken out of the signal). spawn(async move { - sleep_ms(DEBOUNCE_MS).await; - if schedule() != current { + if running() { return; } - // Wait for any in-flight extraction to finish; bail if a newer change arrives. - while running() { - sleep_ms(80).await; - if schedule() != current { - return; - } - } let Some(mut taken) = model.write().take() else { return; }; @@ -156,12 +153,10 @@ fn App() -> Element { error.set(None); status.set("extracting…".to_string()); - let ent_labels = parse_labels(&entity_labels()); - let rel_labels = parse_labels(&relation_labels()); - let cur_text = text(); + let ent = parse_labels(&ent_raw); + let rel = parse_labels(&rel_raw); let mode = taken.choice().mode(); - let outcome = - run_extraction(&mut taken, mode, &cur_text, &ent_labels, &rel_labels).await; + let outcome = run_extraction(&mut taken, mode, &cur_text, &ent, &rel).await; model.set(Some(taken)); match outcome { @@ -206,8 +201,6 @@ fn App() -> Element { Ok(m) => { model.set(Some(m)); status.set("ready".to_string()); - // Kick an extraction now that a model is available. - schedule.with_mut(|s| *s += 1); } Err(e) => { error.set(Some(format!("{e}"))); @@ -411,23 +404,34 @@ async fn run_extraction( } } -fn render_article(text: &str, ex: &Extraction) -> Element { - let spliced = splice_entities(text, &ex.entities, &ex.relations); - rsx! { - Markdown { src: spliced } - } +#[derive(Clone, PartialEq)] +struct EntityView { + text: String, + label: String, + color: String, + rels: Vec<(String, String, String)>, +} + +#[derive(Clone, PartialEq, Default)] +struct Article { + source: String, + entities: Vec, } -/// Splice `` into the markdown source at each -/// entity boundary. The nested `.rels` span renders on hover. -fn splice_entities(text: &str, entities: &[Entity], relations: &[Relation]) -> String { - if entities.is_empty() { - return text.to_string(); +/// Splice `` markers into the markdown source at each entity +/// boundary, and collect the per-entity view the custom component renders. +fn build_article(text: &str, ex: &Extraction) -> Article { + if ex.entities.is_empty() { + return Article { + source: text.to_string(), + entities: Vec::new(), + }; } - let mut sorted: Vec<&Entity> = entities.iter().collect(); + let mut sorted: Vec<&Entity> = ex.entities.iter().collect(); sorted.sort_by_key(|e| e.start_char); - let mut out = String::with_capacity(text.len() + entities.len() * 64); + let mut source = String::with_capacity(text.len() + sorted.len() * 20); + let mut entities: Vec = Vec::new(); let mut cursor = 0usize; let len = text.len(); @@ -440,61 +444,55 @@ fn splice_entities(text: &str, entities: &[Entity], relations: &[Relation]) -> S if !text.is_char_boundary(start) || !text.is_char_boundary(end) { continue; } - out.push_str(&text[cursor..start]); - let color = underline_color(&ent.label); - out.push_str(&format!( - r#""#, - color = color, - label = escape_attr(&ent.label) - )); - // The entity surface text. - out.push_str(&escape_html_content(&text[start..end])); - // Popover: label + relations involving this entity. - out.push_str(r#""#); - out.push_str(r#""#); - out.push_str(&escape_html_content(&ent.label)); - out.push_str(""); + source.push_str(&text[cursor..start]); + let idx = entities.len(); + source.push_str(&format!("")); let ent_text = &text[start..end]; - let mut rel_lines = 0usize; - for rel in relations { - if rel.head.text == ent_text || rel.tail.text == ent_text { - out.push_str(r#""#); - out.push_str(&escape_html_content(&rel.head.text)); - out.push_str(r#""#); - out.push_str(r#""#); - out.push_str(&escape_html_content(&rel.relation)); - out.push_str(""); - out.push_str(r#""#); - out.push_str(&escape_html_content(&rel.tail.text)); - out.push_str(""); - rel_lines += 1; - } - } - if rel_lines == 0 && !relations.is_empty() { - out.push_str(r#"no relations"#); - } - out.push_str(""); - out.push_str(""); + let rels: Vec<(String, String, String)> = ex + .relations + .iter() + .filter(|r| r.head.text == ent_text || r.tail.text == ent_text) + .map(|r| (r.head.text.clone(), r.relation.clone(), r.tail.text.clone())) + .collect(); + + entities.push(EntityView { + text: ent_text.to_string(), + label: ent.label.clone(), + color: underline_color(&ent.label), + rels, + }); cursor = end; } if cursor < len { - out.push_str(&text[cursor..]); + source.push_str(&text[cursor..]); } - out -} - -fn escape_html_content(s: &str) -> String { - s.replace('&', "&") - .replace('<', "<") - .replace('>', ">") + Article { source, entities } } -fn escape_attr(s: &str) -> String { - s.replace('&', "&") - .replace('"', """) - .replace('<', "<") - .replace('>', ">") +fn render_entity(view: EntityView) -> Element { + let has_rels = !view.rels.is_empty(); + rsx! { + span { + class: "entity", + style: "--ec: {view.color};", + "{view.text}" + span { class: "pop", + span { class: "pop-label", "{view.label}" } + if has_rels { + for (i, (head, rel, tail)) in view.rels.iter().cloned().enumerate() { + span { key: "r-{i}", class: "pop-rel", + "{head}" + span { class: "arrow", " → " } + span { class: "rel-name", "{rel}" } + span { class: "arrow", " → " } + "{tail}" + } + } + } + } + } + } } fn underline_color(label: &str) -> String { diff --git a/models/rgliner/src/relex.rs b/models/rgliner/src/relex.rs index f540a16e2..7168dfa5e 100644 --- a/models/rgliner/src/relex.rs +++ b/models/rgliner/src/relex.rs @@ -1578,6 +1578,93 @@ Meta Platforms was founded by Mark Zuckerberg in Cambridge, Massachusetts."; Ok(()) } + fn speedup(cpu: Duration, gpu: Duration) -> f64 { + cpu.as_secs_f64() / gpu.as_secs_f64() + } + + fn print_profile_comparison(variant: &str, cpu: &ExtractProfile, gpu: &ExtractProfile) { + println!( + "COMPARE variant={variant} warm_total_ms cpu={:.2} gpu={:.2} speedup={:.2}x", + cpu.warm_total.as_secs_f64() * 1000.0, + gpu.warm_total.as_secs_f64() * 1000.0, + speedup(cpu.warm_total, gpu.warm_total), + ); + println!( + " cold_total_ms cpu={:.2} gpu={:.2} speedup={:.2}x", + cpu.cold_total.as_secs_f64() * 1000.0, + gpu.cold_total.as_secs_f64() * 1000.0, + speedup(cpu.cold_total, gpu.cold_total), + ); + println!( + " entity_sync_ms cpu={:.2} gpu={:.2} speedup={:.2}x", + cpu.entity_sync.as_secs_f64() * 1000.0, + gpu.entity_sync.as_secs_f64() * 1000.0, + speedup(cpu.entity_sync, gpu.entity_sync), + ); + println!( + " relation_span_sync_ms cpu={:.2} gpu={:.2} speedup={:.2}x", + cpu.relation_span_sync.as_secs_f64() * 1000.0, + gpu.relation_span_sync.as_secs_f64() * 1000.0, + speedup(cpu.relation_span_sync, gpu.relation_span_sync), + ); + println!( + " relation_score_sync_ms cpu={:.2} gpu={:.2} speedup={:.2}x", + cpu.relation_score_sync.as_secs_f64() * 1000.0, + gpu.relation_score_sync.as_secs_f64() * 1000.0, + speedup(cpu.relation_score_sync, gpu.relation_score_sync), + ); + println!( + " tokenize_cpu_ms cpu={:.2} gpu={:.2}", + cpu.tokenize_cpu.as_secs_f64() * 1000.0, + gpu.tokenize_cpu.as_secs_f64() * 1000.0, + ); + println!( + " relation_pair_pack_cpu_ms cpu={:.2} gpu={:.2}", + cpu.relation_pair_pack_cpu.as_secs_f64() * 1000.0, + gpu.relation_pair_pack_cpu.as_secs_f64() * 1000.0, + ); + println!( + " relation_decode_cpu_ms cpu={:.2} gpu={:.2}", + cpu.relation_decode_cpu.as_secs_f64() * 1000.0, + gpu.relation_decode_cpu.as_secs_f64() * 1000.0, + ); + } + + #[tokio::test] + #[ignore = "compare cpu vs gpu rel-ex throughput on remote checkpoints"] + async fn compare_relex_cpu_vs_gpu() -> Result<(), Box> { + let gpu_device = Device::gpu().await?; + let cpu_device = Device::cpu(); + + let variants = [ + ("multi", GlinerRelExSource::relex_multi as fn() -> GlinerRelExSource), + ("base", GlinerRelExSource::relex_base as fn() -> GlinerRelExSource), + ("large", GlinerRelExSource::relex_large as fn() -> GlinerRelExSource), + ]; + + for (variant, source) in variants { + let cpu_model = load_relex(source(), cpu_device.clone()).await?; + let cpu_cold_start = Instant::now(); + let _ = cpu_model + .extract(PROFILE_TEXT, ENTITY_LABELS, RELATION_LABELS) + .await?; + let cpu_cold_total = cpu_cold_start.elapsed(); + let cpu_profile = profile_extract(&cpu_model, variant, cpu_cold_total).await?; + + let gpu_model = load_relex(source(), gpu_device.clone()).await?; + let gpu_cold_start = Instant::now(); + let _ = gpu_model + .extract(PROFILE_TEXT, ENTITY_LABELS, RELATION_LABELS) + .await?; + let gpu_cold_total = gpu_cold_start.elapsed(); + let gpu_profile = profile_extract(&gpu_model, variant, gpu_cold_total).await?; + + print_profile_comparison(variant, &cpu_profile, &gpu_profile); + } + + Ok(()) + } + fn entity_signature(entities: &[Entity]) -> Vec<(String, String, usize, usize, usize, usize)> { entities .iter() From c68e3f344db0b9651ab63aadc351530a6e80f263 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Tue, 14 Apr 2026 21:22:25 -0500 Subject: [PATCH 18/34] optimize bilstm --- demos/rgliner-web/src/main.rs | 237 +++++++++++++++++++++++++------ models/rgliner/src/raw/bilstm.rs | 56 ++++++-- 2 files changed, 241 insertions(+), 52 deletions(-) diff --git a/demos/rgliner-web/src/main.rs b/demos/rgliner-web/src/main.rs index 6f8f43a29..96c1aea13 100644 --- a/demos/rgliner-web/src/main.rs +++ b/demos/rgliner-web/src/main.rs @@ -6,7 +6,7 @@ use components::select::{ Select, SelectItemIndicator, SelectList, SelectOption, SelectTrigger, SelectValue, }; use dioxus::prelude::*; -use pulldown_cmark::{Event, HeadingLevel, Options, Parser, Tag, TagEnd}; +use pulldown_cmark::{Event, HeadingLevel, Options, Parser, Tag}; use std::iter::Peekable; use std::ops::Range; use rgliner::{ @@ -127,26 +127,32 @@ fn App() -> Element { let model_ready = use_memo(move || model.read().is_some()); // One extraction pipeline: watch the inputs, debounce, run. - // `use_future` cancels and restarts whenever any tracked signal changes, - // which gives us the debounce for free. - use_future(move || async move { + // `use_resource` re-runs the async body whenever any tracked signal + // changes, cancelling the previous invocation — the debounce falls + // out of that cancellation behavior. + use_resource(move || async move { let cur_text = text(); let ent_raw = entity_labels(); let rel_raw = relation_labels(); - if !model_ready() { + let ready = model_ready(); + tracing::info!("pipeline tick: ready={ready} text_len={}", cur_text.len()); + if !ready { return; } // Cancelled if any of the above change during the wait. sleep_ms(DEBOUNCE_MS).await; + tracing::info!("pipeline: debounce survived, kicking extraction"); // Detach the extraction so cancellation can't drop it mid-run // (which would lose the model we've taken out of the signal). spawn(async move { if running() { + tracing::info!("extraction already running, skipping"); return; } let Some(mut taken) = model.write().take() else { + tracing::warn!("no model in slot at extract-time"); return; }; running.set(true); @@ -156,11 +162,21 @@ fn App() -> Element { let ent = parse_labels(&ent_raw); let rel = parse_labels(&rel_raw); let mode = taken.choice().mode(); + let mode_name = match mode { + Mode::Ner => "ner", + Mode::Relex => "relex", + }; + tracing::info!("extracting: mode={mode_name} ent={} rel={}", ent.len(), rel.len()); let outcome = run_extraction(&mut taken, mode, &cur_text, &ent, &rel).await; model.set(Some(taken)); match outcome { Ok(e) => { + tracing::info!( + "extraction done: {} entities, {} relations", + e.entities.len(), + e.relations.len() + ); status.set(format!( "{} entities · {} relations", e.entities.len(), @@ -169,6 +185,7 @@ fn App() -> Element { extraction.set(e); } Err(e) => { + tracing::warn!("extraction error: {e}"); error.set(Some(e)); status.set("error".to_string()); } @@ -405,33 +422,19 @@ async fn run_extraction( } #[derive(Clone, PartialEq)] -struct EntityView { - text: String, +struct EntitySpan { + byte_start: usize, + byte_end: usize, label: String, color: String, rels: Vec<(String, String, String)>, } -#[derive(Clone, PartialEq, Default)] -struct Article { - source: String, - entities: Vec, -} - -/// Splice `` markers into the markdown source at each entity -/// boundary, and collect the per-entity view the custom component renders. -fn build_article(text: &str, ex: &Extraction) -> Article { - if ex.entities.is_empty() { - return Article { - source: text.to_string(), - entities: Vec::new(), - }; - } +fn collect_entity_spans(text: &str, ex: &Extraction) -> Vec { let mut sorted: Vec<&Entity> = ex.entities.iter().collect(); sorted.sort_by_key(|e| e.start_char); - let mut source = String::with_capacity(text.len() + sorted.len() * 20); - let mut entities: Vec = Vec::new(); + let mut spans = Vec::new(); let mut cursor = 0usize; let len = text.len(); @@ -444,10 +447,6 @@ fn build_article(text: &str, ex: &Extraction) -> Article { if !text.is_char_boundary(start) || !text.is_char_boundary(end) { continue; } - source.push_str(&text[cursor..start]); - let idx = entities.len(); - source.push_str(&format!("")); - let ent_text = &text[start..end]; let rels: Vec<(String, String, String)> = ex .relations @@ -455,32 +454,190 @@ fn build_article(text: &str, ex: &Extraction) -> Article { .filter(|r| r.head.text == ent_text || r.tail.text == ent_text) .map(|r| (r.head.text.clone(), r.relation.clone(), r.tail.text.clone())) .collect(); - - entities.push(EntityView { - text: ent_text.to_string(), + spans.push(EntitySpan { + byte_start: start, + byte_end: end, label: ent.label.clone(), color: underline_color(&ent.label), rels, }); cursor = end; } - if cursor < len { - source.push_str(&text[cursor..]); + spans +} + +/// Render the markdown article, wrapping entity byte-ranges in interactive +/// spans at the text-event level. We walk pulldown-cmark's flat event stream +/// and build rsx! directly — no HTML serialisation, no custom-component dance. +fn render_article(text: &str, ex: &Extraction) -> Element { + let spans = collect_entity_spans(text, ex); + let mut opts = Options::empty(); + opts.insert(Options::ENABLE_STRIKETHROUGH); + opts.insert(Options::ENABLE_TABLES); + let parser = Parser::new_ext(text, opts).into_offset_iter(); + let mut r = Renderer { + source: text, + spans: &spans, + events: parser.peekable(), + }; + let nodes = r.render_until_end(false); + rsx! { {nodes.into_iter()} } +} + +struct Renderer<'a, I: Iterator, Range)>> { + source: &'a str, + spans: &'a [EntitySpan], + events: Peekable, +} + +impl<'a, I: Iterator, Range)>> Renderer<'a, I> { + /// Consume events until `End(_)` (if `scoped`) or EOF, returning the + /// rsx! elements they expand to. + fn render_until_end(&mut self, scoped: bool) -> Vec { + let mut nodes: Vec = Vec::new(); + while let Some((event, range)) = self.events.next() { + match event { + Event::Start(tag) => nodes.push(self.render_tag(tag)), + Event::End(_) if scoped => return nodes, + Event::End(_) => continue, + Event::Text(_) => nodes.push(self.render_text_range(range)), + Event::Code(s) => { + let s = s.to_string(); + nodes.push(rsx! { code { "{s}" } }); + } + Event::SoftBreak => nodes.push(rsx! { " " }), + Event::HardBreak => nodes.push(rsx! { br {} }), + Event::Rule => nodes.push(rsx! { hr {} }), + Event::Html(s) | Event::InlineHtml(s) => { + // Treat raw HTML as literal text — safer than dangerous_inner_html. + let s = s.to_string(); + nodes.push(rsx! { "{s}" }); + } + Event::FootnoteReference(_) => {} + Event::TaskListMarker(done) => nodes.push(rsx! { + input { r#type: "checkbox", checked: done, disabled: true } + }), + _ => {} + } + } + nodes + } + + fn render_tag(&mut self, tag: Tag<'a>) -> Element { + match tag { + Tag::Paragraph => { + let children = self.render_until_end(true); + rsx! { p { {children.into_iter()} } } + } + Tag::Heading { level, .. } => { + let children = self.render_until_end(true); + match level { + HeadingLevel::H1 => rsx! { h1 { {children.into_iter()} } }, + HeadingLevel::H2 => rsx! { h2 { {children.into_iter()} } }, + HeadingLevel::H3 => rsx! { h3 { {children.into_iter()} } }, + HeadingLevel::H4 => rsx! { h4 { {children.into_iter()} } }, + HeadingLevel::H5 => rsx! { h5 { {children.into_iter()} } }, + HeadingLevel::H6 => rsx! { h6 { {children.into_iter()} } }, + } + } + Tag::BlockQuote(_) => { + let children = self.render_until_end(true); + rsx! { blockquote { {children.into_iter()} } } + } + Tag::CodeBlock(_) => { + let children = self.render_until_end(true); + rsx! { pre { code { {children.into_iter()} } } } + } + Tag::List(Some(_start)) => { + let children = self.render_until_end(true); + rsx! { ol { {children.into_iter()} } } + } + Tag::List(None) => { + let children = self.render_until_end(true); + rsx! { ul { {children.into_iter()} } } + } + Tag::Item => { + let children = self.render_until_end(true); + rsx! { li { {children.into_iter()} } } + } + Tag::Emphasis => { + let children = self.render_until_end(true); + rsx! { em { {children.into_iter()} } } + } + Tag::Strong => { + let children = self.render_until_end(true); + rsx! { strong { {children.into_iter()} } } + } + Tag::Strikethrough => { + let children = self.render_until_end(true); + rsx! { s { {children.into_iter()} } } + } + Tag::Link { dest_url, title, .. } => { + let children = self.render_until_end(true); + let href = dest_url.to_string(); + let title = title.to_string(); + rsx! { a { href: "{href}", title: "{title}", {children.into_iter()} } } + } + Tag::Image { dest_url, title, .. } => { + // Consume inner events (alt text) without rendering them separately. + let _ = self.render_until_end(true); + let src = dest_url.to_string(); + let title = title.to_string(); + rsx! { img { src: "{src}", title: "{title}" } } + } + _ => { + let children = self.render_until_end(true); + rsx! { span { {children.into_iter()} } } + } + } + } + + /// Render the text at `range`, splitting on any entity spans that overlap it. + fn render_text_range(&self, range: Range) -> Element { + let mut parts: Vec = Vec::new(); + let mut cursor = range.start; + let slice_end = range.end; + + for span in self.spans.iter() { + if span.byte_end <= cursor { + continue; + } + if span.byte_start >= slice_end { + break; + } + let s = span.byte_start.max(cursor); + let e = span.byte_end.min(slice_end); + if s > cursor { + let plain = self.source[cursor..s].to_string(); + parts.push(rsx! { "{plain}" }); + } + let ent_text = self.source[s..e].to_string(); + parts.push(render_entity(&ent_text, span)); + cursor = e; + } + if cursor < slice_end { + let tail = self.source[cursor..slice_end].to_string(); + parts.push(rsx! { "{tail}" }); + } + rsx! { {parts.into_iter()} } } - Article { source, entities } } -fn render_entity(view: EntityView) -> Element { - let has_rels = !view.rels.is_empty(); +fn render_entity(text: &str, span: &EntitySpan) -> Element { + let text = text.to_string(); + let label = span.label.clone(); + let color = span.color.clone(); + let rels = span.rels.clone(); + let has_rels = !rels.is_empty(); rsx! { span { class: "entity", - style: "--ec: {view.color};", - "{view.text}" + style: "--ec: {color};", + "{text}" span { class: "pop", - span { class: "pop-label", "{view.label}" } + span { class: "pop-label", "{label}" } if has_rels { - for (i, (head, rel, tail)) in view.rels.iter().cloned().enumerate() { + for (i, (head, rel, tail)) in rels.into_iter().enumerate() { span { key: "r-{i}", class: "pop-rel", "{head}" span { class: "arrow", " → " } diff --git a/models/rgliner/src/raw/bilstm.rs b/models/rgliner/src/raw/bilstm.rs index 3530f988f..15c768719 100644 --- a/models/rgliner/src/raw/bilstm.rs +++ b/models/rgliner/src/raw/bilstm.rs @@ -244,17 +244,42 @@ fn run_direction( reverse: bool, lengths: &[usize], ) -> Tensor<3, f32> { - let [batch, seq_len, _] = input.shape(); + let [batch, seq_len, input_size] = input.shape(); let mut h: Tensor<2, f32> = Tensor::zeros(device, [batch, hidden_size]); let mut c: Tensor<2, f32> = Tensor::zeros(device, [batch, hidden_size]); let mut outputs: Tensor<3, f32> = Tensor::zeros(device, [batch, seq_len, hidden_size]); + let all_active = lengths.iter().all(|&length| length >= seq_len); let bias_broadcast: Tensor<2, f32> = dir .bias .unsqueeze(0) .broadcast_as([batch, 4 * hidden_size]) .to_concrete(); + let input_gates: Tensor<3, f32> = input + .reshape([batch * seq_len, input_size]) + .to_concrete() + .mat_mul(&dir.w_ih_t) + .reshape([batch, seq_len, 4 * hidden_size]) + .to_concrete(); + let active_masks = if all_active { + None + } else { + Some( + Tensor::new( + device, + &lengths + .iter() + .flat_map(|&length| { + (0..seq_len) + .flat_map(move |t| std::iter::repeat_n(if t < length { 1.0 } else { 0.0 }, hidden_size)) + }) + .collect::>(), + ) + .reshape([batch, seq_len, hidden_size]) + .to_concrete(), + ) + }; let iter: Box> = if reverse { Box::new((0..seq_len).rev()) @@ -263,14 +288,14 @@ fn run_direction( }; for t in iter { - let x_t: Tensor<2, f32> = input + let x_gates_t: Tensor<2, f32> = input_gates .narrow(1, t, 1) - .reshape([batch, input.shape()[2]]) + .reshape([batch, 4 * hidden_size]) .to_concrete(); // gates_pre = x_t @ W_ih^T + h @ W_hh^T + bias, shape [batch, 4*hidden] let gates_pre: Tensor<2, f32> = - (x_t.mat_mul(&dir.w_ih_t) + h.mat_mul(&dir.w_hh_t) + bias_broadcast.clone()) + (x_gates_t + h.mat_mul(&dir.w_hh_t) + bias_broadcast.clone()) .to_concrete(); let i_raw: Tensor<2, f32> = gates_pre.narrow(1, 0, hidden_size).to_concrete(); @@ -290,17 +315,24 @@ fn run_direction( let next_c = (f_gate * c.clone() + i_gate * g_gate).to_concrete(); let next_h = (o_gate * next_c.clone().tanh()).to_concrete(); - let active_mask_2d = timestep_mask_2d(device, batch, hidden_size, lengths, t); - c = active_mask_2d.where_cond(&next_c, &c).to_concrete(); - h = active_mask_2d.where_cond(&next_h, &h).to_concrete(); + if let Some(active_masks) = &active_masks { + let active_mask_2d = active_masks + .narrow(1, t, 1) + .reshape([batch, hidden_size]) + .to_concrete(); + c = active_mask_2d.where_cond(&next_c, &c).to_concrete(); + h = active_mask_2d.where_cond(&next_h, &h).to_concrete(); + } else { + c = next_c; + h = next_h; + } + h = h.materialized(); + c = c.materialized(); - let output_t = h.clone().unsqueeze(1).to_concrete().materialized(); - outputs = outputs - .slice_assign([0..batch, t..(t + 1), 0..hidden_size], &output_t) - .materialized(); + let output_t = h.clone().unsqueeze(1).to_concrete(); + outputs = outputs.slice_assign([0..batch, t..(t + 1), 0..hidden_size], &output_t); } - let all_active = lengths.iter().all(|&length| length >= seq_len); if all_active { outputs } else { From 139324fb48e857810bf53e6f4c15fff41bd2782e Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Tue, 14 Apr 2026 21:31:04 -0500 Subject: [PATCH 19/34] ui mostly working --- demos/rgliner-web/src/main.rs | 24 ++++++++++++---------- models/rbert/src/raw/mdeberta/attention.rs | 13 +++++------- models/rbert/src/raw/mdeberta/layer.rs | 4 ++-- models/rbert/src/raw/mdeberta/model.rs | 22 ++++++++++++++++++-- 4 files changed, 40 insertions(+), 23 deletions(-) diff --git a/demos/rgliner-web/src/main.rs b/demos/rgliner-web/src/main.rs index 96c1aea13..b11fe3d80 100644 --- a/demos/rgliner-web/src/main.rs +++ b/demos/rgliner-web/src/main.rs @@ -116,16 +116,16 @@ fn App() -> Element { let mut relation_labels = use_signal(|| "founded by, located in, headquartered in".to_string()); let mut model = use_signal(|| None::); + // Explicit presence flag — written only on load/unload, never by the + // extraction pipeline. Subscribing to `model` directly would loop, since + // the extractor `model.set(Some(_))`s after every run. + let mut has_model = use_signal(|| false); let mut loading = use_signal(|| false); let mut running = use_signal(|| false); let mut error = use_signal(|| None::); let mut extraction = use_signal(Extraction::default); let mut status = use_signal(|| "idle".to_string()); - // Memoised so the effect below only re-runs when a model appears or - // disappears, not every time `run_extraction` writes the model back. - let model_ready = use_memo(move || model.read().is_some()); - // One extraction pipeline: watch the inputs, debounce, run. // `use_resource` re-runs the async body whenever any tracked signal // changes, cancelling the previous invocation — the debounce falls @@ -134,7 +134,7 @@ fn App() -> Element { let cur_text = text(); let ent_raw = entity_labels(); let rel_raw = relation_labels(); - let ready = model_ready(); + let ready = has_model(); tracing::info!("pipeline tick: ready={ready} text_len={}", cur_text.len()); if !ready { return; @@ -202,6 +202,7 @@ fn App() -> Element { choice.set(c); // Unload the current model; user will see a "load" hint. *model.write() = None; + has_model.set(false); extraction.set(Extraction::default()); }; @@ -217,6 +218,7 @@ fn App() -> Element { match build_model(selected).await { Ok(m) => { model.set(Some(m)); + has_model.set(true); status.set("ready".to_string()); } Err(e) => { @@ -228,7 +230,7 @@ fn App() -> Element { }); }; - let has_model = model.read().is_some(); + let is_loaded = has_model(); let model_matches = model .read() .as_ref() @@ -268,13 +270,13 @@ fn App() -> Element { } button { class: "load", - disabled: loading() || (has_model && model_matches), + disabled: loading() || (is_loaded && model_matches), onclick: on_load, if loading() { "loading…" - } else if has_model && model_matches { + } else if is_loaded && model_matches { "loaded" - } else if has_model { + } else if is_loaded { "reload" } else { "load" @@ -308,7 +310,7 @@ fn App() -> Element { } div { class: "status-line", - span { class: "dot", class: if running() || loading() { "busy" } else if error().is_some() { "err" } else if has_model { "ok" } else { "idle" } } + span { class: "dot", class: if running() || loading() { "busy" } else if error().is_some() { "err" } else if is_loaded { "ok" } else { "idle" } } span { class: "msg", "{status()}" } if let Some(e) = error() { span { class: "err-text", " · {e}" } @@ -323,7 +325,7 @@ fn App() -> Element { oninput: move |e: FormEvent| text.set(e.value()), } article { class: "article", - if has_model { + if is_loaded { { render_article(¤t_text, &cur_extraction) } } else { div { class: "placeholder", diff --git a/models/rbert/src/raw/mdeberta/attention.rs b/models/rbert/src/raw/mdeberta/attention.rs index 9fda2f79a..aaf6ff1ae 100644 --- a/models/rbert/src/raw/mdeberta/attention.rs +++ b/models/rbert/src/raw/mdeberta/attention.rs @@ -226,7 +226,7 @@ impl MDebertaAttention { hidden_states: &Tensor<3, f32>, rel_pos_emb: &Tensor<2, f32>, gather_idx: &GatherIndices, - attention_mask: Option<&Tensor<2, u32>>, + attention_bias: Option<&Tensor<4, f32>>, ) -> Tensor<3, f32> { use super::super::utils::split_heads; @@ -302,11 +302,8 @@ impl MDebertaAttention { .mul_scalar(self.scale); // Apply attention mask (broadcast bias to [batch, 1, 1, seq_len]) - let attn_scores = if let Some(mask) = attention_mask { - let mask_bias = super::super::utils::attention_mask_to_bias(mask); - let mask_bias_3d: Tensor<3, f32> = mask_bias.unsqueeze(1).to_concrete(); - let mask_bias_4d: Tensor<4, f32> = mask_bias_3d.unsqueeze(1).to_concrete(); - attn_scores.add_(&mask_bias_4d) + let attn_scores = if let Some(mask_bias) = attention_bias { + attn_scores.add_(mask_bias) } else { attn_scores }; @@ -361,9 +358,9 @@ impl DisentangledSelfAttention { hidden_states: &Tensor<3, f32>, rel_pos_emb: &Tensor<2, f32>, gather_idx: &GatherIndices, - attention_mask: Option<&Tensor<2, u32>>, + attention_bias: Option<&Tensor<4, f32>>, ) -> Tensor<3, f32> { self.attention - .forward_with_indices(hidden_states, rel_pos_emb, gather_idx, attention_mask) + .forward_with_indices(hidden_states, rel_pos_emb, gather_idx, attention_bias) } } diff --git a/models/rbert/src/raw/mdeberta/layer.rs b/models/rbert/src/raw/mdeberta/layer.rs index 0e098c594..450d6cfcd 100644 --- a/models/rbert/src/raw/mdeberta/layer.rs +++ b/models/rbert/src/raw/mdeberta/layer.rs @@ -54,12 +54,12 @@ impl MDebertaLayer { hidden_states: &Tensor<3, f32>, rel_pos_emb: &Tensor<2, f32>, gather_idx: &GatherIndices, - attention_mask: Option<&Tensor<2, u32>>, + attention_bias: Option<&Tensor<4, f32>>, ) -> Tensor<3, f32> { // Self-attention + residual + norm let attn_output = self.attention - .forward_with_rel(hidden_states, rel_pos_emb, gather_idx, attention_mask); + .forward_with_rel(hidden_states, rel_pos_emb, gather_idx, attention_bias); let hidden_states = self .attention_norm .forward(&hidden_states.add_(&attn_output)); diff --git a/models/rbert/src/raw/mdeberta/model.rs b/models/rbert/src/raw/mdeberta/model.rs index 52a6ccf98..63c5bedc2 100644 --- a/models/rbert/src/raw/mdeberta/model.rs +++ b/models/rbert/src/raw/mdeberta/model.rs @@ -104,10 +104,19 @@ impl MDebertaModel { &self.device, ); let rel_pos_emb = self.rel_pos_embedding.get_embeddings(); + let attention_bias = attention_mask.map(|mask| { + let mask_bias = super::super::utils::attention_mask_to_bias(mask); + mask_bias.unsqueeze(1).unsqueeze(1).to_concrete() + }); for layer in &self.layers { hidden_states = - layer.forward_with_rel(&hidden_states, &rel_pos_emb, &gather_idx, attention_mask); + layer.forward_with_rel( + &hidden_states, + &rel_pos_emb, + &gather_idx, + attention_bias.as_ref(), + ); } if let Some(ref proj) = self.output_proj { @@ -137,7 +146,16 @@ impl MDebertaModel { &self.device, ); let rel_pos_emb = self.rel_pos_embedding.get_embeddings(); - self.layers[0].forward_with_rel(hidden_states, &rel_pos_emb, &gather_idx, attention_mask) + let attention_bias = attention_mask.map(|mask| { + let mask_bias = super::super::utils::attention_mask_to_bias(mask); + mask_bias.unsqueeze(1).unsqueeze(1).to_concrete() + }); + self.layers[0].forward_with_rel( + hidden_states, + &rel_pos_emb, + &gather_idx, + attention_bias.as_ref(), + ) } /// Get the embedding dimension. From b33226777151c0521e603eec69e2c9514f5de4be Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Sat, 30 May 2026 14:25:26 -0500 Subject: [PATCH 20/34] update demo --- demos/rgliner-web/assets/style.css | 56 +++++++++++++++++------------- demos/rgliner-web/src/main.rs | 39 ++++++++++++++------- 2 files changed, 57 insertions(+), 38 deletions(-) diff --git a/demos/rgliner-web/assets/style.css b/demos/rgliner-web/assets/style.css index 1a2527e9f..e19f93ece 100644 --- a/demos/rgliner-web/assets/style.css +++ b/demos/rgliner-web/assets/style.css @@ -258,30 +258,30 @@ input:focus-visible { 50% { opacity: 1; transform: scale(1.05); } } -/* ── Split editor / reader ─────────────────────────────────── */ +/* ── Stage (single pane: article ↔ editor) ─────────────────── */ -.split { - display: grid; - grid-template-columns: 1fr 1fr; - gap: 3rem; - align-items: stretch; - min-height: 70vh; +.stage { + max-width: 720px; + margin: 0 auto; + min-height: 60vh; + position: relative; } textarea.editor { - background: transparent; - color: var(--ink-soft); + width: 100%; + background: var(--paper-deep); + color: var(--ink); border: none; - border-right: 1px dashed var(--rule); - padding: 0.5rem 2rem 0.5rem 0; - font-family: var(--mono); - font-size: 0.9rem; + border-left: 3px solid var(--accent); + padding: 1.5rem 1.5rem; + font-family: var(--serif); + font-size: 1.15rem; line-height: 1.75; - resize: none; + resize: vertical; outline: none; min-height: 60vh; + box-shadow: inset 2px 2px 0 rgba(20,17,13,0.04); } -textarea.editor:focus { color: var(--ink); } /* ── Article ───────────────────────────────────────────────── */ @@ -291,7 +291,21 @@ textarea.editor:focus { color: var(--ink); } line-height: 1.75; color: var(--ink); padding: 0.5rem 0; + cursor: text; +} +.article::after { + content: "double-click to edit"; + display: block; + margin-top: 2rem; + font-family: var(--mono); + font-size: 0.65rem; + letter-spacing: 0.2em; + text-transform: uppercase; + color: var(--ink-faint); + opacity: 0; + transition: opacity 0.18s; } +.article:hover::after { opacity: 0.7; } .article .placeholder { font-style: italic; @@ -465,15 +479,7 @@ textarea.editor:focus { color: var(--ink); } @media (max-width: 820px) { .reader { padding: 1.5rem 1.25rem 4rem; } - .split { - grid-template-columns: 1fr; - gap: 1.5rem; - } - textarea.editor { - border-right: none; - border-bottom: 1px dashed var(--rule); - padding: 0 0 1rem; - min-height: 30vh; - } + .stage { padding: 0; } + textarea.editor { padding: 1rem; min-height: 50vh; } [class*="select-trigger"], [class*="selectTrigger"] { min-width: 0 !important; } } diff --git a/demos/rgliner-web/src/main.rs b/demos/rgliner-web/src/main.rs index b11fe3d80..db060fb77 100644 --- a/demos/rgliner-web/src/main.rs +++ b/demos/rgliner-web/src/main.rs @@ -125,6 +125,7 @@ fn App() -> Element { let mut error = use_signal(|| None::); let mut extraction = use_signal(Extraction::default); let mut status = use_signal(|| "idle".to_string()); + let mut editing = use_signal(|| false); // One extraction pipeline: watch the inputs, debounce, run. // `use_resource` re-runs the async body whenever any tracked signal @@ -317,19 +318,31 @@ fn App() -> Element { } } - div { class: "split", - textarea { - class: "editor", - spellcheck: "false", - value: "{current_text}", - oninput: move |e: FormEvent| text.set(e.value()), - } - article { class: "article", - if is_loaded { - { render_article(¤t_text, &cur_extraction) } - } else { - div { class: "placeholder", - "load a model to begin reading." + div { class: "stage", + if editing() { + textarea { + class: "editor", + spellcheck: "false", + value: "{current_text}", + onmounted: move |e: MountedEvent| async move { + // Explicit focus so a click anywhere outside reliably + // fires `blur` and swaps us back into reading mode. + let _ = e.set_focus(true).await; + }, + oninput: move |e: FormEvent| text.set(e.value()), + onblur: move |_| editing.set(false), + } + } else { + article { + class: "article", + ondoubleclick: move |_| editing.set(true), + title: "double-click to edit", + if is_loaded { + { render_article(¤t_text, &cur_extraction) } + } else { + div { class: "placeholder", + "load a model to begin reading." + } } } } From 2fc0ba157d93ce2a7abe336f7af8cedaca8fcf88 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Sat, 30 May 2026 20:48:01 -0500 Subject: [PATCH 21/34] cache label embeddings --- models/rgliner/src/lib.rs | 44 +++++++++++++++++++++++++++++---------- 1 file changed, 33 insertions(+), 11 deletions(-) diff --git a/models/rgliner/src/lib.rs b/models/rgliner/src/lib.rs index c34d1ff91..31d7b4286 100644 --- a/models/rgliner/src/lib.rs +++ b/models/rgliner/src/lib.rs @@ -41,8 +41,9 @@ //! gliner.cache_labels(&labels).await?; //! //! // Fast inference with cached labels +//! let documents = ["Apple Inc. was founded by Steve Jobs.", "Microsoft is in Seattle."]; //! for text in documents { -//! let entities = gliner.extract_with_cached_labels(&text).await?; +//! let entities = gliner.extract_with_cached_labels(text).await?; //! // Process entities... //! } //! # Ok(()) @@ -322,11 +323,14 @@ impl Gliner { /// /// This significantly speeds up inference when using fixed label sets. pub async fn cache_labels(&mut self, labels: &[&str]) -> Result<(), GlinerError> { + // `materialized()` severs the lazy encoder graph into a standalone + // buffer; `to_concrete()` would only clone the lazy GPU tensor and + // re-run the encoder on every reuse. let label_embeddings = self .label_encoder .encode_labels(labels) .await? - .to_concrete(); + .materialized(); self.cached_labels = Some(CachedLabels::new( labels.iter().map(|s| s.to_string()).collect(), label_embeddings, @@ -427,16 +431,34 @@ impl Gliner { return Ok(Vec::new()); } - // Get label embeddings (compute if not cached or labels differ) - let label_embeddings = if let Some(ref cached) = self.cached_labels { - let cached_labels: Vec<&str> = cached.labels.iter().map(|s| s.as_str()).collect(); - if cached_labels == labels { - cached.embeddings.clone() - } else { - self.label_encoder.encode_labels(labels).await? - } + // Get label embeddings, reusing the cache when the label set is + // unchanged. The label encoder is independent of the input text, so a + // fixed label set (the common case when extracting over many texts) + // only needs encoding once. Auto-populate the cache on a miss so the + // default `extract`/`extract_batch` path amortizes label encoding + // without requiring an explicit `cache_labels` call. + let labels_match = self.cached_labels.as_ref().is_some_and(|cached| { + cached.labels.len() == labels.len() + && cached + .labels + .iter() + .zip(labels.iter()) + .all(|(cached, label)| cached == label) + }); + let label_embeddings = if labels_match { + self.cached_labels.as_ref().unwrap().embeddings.clone() } else { - self.label_encoder.encode_labels(labels).await? + // `materialized()` (not `to_concrete()`) resolves the encoder graph + // into a standalone output buffer, severing the lazy graph so reuse + // does not re-run the label encoder. `to_concrete()` only clones the + // lazy GPU tensor, which would drag the whole encoder into every + // subsequent extract. + let embeddings = self.label_encoder.encode_labels(labels).await?.materialized(); + self.cached_labels = Some(CachedLabels::new( + labels.iter().map(|s| s.to_string()).collect(), + embeddings.clone(), + )); + embeddings }; self.extract_internal_batch(texts, labels, &label_embeddings) From 4b12376945fececa85cace4cee04e88aeed58745 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Sat, 30 May 2026 20:48:01 -0500 Subject: [PATCH 22/34] mmap model weights --- Cargo.lock | 1 + models/kalosm-llama/Cargo.toml | 3 +++ models/kalosm-llama/src/model/mod.rs | 2 +- models/kalosm-llama/src/source.rs | 39 +++++++++++++++++++++++++--- 4 files changed, 41 insertions(+), 4 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 0a52a39f4..7900d9d66 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5510,6 +5510,7 @@ dependencies = [ "kalosm-sample", "kalosm-streams", "kalosm-tokenizer", + "memmap2", "minijinja", "minijinja-contrib", "pollster", diff --git a/models/kalosm-llama/Cargo.toml b/models/kalosm-llama/Cargo.toml index ee7180056..c88718c42 100644 --- a/models/kalosm-llama/Cargo.toml +++ b/models/kalosm-llama/Cargo.toml @@ -39,6 +39,9 @@ image = { workspace = true, optional = true } [target.'cfg(target_arch = "wasm32")'.dependencies] getrandom = { version = "0.3", features = ["wasm_js"], optional = true } +[target.'cfg(not(target_arch = "wasm32"))'.dependencies] +memmap2 = "0.9" + [dev-dependencies] ahash = "0.8.12" pollster = "0.4" diff --git a/models/kalosm-llama/src/model/mod.rs b/models/kalosm-llama/src/model/mod.rs index 06351fdf6..416c09809 100644 --- a/models/kalosm-llama/src/model/mod.rs +++ b/models/kalosm-llama/src/model/mod.rs @@ -512,7 +512,7 @@ where // Read metadata from all model files let mut files_with_metadata = Vec::new(); for bytes in &model_bytes { - let mut cursor = std::io::Cursor::new(bytes); + let mut cursor = std::io::Cursor::new(&bytes[..]); let metadata = GgufMetadata::read(&mut cursor)?; files_with_metadata.push((metadata, cursor)); } diff --git a/models/kalosm-llama/src/source.rs b/models/kalosm-llama/src/source.rs index b48dcbbbb..0d4019cf6 100644 --- a/models/kalosm-llama/src/source.rs +++ b/models/kalosm-llama/src/source.rs @@ -2,6 +2,14 @@ use fusor_gguf::GgufReadError; use kalosm_common::CacheError; use kalosm_model_types::{FileLoadingProgress, FileSource}; +/// Backing storage for a loaded model file. On native targets this is a +/// memory-mapped view of the cached GGUF (lazy, zero upfront copy); on WASM it +/// is an in-memory buffer. Both deref to `[u8]`. +#[cfg(not(target_arch = "wasm32"))] +pub(crate) type ModelBytes = memmap2::Mmap; +#[cfg(target_arch = "wasm32")] +pub(crate) type ModelBytes = Vec; + #[cfg(feature = "hf-config-json")] use crate::raw::RopeScalingConfig; @@ -225,11 +233,36 @@ impl LlamaSource { pub(crate) async fn model( &self, - mut progress: impl FnMut(FileLoadingProgress), - ) -> Result>, LlamaSourceError> { + #[cfg_attr(target_arch = "wasm32", allow(unused_mut))] mut progress: impl FnMut( + FileLoadingProgress, + ), + ) -> Result, LlamaSourceError> { let mut model_bytes = Vec::new(); for file in &self.model { - model_bytes.push(self.cache.get_bytes(file, &mut progress).await?); + // Memory-map the (cached) model file rather than reading the whole + // multi-GB GGUF into a Vec up front. The OS pages weights in lazily + // as each tensor is read during upload, eliminating one full read + + // allocation pass over the entire file. WASM has no mmap, so it + // falls back to reading the bytes into memory. + #[cfg(not(target_arch = "wasm32"))] + { + let path = self.cache.get(file, &mut progress).await?; + let handle = std::fs::File::open(&path).map_err(CacheError::from)?; + // SAFETY: model cache files are treated as immutable for the + // lifetime of the mapping (we only ever read them). + let mmap = unsafe { memmap2::Mmap::map(&handle) }.map_err(CacheError::from)?; + // We read the whole file (every tensor) exactly once, in order. + // Kick off sequential read-ahead so a cold page cache prefetches + // from disk in the background and overlaps with parsing/upload, + // instead of stalling on a page fault per tensor. + let _ = mmap.advise(memmap2::Advice::Sequential); + let _ = mmap.advise(memmap2::Advice::WillNeed); + model_bytes.push(mmap); + } + #[cfg(target_arch = "wasm32")] + { + model_bytes.push(self.cache.get_bytes(file, &mut progress).await?); + } } Ok(model_bytes) } From f1b63610f00f7a31c65304993dcfe3519a4d1757 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 1 Jun 2026 20:42:56 -0500 Subject: [PATCH 23/34] use mapped buffer init for model weights --- fusor-ml/core/src/device.rs | 11 ++++++++ fusor-ml/core/src/quantized/mod.rs | 5 +++- fusor-ml/tile-ir-runtime/src/buffer_pool.rs | 29 +++++++++++++++++++++ 3 files changed, 44 insertions(+), 1 deletion(-) diff --git a/fusor-ml/core/src/device.rs b/fusor-ml/core/src/device.rs index e0c17ec28..964e588b3 100644 --- a/fusor-ml/core/src/device.rs +++ b/fusor-ml/core/src/device.rs @@ -371,6 +371,17 @@ impl Device { self.inner.buffer_pool.create_buffer_init(data, usage) } + /// Create a fresh, write-once buffer initialized from `data` via + /// `mapped_at_creation` (single memcpy, no staging belt). For long-lived + /// data like model weights — not transient inference allocations. + pub fn create_buffer_init_mapped( + &self, + data: &[u8], + usage: wgpu::BufferUsages, + ) -> Arc { + self.inner.buffer_pool.create_buffer_init_mapped(data, usage) + } + /// Get or create a buffer of the specified size. pub fn create_buffer_init_iter( &self, diff --git a/fusor-ml/core/src/quantized/mod.rs b/fusor-ml/core/src/quantized/mod.rs index 49848d53c..44777d08f 100644 --- a/fusor-ml/core/src/quantized/mod.rs +++ b/fusor-ml/core/src/quantized/mod.rs @@ -474,7 +474,10 @@ impl QMatrix { (unsupported, _) => return Err(GgufReadError::UnsupportedDType(unsupported as u32)), }; let datatype = ty; - let buffer = device.create_buffer_init( + // Weights are written once and live for the model's lifetime, so map the + // buffer at creation and copy straight in (one memcpy) rather than going + // through the queue staging belt. + let buffer = device.create_buffer_init_mapped( &storage_bytes, wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC diff --git a/fusor-ml/tile-ir-runtime/src/buffer_pool.rs b/fusor-ml/tile-ir-runtime/src/buffer_pool.rs index 06885d7fa..be44f366e 100644 --- a/fusor-ml/tile-ir-runtime/src/buffer_pool.rs +++ b/fusor-ml/tile-ir-runtime/src/buffer_pool.rs @@ -184,6 +184,35 @@ impl BufferPool { buffer } + /// Create a fresh buffer initialized from `data` via `mapped_at_creation`, + /// copying the bytes straight into the buffer's own memory. + /// + /// Unlike [`create_buffer_init`], this skips the queue staging belt (no + /// extra staging copy and no queued GPU copy on submit) — a single memcpy. + /// It always allocates a new buffer (no pooling), so it is meant for + /// long-lived, write-once data like model weights, not transient + /// inference-time allocations. + pub fn create_buffer_init_mapped( + &self, + data: &[u8], + usage: wgpu::BufferUsages, + ) -> Arc { + let padded_len = padded_copy_size(data.len() as u64); + let buffer = self.device.create_buffer(&wgpu::BufferDescriptor { + label: Some("Weight Buffer"), + size: padded_len, + usage, + mapped_at_creation: true, + }); + { + let mut view = buffer.slice(..).get_mapped_range_mut(); + view[..data.len()].copy_from_slice(data); + view[data.len()..].fill(0); + } + buffer.unmap(); + Arc::new(buffer) + } + /// Get or create a buffer initialized from a byte iterator. pub fn create_buffer_init_iter( &self, From 0e327f79d6795e4df2d5be0a11b5eee7f73aea31 Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Mon, 1 Jun 2026 21:08:15 -0500 Subject: [PATCH 24/34] trim down the diff --- fusor-ml/core/src/device.rs | 11 ---- fusor-ml/core/src/quantized/mod.rs | 5 +- fusor-ml/tile-ir-runtime/src/buffer_pool.rs | 29 ---------- .../kalosm-llama/examples/startup_timing.rs | 29 ++++++++++ models/kalosm-llama/examples/throughput.rs | 53 +++++++++++++++++++ 5 files changed, 83 insertions(+), 44 deletions(-) create mode 100644 models/kalosm-llama/examples/startup_timing.rs create mode 100644 models/kalosm-llama/examples/throughput.rs diff --git a/fusor-ml/core/src/device.rs b/fusor-ml/core/src/device.rs index d2ec8730b..bb9e8f5f2 100644 --- a/fusor-ml/core/src/device.rs +++ b/fusor-ml/core/src/device.rs @@ -371,17 +371,6 @@ impl Device { self.inner.buffer_pool.create_buffer_init(data, usage) } - /// Create a fresh, write-once buffer initialized from `data` via - /// `mapped_at_creation` (single memcpy, no staging belt). For long-lived - /// data like model weights — not transient inference allocations. - pub fn create_buffer_init_mapped( - &self, - data: &[u8], - usage: wgpu::BufferUsages, - ) -> Arc { - self.inner.buffer_pool.create_buffer_init_mapped(data, usage) - } - /// Get or create a buffer of the specified size. pub fn create_buffer_init_iter( &self, diff --git a/fusor-ml/core/src/quantized/mod.rs b/fusor-ml/core/src/quantized/mod.rs index f3ee11248..38ba34760 100644 --- a/fusor-ml/core/src/quantized/mod.rs +++ b/fusor-ml/core/src/quantized/mod.rs @@ -474,10 +474,7 @@ impl QMatrix { (unsupported, _) => return Err(GgufReadError::UnsupportedDType(unsupported as u32)), }; let datatype = ty; - // Weights are written once and live for the model's lifetime, so map the - // buffer at creation and copy straight in (one memcpy) rather than going - // through the queue staging belt. - let buffer = device.create_buffer_init_mapped( + let buffer = device.create_buffer_init( &storage_bytes, wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC diff --git a/fusor-ml/tile-ir-runtime/src/buffer_pool.rs b/fusor-ml/tile-ir-runtime/src/buffer_pool.rs index be44f366e..06885d7fa 100644 --- a/fusor-ml/tile-ir-runtime/src/buffer_pool.rs +++ b/fusor-ml/tile-ir-runtime/src/buffer_pool.rs @@ -184,35 +184,6 @@ impl BufferPool { buffer } - /// Create a fresh buffer initialized from `data` via `mapped_at_creation`, - /// copying the bytes straight into the buffer's own memory. - /// - /// Unlike [`create_buffer_init`], this skips the queue staging belt (no - /// extra staging copy and no queued GPU copy on submit) — a single memcpy. - /// It always allocates a new buffer (no pooling), so it is meant for - /// long-lived, write-once data like model weights, not transient - /// inference-time allocations. - pub fn create_buffer_init_mapped( - &self, - data: &[u8], - usage: wgpu::BufferUsages, - ) -> Arc { - let padded_len = padded_copy_size(data.len() as u64); - let buffer = self.device.create_buffer(&wgpu::BufferDescriptor { - label: Some("Weight Buffer"), - size: padded_len, - usage, - mapped_at_creation: true, - }); - { - let mut view = buffer.slice(..).get_mapped_range_mut(); - view[..data.len()].copy_from_slice(data); - view[data.len()..].fill(0); - } - buffer.unmap(); - Arc::new(buffer) - } - /// Get or create a buffer initialized from a byte iterator. pub fn create_buffer_init_iter( &self, diff --git a/models/kalosm-llama/examples/startup_timing.rs b/models/kalosm-llama/examples/startup_timing.rs new file mode 100644 index 000000000..324d7467f --- /dev/null +++ b/models/kalosm-llama/examples/startup_timing.rs @@ -0,0 +1,29 @@ +#![recursion_limit = "512"] + +use fusor::Device; +use kalosm_llama::prelude::*; +use std::time::Instant; + +fn main() { + pollster::block_on(async { + let device = Device::new().await.expect("gpu"); + let t = Instant::now(); + let model = Llama::builder() + .with_source(LlamaSource::llama_3_1_8b_chat()) + .with_device(device) + .build() + .await + .unwrap(); + println!("startup {:.3}s", t.elapsed().as_secs_f64()); + + let mut stream = model("The capital of France is"); + let mut out = String::new(); + for _ in 0..10 { + match stream.next().await { + Some(tok) => out.push_str(&tok), + None => break, + } + } + println!("GEN: \"The capital of France is{out}\""); + }); +} diff --git a/models/kalosm-llama/examples/throughput.rs b/models/kalosm-llama/examples/throughput.rs new file mode 100644 index 000000000..577c3e93c --- /dev/null +++ b/models/kalosm-llama/examples/throughput.rs @@ -0,0 +1,53 @@ +//! Decode throughput probe against the locally-cached Llama 3.1 8B Q4_K_M. +//! Warms up, then measures steady-state decode tok/s. Release builds only. + +use kalosm_llama::*; +use kalosm_model_types::ModelLoadingProgress; +use prelude::{StreamExt, TextCompletionModelExt}; + +fn env_usize(name: &str, default: usize) -> usize { + std::env::var(name) + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(default) +} + +fn main() { + pollster::block_on(async { + let warmup = env_usize("WARMUP", 8); + let measured = env_usize("TOKENS", 128); + let prompt = "Write a detailed explanation of how a CPU executes instructions:"; + + let load = std::time::Instant::now(); + let model = Llama::builder() + .with_source(LlamaSource::llama_3_1_8b_chat()) + .build_with_loading_handler(|_: ModelLoadingProgress| {}) + .await + .unwrap(); + println!("load: {:.2}s", load.elapsed().as_secs_f64()); + + let mut stream = model.complete(prompt).take(warmup + measured); + for _ in 0..warmup { + if stream.next().await.is_none() { + eprintln!("stream ended during warmup"); + return; + } + } + + let start = std::time::Instant::now(); + let mut tokens = 0usize; + while tokens < measured { + if stream.next().await.is_none() { + break; + } + tokens += 1; + } + let elapsed = start.elapsed(); + let tps = tokens as f64 / elapsed.as_secs_f64(); + let per_token_ms = elapsed.as_secs_f64() * 1_000.0 / tokens.max(1) as f64; + println!( + "decode: {tokens} tok in {:.3}s => {tps:.2} tok/s ({per_token_ms:.3} ms/tok)", + elapsed.as_secs_f64() + ); + }); +} From 5e3b91cefd106001440e2b05f66dea9e37090a4d Mon Sep 17 00:00:00 2001 From: Evan Almloff Date: Sat, 20 Jun 2026 14:45:52 -0500 Subject: [PATCH 25/34] clean up --- models/rgliner/.gitignore | 2 + .../convert_to_gguf.cpython-311.pyc | Bin 27125 -> 0 bytes .../rgliner/scripts/convert_relex_to_gguf.py | 228 +--- .../rgliner/{ => scripts}/convert_to_gguf.py | 231 +--- models/rgliner/scripts/debug_compare.py | 173 --- models/rgliner/scripts/gguf_common.py | 280 ++++ models/rgliner/src/lib.rs | 1192 +---------------- models/rgliner/src/raw/bilstm.rs | 145 -- models/rgliner/src/raw/label_encoder.rs | 76 -- models/rgliner/src/raw/mod.rs | 3 - models/rgliner/src/raw/scorer.rs | 45 - models/rgliner/src/raw/span_layer.rs | 133 +- models/rgliner/src/raw/text_encoder.rs | 9 - models/rgliner/src/relation_decoding.rs | 39 - models/rgliner/src/relex.rs | 1131 +--------------- models/rgliner/src/relex_tokenization.rs | 103 +- models/rgliner/src/source.rs | 126 +- models/rgliner/src/tokenization.rs | 32 +- models/rgliner/tests/relex_batch.rs | 141 ++ 19 files changed, 562 insertions(+), 3527 deletions(-) create mode 100644 models/rgliner/.gitignore delete mode 100644 models/rgliner/__pycache__/convert_to_gguf.cpython-311.pyc rename models/rgliner/{ => scripts}/convert_to_gguf.py (68%) delete mode 100644 models/rgliner/scripts/debug_compare.py create mode 100644 models/rgliner/scripts/gguf_common.py delete mode 100644 models/rgliner/src/raw/scorer.rs create mode 100644 models/rgliner/tests/relex_batch.rs diff --git a/models/rgliner/.gitignore b/models/rgliner/.gitignore new file mode 100644 index 000000000..7a60b85e1 --- /dev/null +++ b/models/rgliner/.gitignore @@ -0,0 +1,2 @@ +__pycache__/ +*.pyc diff --git a/models/rgliner/__pycache__/convert_to_gguf.cpython-311.pyc b/models/rgliner/__pycache__/convert_to_gguf.cpython-311.pyc deleted file mode 100644 index 9af6e8f7d95b1773518b6ecd0101655f766c5e0e..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 27125 zcmeHwS!^81m0(s?)?G!`Nfr-LMN$XP;!ROJ#HK`vIw(plQQa+eHwq0`c}m3$c)@#00qiz6f6dlCFwr_)Bj75?^rkN$3iApQxxWG@^7 zFJCO{2;vbz5fcQZqx3P|gbqLT6MFnKOc?N!oFMVjIAMgRA!eF2PndNm4jHq|S|_Zt zwh7y;eZoHLm~dc!eatyggyR%X6vNXPbIrOZ+_Rnu4}_a0yp(yuN9iW~lx3oXvQCu3 zdw{Y{lu`DHa>_ALK{+QXsiKJkRPjU=<(fE1xhH~@XQCS7)IgkCh;sj1Z&@=i49i4lVGy-QI3kM#uc3H-^|L?cy0mC#L8Df|U!2W_L9=;E8r+3+cSS=LF= zHs#xti7I>7I1!@Cam=PElB#%@oM^EUR)VUk%-y6X*(ek(gBPqz&~X-s(?OPr#1n`aKn3y@ zn~P916-?0c5hlXYv3vEl&k(Q2>U9!%Il|7=lagU5eorEYqm!&ejz$x#WE@|ZkJ0sd zsd#28JhN~kOwHYmBTAAx5s%C#X6D$OR}^voGXs$P3@VHG3{~uBBmzbRObD0}upnSX z0QI)d>+nvGBqHX@m!MiX;P0xE(`(6K2Tom{2gf5K zsZSC&vtdM1hZ3N4zy#fv>nW9{&TwQ08(G%K)K53XTd<>zVRc=1D?0>Br)cTqv$xC(T|!l(s3+rc{~R2sMpN)CyeOTzoQsB}ZhOA4 zh?AM?CZdclO3$L#+dZT2LM|H`F2M5Unu@}%-Iv<(5 zB^iK5rs^c!6w?I7k<4-gN>h2xFeVU|Q|AJT`cs9R0~(1-0H6vz-bZy0>Q;{lo?6jU zyKGKdtlu!dXa1(?p?!7uacye)drcoT39c5=)goA0MN2DhY0coQm<3Bfv;_F(XI*(eql-rPr*0h)0YFeCUr@Z_gwSb2I%(Xk$_Q^5x?@_LjY)-T-f;qj<)OBoNCtBLpkR(#h zJZ|2FzEWzAz$jVRJJ82CbFt`)oVlhnaaZ)r-|xE&!7Mc`KeT0FYx9_pYt?rZdMF+# zevh_sC{I<*33UqiJF}nB$h7}ah$8cAFm~q=>HU54_)j$onRcs_mvB=*OuXHI=L!4a zsu8~GfS+I2YHR6!V`2MM@pb2MW&eErQ;jSC{CqY4!F=6$T>I+2+w0!Pwf)UK{!qRK zpuO%ql4@F}j<{cj+aj94XD-UZ)@R*lUZrfA>X!T+;zP|C_l}NLQnIQky*zfkojaVJ z`DRy5L3!167Z#DhC<_u6Soefru6vHt_n&0Yx}8D26W(q!A=xp-1+T>u3-j}H3=7lb z2*X6~1?d<)3sMOjzczHkv3kE`izcG+1RDV%hD6?&n~O1MA}MI5!B z*QFiKZ;ZS*^35|3FQmxFalzLm`nm*1x9I5R9o?$XwJ!e1u~Z}#S&s<59?{n$IF5^s z<9zmh?krt3ewYxPwW70ji#UX+7yDLjzjtr>UfSh;WPV`&(DZ>lHT-uiLZDX+^a`#% z(bdPh`c$E7hlN0^7-$t-ZKA7<&)zST6n@pvD+Kz)K%e017hV0ltA8&f8unKGgR_q(1$Tq!ZrCDR z&ewD*%DY$N-_@?sKWKQ;AOyO^K$qa^7G2%Et2WTMWF( zTp-g!(k^_d>8g374*P@C%*nmfbYGGtaWfY$rI4p&^ABq&aOO2FFxPH-3u<^#k7qFG z!(MQ0w-$uij$$2>j{>-bvz>N|TqLK7MWd3F!^Dd2?f>}}w82HwqIuDhw6I7xoVhZq zwrUnhKQq_OT%9D%tYYPPGbf#RVpMZy(H#E{(5=Z*Ac%d@&8-Nwp*sh3krXy_W$03W zXnvFClw$3Wr4BMPWqJKse|5`PG$+lJiG%v5%pCOKlw#(cjiI4bn)NTBM5K0V3nn#> z_ubC*s=j5unVZwA^Po;e9i}5#H2KP0$sK9nm85l_Eu8kpv@`Umro_B8co#vg?fJz zjcRygGcOipMwucQWp-E>LcdeUukxv=M@wyffJ+L_ZNsk}QiJl@wkzNLw?K1lLqBzS zbVS@&#osPIp=-|f==JQI~VSLSfxeM2cvIMW*!TIFX9t&D&v>pmh zfl@Oh>F?fPdVm^U&*9bG4|I}|LMxch&$pDrocu>w^+G z2Z}1JrPQw7uu?UAacJb+DanBsVy~UPa`oJ$F}#w(dM>QplZ>EuolW2+7;_%dlbjErW~JCXoBCL5hzm|IBTwWegoYdi>7qPKNt{J{Z>$1ItXXW*T@6Q%D; zrUh6-LM&4xfyyesL#Qz^6MG${EcT{wY0lSmt=xUgOn*9V6eteG|C zU8Kn706p=K@EwLVq|?1}=wbcpxZtc2ov?~LW^GJGUZ|FlH{}_;;Aj^e?YyI1wTyZk z6MP+_uS0NjijGd+(WwrdTs!onNjP#$+xOv0h!uwlx#R+ZrTX}!$r`yv> zR~-EA@oyi0eC4C-LPZOt1cVP%r>=Z!Vy#=KY}>3jzEN@fr!_yR7b=Fu3P^U#M!5aC z0#kykQFJx(u14+RhJnXDysJwU$~W{2fdMfvAh=G5t`ofL#0yvDrmJ?tRht^;8wdD9 zCj{4^=o;i*gK4NY>Z4mAAmCj!FDef3)g9|Mgz6rl;`q}mLiyl|8DUJrk`}@RR}EBO zMdf#oeEZ0^>OQDjzK{--FP}?S9bM~tS|&7}6sk@=yC)pDygZg^h$|zjAP31W^I>7t zd2BuTZ2aeM2_09&jw@RP>1@OR6sGT6S3t%A%vl{yRj00J3bjgp*YRz~w~9U}5?n_` z*HPYe6f&CfeN?f=2-R&uMLXoUya(kJ3r1GWsiQ)mKGzhXHEvl^D$_dXU$~QD@Hy&k zzeXVMx1H-ZzxmB}-d{30!dk=^HKz`JR1Xih1xrY@g!t^8Koi=3HP&*;WB8?qyi{WT zWtk4(7#O^O0SK~n2!|zGI1Kh83o+z&gu}NNBC$+}ISd9rli@Hkga)c00w8%J7-UGm zyoSK*03>sq!s9X1kGvWYK*BUUlH(!h2LxV@%3~6m{V@m}hrh&s1F)3)rO6tetl1(B z#(~wVTLe6`z@|6WrVP;ejkRm1@IwnM#97@n9r8tEgt2Qifgfu%PeAMe@6WvXFonxP z8p~23=z@n9IO>h(b!(Gb1bS`*U(@v&TegToAeS1I73_PE;aWIKN!H28 zXm){#s|<=@0wyohr(imxUZzuET(sRso<+<0%#{Qqh-omS`@Zgp;bZXsBp0@*XHn&4 zuDroTO5vLgl?Br3DaoRNRV)ya21>K^0zm=EDpZR%b4q`hyu$n;Z~C%mq)h0iroYRS zN0Gxl+_@!wzSK}|lRB?VgHjm{(?8R;HQ>$DP?Xz5St;8SyQUt@@l%w8 za>l?^bkPD6y8V@|sGxMY2_TxbZ=XLsk%SO+$HZ;keFmKy#&fzZh9vhy?Jn%RP18z z1KLygODsaGfax7+VMFmvhs~51n;2e9b6iH zZnnOA_pkr*oxfa(3FcbS3`4iY*!`p1`T%Ql!xSiv@`MR`9jEqi$PQSw1Db&hX2DU_ZPOHVC3b|4H>z?JK(Nx|PF z`kR(*FF+-8hqoW#$piS)`?FI2P>*co`DcSVfRZCyVdxt>spz`|8VAs5{l&k5 zg)lNK#DZ7wD7=tKxn>+dLfS$^@SzM^^MhUJ!pOX5x@Jlvz;Q-=1;(RLP!kOR2z~am%RkJ>ynTm3ziLKvnG-_aGJAGj26i zvuE5|>d>BXVI#zzagR_(_l#Rd)$bX%VIS?Rk!spA-DWDZXWSO5bEQm+)^!c1z2YrYtn&rPJ{sYnnF_4z$3 z@WH@4A=L!zF-1Suuedv0EXo{iTtnqOg-el~5^; zF7Lom1=vF{dyI!(zkTaYh_m6BFm_ur0dpbFDSOMYL!2cv8=DWk9z_9DQ}HnD`JuxL z^SJuA391XZsfW;3meB15nz@H~kLodeJH+|47_k2%7Y@y*fC4s? z$zeb<2q-9IDxHF@6tG&UEDTYU$|L7N*~%a{AkC#=6BQMOW)jNKvEcTB+DuVMoq{L( zuAu7;VMc^iguOCT=E1%Rx_E}JQ{~B7nRq-DOR&Ik5;?gL+8L%888ok%{a~x zcSQT0rQx)_SoVCCzj{-!-xBS&0N0emJ^Ac{V81NdFDp<-`LSuiJ|o&^6sW7tR8p{a ziuTSN=Bhmv7VKT3y-SH_OI_09*;1o|y-l>Y<)b=MF~Qy;+B+0c`M}__fM7o_+Rv+e z`}l#&g8fy|{%S5hpz=en3-&ie`x^?+A^udMXBU4AawL~gQS#8bL$D8s_5oE|t?Lb0 z-yG{g>o&pOC))cIDQ$<>z7FXOi}qm!UN_$ZWxpueFRGH+j%3|hlVCq4+K(xsH}RoU zS(=KxH?MtFun&s%K_Hr~cc_im1bdTcZ&IanmOuZdV1G-rzoj6%t>>9v6c;e`_wd_3 zn>l7Gvi`zeymSG}J|)8x=-KShxnMIb7J5h)kT1jJhf{B=L*fm*_I=^jN=??`cu*6H{HNuWli z28O8JthER)^cnHEwxiEz}4TJ?D%$y5TnCU~9hXt-=FOc!+3uTb+=%U1; z75+eFzsGX>aJlV^W>8Ms?+}a$y>f0vxN!IVN#Gm4yeD6cw2_1i8qeLK=D zkO%CUX6cSJ15}xsiuw)mxl;+tqna&!V6frbfg6QVgrXgB3#ACfJK`2f5nMas7D^G^ zJK`2f5j;EM9;S|{dA8`?5%(xnr@{5@*vjjZ#cJGL+DyYfa2t~v$<{7(o01xd*)F)v zN%x-f8%k<=$1ZeR*qreZXk2{i9=_)kTdB4^RvWjd5*d`X)+M=Lr?WLYx@m+8XUzJBTq?# zlQVbBV`$S_EQ$Y@T0&{~c!*$r28Kl--HG`YOaGa3oHK;nCC34_IUjnnH0w^5Zf|Hw z9ZLpMNJm?6X}FdD@F+0VIO_JuHsuj-=FIb+Xe3=JMPJC{CjT|zbs zQ?+7!(hhC)c+!t&DkXUe+s#y9l&Z|9mv<%o$r3PMgx-R+ZLf%Cc%JIrn zvz;;x@5^`mb`BMMJ1dgqJI;(ws#^+-COhH&UNa-=6teLuNTXuU-}1`jPPhlOv!KNT zJK`QJ&>mFnh+9n!tMj{fa7Wx)>a+$oxFhai>Wl`rdPiQ5BoC-@cd3`NNsVS?7u<76 zjmBjc-1EtUd(Oj!Wc8kLN2!Z@#vMy)G&!37s{UG9)qMoUo=e;GKiU>kU-$MMvkzqQ2TAM0CNzRz#%yD&!W|6BWN zvh0uit<`-0_VTS&%}B1O`y$L6yFp8KmHV%!)W^GF{Yl}eZph;06hnhB9@J-dvuD16 z6;DDw>a98W7Yf!P^{GB6XVomK%!31E;eIv zZgW<Fv_>z=gt%hcMP#;=wN&cm^uZvf?+yM|7i1nl7XyJ9 zGDE)skSkYijm}8Lun7}vo37ej-^BT*X*lOA5(`DZR0>SGC&7*o4B2IyD$bD&MT*DY z>RLYg9b8;x)>oNU7@4{QBN!y};9N3+&}`%@;k!`^43VrkCxga=90pbeAVy2pywfrp zVQ$gPYpq;qhEQVRhFlV)`ao9gt^yz9RFr|!=ZZ3%&KQJA_L(RJijr^wM1|2Mu#=Iz z@rBuNPN5W@p(9j+L8kAL8--->vGygwAd!{iR;qJkO0R%qU9#N)W3d~u$FrYpj$C1J26G6_a=&f=$hL&%-3KiZPj3RR=&Jb#!;tauy zCpjZNc9HXL(=8=G?y5OuY_vUGKQEw&t#%^k^U&9quWc zsAm1Jl20<>9X1T}Jw)#Nh@2GpN$)3#0{70Zs0V90E32LAv;>${+ylL(1 zdQ>QTMJ#(oAWN3cub9AQ)o9~wl>&J{B$0t-I#9D2sM`qC34sPN0Ji-)V=o3@cPx|3 zQ`lg$XQh)b>VbQ;`-6VjwRYr5lk5^K$3@F=9^atcvba{dHp~Iu9C+@lUbZe<(^mJ& zwGC@IZ!K5&ytZMj;H?$UeKiQ@)1m-h_R9JU0Jw#sA+czPw+sQf0|$AsEL~Q?lO->L z4V%H%jbN(~Y!`!YSORJ6#z3%CE<0AH()Px+ByaD7JKYXPN>ry{yGjf1XnE1n1^bl1 zywuUeJDPw4MS)ej;5;Nc4=tU6G>rC5vUGziThEJWbkU@71C;%a4kF zShSuH>iWdGzNaU}xM0$yah{e4zWOP9S?lvX>{(eycp$hP5+4I{(A->wod= zq|h`XHjN16S&=-;lV=sN^GN4~rb}YeC4sywl9ze%vI41Xz5Czx{_9@;#AW`AAOwcR zz%WnxbI7UD^}~X@M|Af9uD^EEe{{otl&>3l)+zYUivF{Zg*@Hvrxw9IB)W%|&Z{`o zwSG?^heUFSCx;Zp?Ri=+ki#Mg2@Wf~o7Qd?Rudm+TkjLdL6IEf$w38nC*Qp-!&}y? zwyVXKfAFt!Lfc8P?W91S63J6Mc}hVIm1~>YYg~hh=n=^tp6to=u~$E=<*QnS@;0%& zjkmoj-vW6>B(Lz~l^4PG&0x<)utx~?ib3dd#m4T>5y-UW=e}An3&4kdEL*os7VE*Z zx9ritgMs%49}aGM>o&Y~g115RHf(xB8{UxMZ56$(o8ImXZ@1v>5xqUj!^^|Jd0t$x zMVOF*c^0G{f#p&3OS?)QIUhLRFMe3O={mgOIxKI>+H^H;xEke?Pc~gW8?GL~)hoJs zmxrE%9XlKy;(ug+U|&6za{q1F_shOl@lnNQb;m|^hfv)mR%ebhd0rBDcyr5W@D+Vd z0QfZm|Ib7em%KN<5_vc+6jg~uRa<7l=|-8{GDF_LZ2e0Y-i9@>0z=FI+^OsGbKUpj za@i$VPKcHhyyZmNQnYC)+pv_a)~CkTjtiD9(bC0Rz*4>I1!?>1(RW6_Humn=(%5rk zxZj+z1Ato~Ln0aC$q@EF^UZd!haY?}nCcU}Euyz&{gB}8e(Dpv1HA1_?iR=qksRU4 z5wIjEJG2>S+z2$TjqqLPg}?;&p&>Ce^sGh)&PsVk;G@AAt-1O9Lc%aW7f7&j1PKcfp(7&^&{9)9;z5SxM zA8^Z#Y?d`_l;NIC3T3CovePiG=1Ab4oE;&2)tkP<8@|K*k-=wP!FO8ponE^5B43Q- zPocAASd%LTwD-}ATNsKMV?ghsAK(YMc>Pz4?XJ=$a5lj zjwjD4(rj6;$Z<5&_g+VRuR|#B6w5n#+w1ZzkZ*|O8$9{Oi(uDgux}&SCjTPFW&W<3Cz-04D)AccIN~NWuBV_hDkC-bR=|F@VxN z(VPvj2-{A!7u2F{BrL3*vKIwgMV^>HhOH`}z@L19SbL|~7RsuLZQDm|*b<_O?Sw6k zMc>y^#ZO!xL!Lg#p=b8lH+^hDG2PjViMI@Lx<-OXI; zNu6TzpDurf5cnelCwx8eX;Vcf!Poo8xWlK^GO_Ldl=ZWaNpu3j?S&`v5?DGS4+ zq)}8LjJM>e7k~9O$r53v;r}EgXsKw3nO;DLEL=vfjC`y-gknFu30ksY%rJN+@VZ4Z zVxfm*l2_RIRVD08G+|sg6A)pmEI&!Kd}JiZ)sNA-rxd^v@!aW$i9O%;2JgHsIIoM& z>q}>H^S89cz0&#I=}No(Y3Jycl`sZACyYketOsBV0|$|vxTh_imEkP%p_d~2Rw#t* z@~2%#czmawEn9Y+gb5`90Onx;iWEusaU5UT347bO$xRXA@NC*EH|&*yy-KuKEe(H4 zmP76zmc!|14=*6sfwZUOVQ9+=CD4{7S1Q?z5H_ph1$UT1%5pHMMXhSAzJp1?li5Jv z9R!vTXafL;Jl;#7vozS5So8*iMnYKE-&YbK~9#x6oH77-a;8%TaWm`9MMneeVb+j|D zpfu3QH_S-{P9tyvfdK#oM_P7w=4CX>6Hi|447Y0D+f?sen((f$I=n{{-V<(WgPgfNOGRm~k0)Rstry-EPM@DA+=U3>3%Dq* zuLKZ)5YTYKp%DNuwVi=?sc6fT)al@yf$ibv*#Fnc(7!ZP>R=--fE1j00FQO~@yz&) z`b7tS`I>n7E&jq=;+ePk@Fag`QhbHNXj*V$NT+k5vk}m9E%TJ$RB%jB_{-D&vMuXt zdL10eiD0_Oy=6d-7jEB{5&Oe>%Clw0P7C4kDzW7_8OV^eIv96yU{pttj%9AcP{td| Z(xhc+_-kk0J+p*=TLywG2XEZe{y#v+kaz$9 diff --git a/models/rgliner/scripts/convert_relex_to_gguf.py b/models/rgliner/scripts/convert_relex_to_gguf.py index aa6303a1f..37b2e6b4f 100644 --- a/models/rgliner/scripts/convert_relex_to_gguf.py +++ b/models/rgliner/scripts/convert_relex_to_gguf.py @@ -24,223 +24,12 @@ from pathlib import Path from typing import Any, Dict, List, Tuple -# In-process formats: we quantise per-tensor via `gguf.quants.quantize` so the -# file layout is byte-identical to what fusor expects. -IN_PROCESS_QUANTS = { - "f32", - "f16", - "bf16", - "q4_0", - "q4_1", - "q5_0", - "q5_1", - "q8_0", -} -# k-quants (Q4_K, Q5_K, Q6_K) need `llama-quantize` - the Python `gguf` package -# doesn't implement their `quantize_blocks` path. We write an f32 GGUF (with a -# temporary `general.architecture = "bert"` masquerade and the minimal bert.* -# metadata llama-quantize's loader validates) and then shell out to the binary. -LLAMA_QUANT_TYPES = { - "q2_k", - "q3_k", - "q3_k_s", - "q3_k_m", - "q3_k_l", - "q4_k", - "q4_k_s", - "q4_k_m", - "q5_k", - "q5_k_s", - "q5_k_m", - "q6_k", -} -QUANT_TYPES = IN_PROCESS_QUANTS | LLAMA_QUANT_TYPES - import numpy as np import torch from huggingface_hub import snapshot_download -# GGUF constants (same as convert_to_gguf.py) -GGUF_MAGIC = 0x46554747 -GGUF_VERSION = 3 - -GGUF_TYPE_UINT8 = 0 -GGUF_TYPE_INT8 = 1 -GGUF_TYPE_UINT16 = 2 -GGUF_TYPE_INT16 = 3 -GGUF_TYPE_UINT32 = 4 -GGUF_TYPE_INT32 = 5 -GGUF_TYPE_FLOAT32 = 6 -GGUF_TYPE_BOOL = 7 -GGUF_TYPE_STRING = 8 -GGUF_TYPE_ARRAY = 9 -GGUF_TYPE_UINT64 = 10 -GGUF_TYPE_INT64 = 11 -GGUF_TYPE_FLOAT64 = 12 - -GGML_TYPE_F32 = 0 -GGML_TYPE_F16 = 1 -GGML_TYPE_Q4_0 = 2 -GGML_TYPE_Q4_1 = 3 -GGML_TYPE_Q5_0 = 6 -GGML_TYPE_Q5_1 = 7 -GGML_TYPE_Q8_0 = 8 -GGML_TYPE_BF16 = 30 - - -def _ggml_type_for(quant: str) -> int: - return { - "f32": GGML_TYPE_F32, - "f16": GGML_TYPE_F16, - "bf16": GGML_TYPE_BF16, - "q4_0": GGML_TYPE_Q4_0, - "q4_1": GGML_TYPE_Q4_1, - "q5_0": GGML_TYPE_Q5_0, - "q5_1": GGML_TYPE_Q5_1, - "q8_0": GGML_TYPE_Q8_0, - }[quant] - - -def _gguf_block_quant(data: np.ndarray, quant: str) -> np.ndarray: - """Return the packed byte array for a block-quantised tensor. - - Uses `gguf.quants.quantize` (Python package from llama.cpp). The inner row - dimension must be a multiple of 32 for q4_0/q5_0/q8_0, else we fall back to - f16 for that tensor (caller decides). - """ - import gguf # Deferred import - only needed for block quants. - qtype = { - "q4_0": gguf.GGMLQuantizationType.Q4_0, - "q4_1": gguf.GGMLQuantizationType.Q4_1, - "q5_0": gguf.GGMLQuantizationType.Q5_0, - "q5_1": gguf.GGMLQuantizationType.Q5_1, - "q8_0": gguf.GGMLQuantizationType.Q8_0, - }[quant] - return gguf.quants.quantize(data, qtype) - - -class _U32(int): - """Force a metadata int to be written as GGUF u32 (needed by llama-quantize for some keys).""" - - -class GGUFWriter: - """Simple GGUF file writer.""" - - def __init__(self, path: str): - self.path = path - self.metadata: Dict[str, Any] = {} - self.tensors: List[Tuple[str, np.ndarray, int]] = [] - - def add_metadata(self, key: str, value: Any): - self.metadata[key] = value - - def add_tensor( - self, - name: str, - data: np.ndarray, - ggml_type: int = GGML_TYPE_F32, - shape: Tuple[int, ...] = None, - ): - """Add a tensor. `shape` is the logical (un-packed) shape; defaults to data.shape.""" - if shape is None: - shape = tuple(data.shape) - self.tensors.append((name, data, ggml_type, tuple(shape))) - - def _write_string(self, f, s: str): - encoded = s.encode('utf-8') - f.write(struct.pack('> 16) & 0xFFFF).astype(np.uint16) - - self._write_string(f, name) - f.write(struct.pack(' Tuple[Dict[str, torch.Tensor], Dict, str, str]: @@ -542,19 +331,6 @@ def convert_relex_to_gguf( print("\nConversion complete!") -def _name_for_type(ggml_type: int) -> str: - return { - GGML_TYPE_F32: "F32", - GGML_TYPE_F16: "F16", - GGML_TYPE_BF16: "BF16", - GGML_TYPE_Q4_0: "Q4_0", - GGML_TYPE_Q4_1: "Q4_1", - GGML_TYPE_Q5_0: "Q5_0", - GGML_TYPE_Q5_1: "Q5_1", - GGML_TYPE_Q8_0: "Q8_0", - }.get(ggml_type, f"type={ggml_type}") - - def main(): parser = argparse.ArgumentParser(description="Convert GLiNER-RelEx PyTorch models to GGUF") parser.add_argument( diff --git a/models/rgliner/convert_to_gguf.py b/models/rgliner/scripts/convert_to_gguf.py similarity index 68% rename from models/rgliner/convert_to_gguf.py rename to models/rgliner/scripts/convert_to_gguf.py index 9c01c028c..8b408567f 100644 --- a/models/rgliner/convert_to_gguf.py +++ b/models/rgliner/scripts/convert_to_gguf.py @@ -3,7 +3,7 @@ Convert GLiNER PyTorch models to GGUF format. Usage: - python convert_to_gguf.py --model knowledgator/gliner-bi-edge-v2.0 --output gliner-bi-edge-v2.0.gguf + python scripts/convert_to_gguf.py --model knowledgator/gliner-bi-edge-v2.0 --output gliner-bi-edge-v2.0.gguf This script converts both: 1. The main text encoder (ModernBERT/Ettin) + span layer weights @@ -25,233 +25,8 @@ import torch from huggingface_hub import hf_hub_download, snapshot_download - -# Quantization targets. In-process ones are packed directly by `gguf.quants.quantize`; -# k-quants are produced by shelling out to `llama-quantize` post-hoc. -IN_PROCESS_QUANTS = { - "f32", "f16", "bf16", - "q4_0", "q4_1", "q5_0", "q5_1", "q8_0", -} -LLAMA_QUANT_TYPES = { - "q2_k", "q3_k", "q3_k_s", "q3_k_m", "q3_k_l", - "q4_k", "q4_k_s", "q4_k_m", - "q5_k", "q5_k_s", "q5_k_m", - "q6_k", -} -QUANT_TYPES = IN_PROCESS_QUANTS | LLAMA_QUANT_TYPES - - -# GGUF constants -GGUF_MAGIC = 0x46554747 # "GGUF" in little-endian -GGUF_VERSION = 3 - -# GGUF data types -GGUF_TYPE_UINT8 = 0 -GGUF_TYPE_INT8 = 1 -GGUF_TYPE_UINT16 = 2 -GGUF_TYPE_INT16 = 3 -GGUF_TYPE_UINT32 = 4 -GGUF_TYPE_INT32 = 5 -GGUF_TYPE_FLOAT32 = 6 -GGUF_TYPE_BOOL = 7 -GGUF_TYPE_STRING = 8 -GGUF_TYPE_ARRAY = 9 -GGUF_TYPE_UINT64 = 10 -GGUF_TYPE_INT64 = 11 -GGUF_TYPE_FLOAT64 = 12 - -# GGML tensor types -GGML_TYPE_F32 = 0 -GGML_TYPE_F16 = 1 -GGML_TYPE_Q4_0 = 2 -GGML_TYPE_Q4_1 = 3 -GGML_TYPE_Q5_0 = 6 -GGML_TYPE_Q5_1 = 7 -GGML_TYPE_Q8_0 = 8 -GGML_TYPE_Q8_1 = 9 -GGML_TYPE_BF16 = 30 - - -def _ggml_type_for(quant: str) -> int: - return { - "f32": GGML_TYPE_F32, - "f16": GGML_TYPE_F16, - "bf16": GGML_TYPE_BF16, - "q4_0": GGML_TYPE_Q4_0, - "q4_1": GGML_TYPE_Q4_1, - "q5_0": GGML_TYPE_Q5_0, - "q5_1": GGML_TYPE_Q5_1, - "q8_0": GGML_TYPE_Q8_0, - }[quant] - - -def _name_for_type(ggml_type: int) -> str: - return { - GGML_TYPE_F32: "F32", - GGML_TYPE_F16: "F16", - GGML_TYPE_BF16: "BF16", - GGML_TYPE_Q4_0: "Q4_0", - GGML_TYPE_Q4_1: "Q4_1", - GGML_TYPE_Q5_0: "Q5_0", - GGML_TYPE_Q5_1: "Q5_1", - GGML_TYPE_Q8_0: "Q8_0", - }.get(ggml_type, f"type={ggml_type}") - - -def _gguf_block_quant(data: np.ndarray, quant: str) -> np.ndarray: - """Return the packed byte array for a block-quantised tensor.""" - import gguf - qtype = { - "q4_0": gguf.GGMLQuantizationType.Q4_0, - "q4_1": gguf.GGMLQuantizationType.Q4_1, - "q5_0": gguf.GGMLQuantizationType.Q5_0, - "q5_1": gguf.GGMLQuantizationType.Q5_1, - "q8_0": gguf.GGMLQuantizationType.Q8_0, - }[quant] - return gguf.quants.quantize(data, qtype) - - -class _U32(int): - """Force a metadata int to be written as GGUF u32 (required by llama-quantize's arch loader).""" - - -class GGUFWriter: - """Simple GGUF file writer.""" - - def __init__(self, path: str): - self.path = path - self.metadata: Dict[str, Any] = {} - # (name, data, ggml_type, logical_shape) - self.tensors: List[Tuple[str, np.ndarray, int, Tuple[int, ...]]] = [] - - def add_metadata(self, key: str, value: Any): - """Add metadata key-value pair.""" - self.metadata[key] = value - - def add_tensor( - self, - name: str, - data: np.ndarray, - ggml_type: int = GGML_TYPE_F32, - shape: Tuple[int, ...] = None, - ): - """Add a tensor. `shape` is the logical (un-packed) shape; defaults to data.shape.""" - if shape is None: - shape = tuple(data.shape) - self.tensors.append((name, data, ggml_type, tuple(shape))) - - def _write_string(self, f, s: str): - """Write a GGUF string (length-prefixed UTF-8).""" - encoded = s.encode('utf-8') - f.write(struct.pack('> 16) & 0xFFFF).astype(np.uint16) - - # Write tensor info - # GGUF stores dimensions in reverse order (column-major) - # Reader reverses them back, so we write reversed to get original order - self._write_string(f, name) - f.write(struct.pack(' Tuple[Dict[str, torch.Tensor], Dict]: diff --git a/models/rgliner/scripts/debug_compare.py b/models/rgliner/scripts/debug_compare.py deleted file mode 100644 index 342b1a98b..000000000 --- a/models/rgliner/scripts/debug_compare.py +++ /dev/null @@ -1,173 +0,0 @@ -#!/usr/bin/env python3 -"""Debug script to compare encoder outputs between Python and Rust implementations.""" - -import torch -import numpy as np -from transformers import AutoModel, AutoTokenizer, DebertaV2Model -from gliner import GLiNER - -# Load the model -model = GLiNER.from_pretrained("knowledgator/gliner-relex-multi-v1.0") -tokenizer = model.data_processor.transformer_tokenizer - -# Text and labels -text = "Apple was founded by Steve Jobs in California." -entity_labels = ["person", "organization", "location"] -relation_labels = ["founded by", "located in"] - -# Build the same prompt as our Rust code -# IMPORTANT: Don't add [CLS] - tokenizer adds it automatically -def build_prompt(): - """Build the prompt in the same format as Rust (without [CLS]).""" - parts = [] - for label in entity_labels: - parts.append("<>") - parts.append(label) - parts.append("[SEP]") - for label in relation_labels: - parts.append("<>") - parts.append(label) - parts.append("[SEP]") - parts.append(text) - return " ".join(parts) - -prompt = build_prompt() -print(f"Prompt: {prompt}") - -# Tokenize (tokenizer should add [CLS] automatically) -encoding = tokenizer( - prompt, - return_tensors="pt", - padding=False, - truncation=True, - max_length=512, - add_special_tokens=True, -) -input_ids = encoding["input_ids"] -attention_mask = encoding["attention_mask"] - -print(f"\nToken IDs (first 30): {input_ids[0, :30].tolist()}") -print(f"Token count: {input_ids.shape[1]}") - -# Decode tokens to see what they are -tokens = tokenizer.convert_ids_to_tokens(input_ids[0].tolist()) -print(f"\nFirst 30 tokens: {tokens[:30]}") - -# Find <> token positions -ent_token_id = tokenizer.convert_tokens_to_ids("<>") -print(f"\n<> token ID: {ent_token_id}") - -ent_positions = [] -for i, tok in enumerate(input_ids[0]): - if tok.item() == ent_token_id: - ent_positions.append(i) -print(f"<> positions: {ent_positions}") - -# Find the encoder - try different attribute names -print(f"\nModel type: {type(model)}") -print(f"Model attributes: {[a for a in dir(model) if not a.startswith('_')]}") - -# The encoder is typically model.model in GLiNER -encoder = None -if hasattr(model, 'model'): - encoder = model.model - print(f"Found encoder at model.model: {type(encoder)}") -elif hasattr(model, 'token_rep_layer'): - encoder = model.token_rep_layer - print(f"Found encoder at model.token_rep_layer: {type(encoder)}") - -# Check what encoder contains -if encoder is not None: - print(f"Encoder attributes: {[a for a in dir(encoder) if not a.startswith('_')]}") - - # Try to find the actual DeBERTa model - if hasattr(encoder, 'deberta'): - deberta = encoder.deberta - print(f"Found DeBERTa at encoder.deberta: {type(deberta)}") - elif hasattr(encoder, 'model'): - deberta = encoder.model - print(f"Found DeBERTa at encoder.model: {type(deberta)}") - else: - deberta = encoder - -# Run the encoder -with torch.no_grad(): - # Get the token_rep_layer (DeBERTa encoder) - token_rep_layer = model.model.token_rep_layer - print(f"\ntoken_rep_layer type: {type(token_rep_layer)}") - print(f"token_rep_layer children: {list(token_rep_layer.named_children())}") - - # Call the token_rep_layer - outputs = token_rep_layer(input_ids, attention_mask=attention_mask) - if hasattr(outputs, 'last_hidden_state'): - hidden_states = outputs.last_hidden_state - elif isinstance(outputs, tuple): - hidden_states = outputs[0] - else: - hidden_states = outputs - -print(f"\nEncoder (token_rep_layer) output shape: {hidden_states.shape}") - -# Stats -hs = hidden_states[0].numpy() -print(f"Encoder output stats: mean={hs.mean():.6f}, std={hs.std():.6f}, min={hs.min():.6f}, max={hs.max():.6f}") - -# Print values at <> positions -print(f"\nEncoder output at <> positions (first 5 values):") -for pos in ent_positions: - vals = hs[pos, :5] - print(f" pos {pos}: [{', '.join(f'{v:.4f}' for v in vals)}]") - -# Also check other positions for comparison -print(f"\nEncoder output at other positions:") -for pos in [0, 2, 4, 10, 17]: - if pos < hs.shape[0]: - vals = hs[pos, :5] - print(f" pos {pos}: [{', '.join(f'{v:.4f}' for v in vals)}]") - -# Check prompt_rep_layer -print("\n--- Prompt Rep Layer ---") -ent_embs = hidden_states[0, ent_positions, :] # [n_labels, hidden] -print(f"Entity embeddings shape: {ent_embs.shape}") -print(f"Entity embeddings stats: mean={ent_embs.mean():.6f}, std={ent_embs.std():.6f}") - -# Apply prompt_rep_layer from model.model -prompt_rep = model.model.prompt_rep_layer # This should be a Sequential or MLP -projected = prompt_rep(ent_embs) -print(f"After prompt_rep_layer: shape={projected.shape}") -print(f"After prompt_rep_layer stats: mean={projected.mean():.6f}, std={projected.std():.6f}") - -# Check what prompt_rep_layer consists of -print(f"\nPrompt rep layer structure:") -for name, module in prompt_rep.named_modules(): - if name: - print(f" {name}: {module}") - -# Print projected values for first 5 dims -print(f"\nProjected entity embeddings (first 5 values):") -for i, label in enumerate(entity_labels): - vals = projected[i, :5].detach().numpy() - print(f" {label}: [{', '.join(f'{v:.4f}' for v in vals)}]") - -# Now let's check the raw token embeddings (before any attention) -print("\n--- Raw Token Embeddings (before transformer) ---") -bert_layer = token_rep_layer.bert_layer -deberta_model = bert_layer.model - -# Get raw embeddings -word_embs = deberta_model.embeddings(input_ids) -print(f"Raw embeddings shape: {word_embs.shape}") - -word_embs_np = word_embs[0].detach().numpy() -print(f"Raw embeddings stats: mean={word_embs_np.mean():.6f}, std={word_embs_np.std():.6f}") - -print(f"\nRaw embeddings at <> positions (first 5 values):") -for pos in ent_positions: - vals = word_embs_np[pos, :5] - print(f" pos {pos}: [{', '.join(f'{v:.4f}' for v in vals)}]") - -print(f"\nRaw embeddings at other positions (first 5 values):") -for pos in [0, 2, 4, 10, 17]: - if pos < word_embs_np.shape[0]: - vals = word_embs_np[pos, :5] - print(f" pos {pos}: [{', '.join(f'{v:.4f}' for v in vals)}]") diff --git a/models/rgliner/scripts/gguf_common.py b/models/rgliner/scripts/gguf_common.py new file mode 100644 index 000000000..7bb4cba45 --- /dev/null +++ b/models/rgliner/scripts/gguf_common.py @@ -0,0 +1,280 @@ +"""Shared GGUF-writing infrastructure for the GLiNER converters. + +Both `convert_to_gguf.py` (bi-encoder) and `convert_relex_to_gguf.py` (RelEx) +emit GGUF files with the byte layout fusor expects. This module holds the +quantisation tables, GGUF/GGML constants, and the `GGUFWriter` they share; each +converter keeps its own model-specific weight mapping and `main()`. +""" + +import struct +from typing import Any, Dict, List, Tuple + +import numpy as np + +__all__ = [ + "IN_PROCESS_QUANTS", + "LLAMA_QUANT_TYPES", + "QUANT_TYPES", + "GGUF_MAGIC", + "GGUF_VERSION", + "GGUF_TYPE_UINT8", + "GGUF_TYPE_INT8", + "GGUF_TYPE_UINT16", + "GGUF_TYPE_INT16", + "GGUF_TYPE_UINT32", + "GGUF_TYPE_INT32", + "GGUF_TYPE_FLOAT32", + "GGUF_TYPE_BOOL", + "GGUF_TYPE_STRING", + "GGUF_TYPE_ARRAY", + "GGUF_TYPE_UINT64", + "GGUF_TYPE_INT64", + "GGUF_TYPE_FLOAT64", + "GGML_TYPE_F32", + "GGML_TYPE_F16", + "GGML_TYPE_Q4_0", + "GGML_TYPE_Q4_1", + "GGML_TYPE_Q5_0", + "GGML_TYPE_Q5_1", + "GGML_TYPE_Q8_0", + "GGML_TYPE_Q8_1", + "GGML_TYPE_BF16", + "_ggml_type_for", + "_name_for_type", + "_gguf_block_quant", + "_U32", + "GGUFWriter", +] + + +# Quantization targets. In-process ones are packed directly by `gguf.quants.quantize`; +# k-quants are produced by shelling out to `llama-quantize` post-hoc. +IN_PROCESS_QUANTS = { + "f32", "f16", "bf16", + "q4_0", "q4_1", "q5_0", "q5_1", "q8_0", +} +LLAMA_QUANT_TYPES = { + "q2_k", "q3_k", "q3_k_s", "q3_k_m", "q3_k_l", + "q4_k", "q4_k_s", "q4_k_m", + "q5_k", "q5_k_s", "q5_k_m", + "q6_k", +} +QUANT_TYPES = IN_PROCESS_QUANTS | LLAMA_QUANT_TYPES + + +# GGUF constants +GGUF_MAGIC = 0x46554747 # "GGUF" in little-endian +GGUF_VERSION = 3 + +# GGUF data types +GGUF_TYPE_UINT8 = 0 +GGUF_TYPE_INT8 = 1 +GGUF_TYPE_UINT16 = 2 +GGUF_TYPE_INT16 = 3 +GGUF_TYPE_UINT32 = 4 +GGUF_TYPE_INT32 = 5 +GGUF_TYPE_FLOAT32 = 6 +GGUF_TYPE_BOOL = 7 +GGUF_TYPE_STRING = 8 +GGUF_TYPE_ARRAY = 9 +GGUF_TYPE_UINT64 = 10 +GGUF_TYPE_INT64 = 11 +GGUF_TYPE_FLOAT64 = 12 + +# GGML tensor types +GGML_TYPE_F32 = 0 +GGML_TYPE_F16 = 1 +GGML_TYPE_Q4_0 = 2 +GGML_TYPE_Q4_1 = 3 +GGML_TYPE_Q5_0 = 6 +GGML_TYPE_Q5_1 = 7 +GGML_TYPE_Q8_0 = 8 +GGML_TYPE_Q8_1 = 9 +GGML_TYPE_BF16 = 30 + + +def _ggml_type_for(quant: str) -> int: + return { + "f32": GGML_TYPE_F32, + "f16": GGML_TYPE_F16, + "bf16": GGML_TYPE_BF16, + "q4_0": GGML_TYPE_Q4_0, + "q4_1": GGML_TYPE_Q4_1, + "q5_0": GGML_TYPE_Q5_0, + "q5_1": GGML_TYPE_Q5_1, + "q8_0": GGML_TYPE_Q8_0, + }[quant] + + +def _name_for_type(ggml_type: int) -> str: + return { + GGML_TYPE_F32: "F32", + GGML_TYPE_F16: "F16", + GGML_TYPE_BF16: "BF16", + GGML_TYPE_Q4_0: "Q4_0", + GGML_TYPE_Q4_1: "Q4_1", + GGML_TYPE_Q5_0: "Q5_0", + GGML_TYPE_Q5_1: "Q5_1", + GGML_TYPE_Q8_0: "Q8_0", + }.get(ggml_type, f"type={ggml_type}") + + +def _gguf_block_quant(data: np.ndarray, quant: str) -> np.ndarray: + """Return the packed byte array for a block-quantised tensor. + + Uses `gguf.quants.quantize` (Python package from llama.cpp). The inner row + dimension must be a multiple of 32 for q4_0/q5_0/q8_0, else we fall back to + f16 for that tensor (caller decides). + """ + import gguf # Deferred import - only needed for block quants. + qtype = { + "q4_0": gguf.GGMLQuantizationType.Q4_0, + "q4_1": gguf.GGMLQuantizationType.Q4_1, + "q5_0": gguf.GGMLQuantizationType.Q5_0, + "q5_1": gguf.GGMLQuantizationType.Q5_1, + "q8_0": gguf.GGMLQuantizationType.Q8_0, + }[quant] + return gguf.quants.quantize(data, qtype) + + +class _U32(int): + """Force a metadata int to be written as GGUF u32 (required by llama-quantize's arch loader).""" + + +class GGUFWriter: + """Simple GGUF file writer.""" + + def __init__(self, path: str): + self.path = path + self.metadata: Dict[str, Any] = {} + # (name, data, ggml_type, logical_shape) + self.tensors: List[Tuple[str, np.ndarray, int, Tuple[int, ...]]] = [] + + def add_metadata(self, key: str, value: Any): + """Add metadata key-value pair.""" + self.metadata[key] = value + + def add_tensor( + self, + name: str, + data: np.ndarray, + ggml_type: int = GGML_TYPE_F32, + shape: Tuple[int, ...] = None, + ): + """Add a tensor. `shape` is the logical (un-packed) shape; defaults to data.shape.""" + if shape is None: + shape = tuple(data.shape) + self.tensors.append((name, data, ggml_type, tuple(shape))) + + def _write_string(self, f, s: str): + """Write a GGUF string (length-prefixed UTF-8).""" + encoded = s.encode('utf-8') + f.write(struct.pack('> 16) & 0xFFFF).astype(np.uint16) + + # Write tensor info + # GGUF stores dimensions in reverse order (column-major) + # Reader reverses them back, so we write reversed to get original order + self._write_string(f, name) + f.write(struct.pack(') { use fusor::{Device, Tensor, VarBuilder}; use kalosm_common::Cache; -use kalosm_model_types::ModelLoadingProgress; +use kalosm_model_types::{FileSource, ModelLoadingProgress}; use rbert::BertSource; use std::sync::Arc; use tokenizers::Tokenizer; @@ -126,6 +126,21 @@ async fn default_device() -> Device { Device::gpu().await.unwrap_or_else(|_| Device::cpu()) } +/// Download a model artifact from `source` through `cache`, reporting progress +/// under the label `"