Skip to main content

problemreductions/models/algebraic/
closest_vector_problem.rs

1//! Closest Vector Problem (CVP).
2//!
3//! Given an integer lattice basis `B` and a target vector `t`, find integer
4//! coefficients `x` minimizing `||Bx - t||_2`.
5
6use crate::registry::{ConstructionError, CreateSpec, ProblemSchemaEntry, VariantDimension};
7use crate::traits::{EvaluationError, Problem};
8use crate::types::Min;
9use serde::{Deserialize, Serialize};
10
11/// Target coordinate domains supported by [`ClosestVectorProblem`].
12pub trait ClosestVectorTarget: Clone + std::fmt::Debug + 'static {
13    /// Registered value of the `target` variant dimension.
14    const NAME: &'static str;
15
16    /// Validate one stored target coordinate.
17    fn validate(&self, index: usize) -> Result<(), ConstructionError>;
18
19    /// Convert one coordinate for numerical evaluation and solving.
20    fn to_f64(&self) -> Result<f64, EvaluationError>;
21}
22
23impl ClosestVectorTarget for i64 {
24    const NAME: &'static str = "i64";
25
26    fn validate(&self, _index: usize) -> Result<(), ConstructionError> {
27        Ok(())
28    }
29
30    fn to_f64(&self) -> Result<f64, EvaluationError> {
31        crate::types::i64_to_exact_f64(*self)
32            .map_err(|error| EvaluationError::InexactFloatConversion(error.to_string()))
33    }
34}
35
36impl ClosestVectorTarget for f64 {
37    const NAME: &'static str = "f64";
38
39    fn validate(&self, index: usize) -> Result<(), ConstructionError> {
40        if self.is_finite() {
41            Ok(())
42        } else {
43            Err(ConstructionError::NonFiniteFloat(format!(
44                "target coordinate at index {index} must be finite"
45            )))
46        }
47    }
48
49    fn to_f64(&self) -> Result<f64, EvaluationError> {
50        Ok(*self)
51    }
52}
53
54macro_rules! cvp_create_spec {
55    ($name:ident, $target:ty) => {
56        #[derive(Debug, Deserialize, crate::CreateSpec)]
57        struct $name {
58            /// Integer basis matrix as semicolon-separated column vectors.
59            #[create(codec = "semicolon-separated")]
60            basis: Vec<Vec<i64>>,
61            /// Target vector.
62            #[create(name = "target_vec", codec = "comma-separated")]
63            target: Vec<$target>,
64        }
65
66        impl TryFrom<$name> for ClosestVectorProblem<$target> {
67            type Error = ConstructionError;
68
69            fn try_from(spec: $name) -> Result<Self, Self::Error> {
70                ClosestVectorProblem::new(spec.basis, spec.target)
71            }
72        }
73    };
74}
75
76cvp_create_spec!(ClosestVectorProblemI64CreateSpec, i64);
77cvp_create_spec!(ClosestVectorProblemF64CreateSpec, f64);
78
79inventory::submit! {
80    ProblemSchemaEntry {
81        name: "ClosestVectorProblem",
82        display_name: "Closest Vector Problem",
83        aliases: &["CVP"],
84        dimensions: &[VariantDimension::new("target", "i64", &["i64", "f64"])],
85        category: crate::registry::ProblemCategory::Algebraic,
86        module_path: module_path!(),
87        description: "Find the closest point in an integer lattice to a target vector",
88        fields: ClosestVectorProblemI64CreateSpec::FIELDS,
89    }
90}
91
92/// Euclidean Closest Vector Problem over an integer lattice basis.
93#[derive(Debug, Clone, Serialize)]
94pub struct ClosestVectorProblem<T = i64> {
95    /// Basis matrix stored as column vectors.
96    basis: Vec<Vec<i64>>,
97    /// Target vector in the ambient space.
98    target: Vec<T>,
99}
100
101impl<T: ClosestVectorTarget> ClosestVectorProblem<T> {
102    /// Construct a CVP instance with a full-column-rank integer basis.
103    pub fn new(basis: Vec<Vec<i64>>, target: Vec<T>) -> Result<Self, ConstructionError> {
104        let ambient_dimension = target.len();
105        for (index, coordinate) in target.iter().enumerate() {
106            coordinate.validate(index)?;
107        }
108        for (index, column) in basis.iter().enumerate() {
109            if column.len() != ambient_dimension {
110                return Err(ConstructionError::Conversion(format!(
111                    "basis vector {index} has length {}, expected {ambient_dimension}",
112                    column.len()
113                )));
114            }
115        }
116        if basis.len() > ambient_dimension {
117            return Err(ConstructionError::Conversion(format!(
118                "{} basis vectors cannot be independent in ambient dimension {ambient_dimension}",
119                basis.len()
120            )));
121        }
122        if independent_rows(&basis, ambient_dimension)?.is_none() {
123            return Err(ConstructionError::Conversion(
124                "closest-vector basis columns must be linearly independent".into(),
125            ));
126        }
127        Ok(Self { basis, target })
128    }
129
130    /// Number of basis vectors.
131    pub fn num_basis_vectors(&self) -> usize {
132        self.basis.len()
133    }
134
135    /// Dimension of the ambient space.
136    pub fn ambient_dimension(&self) -> usize {
137        self.target.len()
138    }
139
140    /// Integer basis columns.
141    pub fn basis(&self) -> &[Vec<i64>] {
142        &self.basis
143    }
144
145    /// Target coordinates.
146    pub fn target(&self) -> &[T] {
147        &self.target
148    }
149
150    pub(crate) fn independent_rows(&self) -> Result<Vec<usize>, ConstructionError> {
151        independent_rows(&self.basis, self.ambient_dimension())?.ok_or_else(|| {
152            ConstructionError::Conversion(
153                "closest-vector basis columns must be linearly independent".into(),
154            )
155        })
156    }
157}
158
159fn independent_rows(
160    basis: &[Vec<i64>],
161    ambient_dimension: usize,
162) -> Result<Option<Vec<usize>>, ConstructionError> {
163    let num_columns = basis.len();
164    if num_columns == 0 {
165        return Ok(Some(Vec::new()));
166    }
167
168    let mut matrix = (0..ambient_dimension)
169        .map(|row| basis.iter().map(|column| column[row]).collect::<Vec<_>>())
170        .collect::<Vec<_>>();
171    let mut previous_pivot = 1_i64;
172    let mut row_indices = (0..ambient_dimension).collect::<Vec<_>>();
173
174    for column in 0..num_columns {
175        let Some(pivot_row) = (column..ambient_dimension).find(|&row| matrix[row][column] != 0)
176        else {
177            return Ok(None);
178        };
179        matrix.swap(column, pivot_row);
180        row_indices.swap(column, pivot_row);
181        let pivot = matrix[column][column];
182
183        for row in (column + 1)..ambient_dimension {
184            for next_column in (column + 1)..num_columns {
185                let left = matrix[row][next_column]
186                    .checked_mul(pivot)
187                    .ok_or_else(rank_overflow)?;
188                let right = matrix[row][column]
189                    .checked_mul(matrix[column][next_column])
190                    .ok_or_else(rank_overflow)?;
191                let numerator = left.checked_sub(right).ok_or_else(rank_overflow)?;
192                matrix[row][next_column] = numerator
193                    .checked_div(previous_pivot)
194                    .ok_or_else(rank_overflow)?;
195            }
196            matrix[row][column] = 0;
197        }
198        previous_pivot = pivot;
199    }
200    row_indices.truncate(num_columns);
201    Ok(Some(row_indices))
202}
203
204fn rank_overflow() -> ConstructionError {
205    ConstructionError::IntegerOverflow("checking closest-vector basis rank".into())
206}
207
208impl<'de, T> Deserialize<'de> for ClosestVectorProblem<T>
209where
210    T: ClosestVectorTarget + Deserialize<'de>,
211{
212    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
213    where
214        D: serde::Deserializer<'de>,
215    {
216        #[derive(Deserialize)]
217        struct Raw<T> {
218            basis: Vec<Vec<i64>>,
219            target: Vec<T>,
220        }
221
222        let raw = Raw::deserialize(deserializer)?;
223        Self::new(raw.basis, raw.target).map_err(serde::de::Error::custom)
224    }
225}
226
227impl<T> Problem for ClosestVectorProblem<T>
228where
229    T: ClosestVectorTarget + Serialize + for<'de> Deserialize<'de>,
230{
231    const NAME: &'static str = "ClosestVectorProblem";
232    type Solution = Vec<i64>;
233    type Value = Min<f64>;
234
235    crate::problem_parameters![
236        ("ambient_dimension", ambient_dimension),
237        ("num_basis_vectors", num_basis_vectors),
238    ];
239
240    fn evaluate(&self, solution: &Self::Solution) -> Result<Min<f64>, EvaluationError> {
241        if solution.len() != self.num_basis_vectors() {
242            return Err(EvaluationError::InvalidConfiguration(format!(
243                "expected {} closest-vector coefficients, got {}",
244                self.num_basis_vectors(),
245                solution.len()
246            )));
247        }
248
249        let mut displacement = self
250            .target
251            .iter()
252            .map(ClosestVectorTarget::to_f64)
253            .collect::<Result<Vec<_>, _>>()?;
254        for value in &mut displacement {
255            *value = -*value;
256        }
257
258        for (&coefficient, column) in solution.iter().zip(&self.basis) {
259            let coefficient = crate::types::i64_to_exact_f64(coefficient)
260                .map_err(|error| EvaluationError::InexactFloatConversion(error.to_string()))?;
261            for (value, &basis_entry) in displacement.iter_mut().zip(column) {
262                let basis_entry = crate::types::i64_to_exact_f64(basis_entry)
263                    .map_err(|error| EvaluationError::InexactFloatConversion(error.to_string()))?;
264                let next = *value + coefficient * basis_entry;
265                if !next.is_finite() {
266                    return Err(EvaluationError::NonFiniteResult(
267                        "computing closest-vector displacement".into(),
268                    ));
269                }
270                *value = next;
271            }
272        }
273
274        let squared_norm = displacement.into_iter().try_fold(0.0, |total, value| {
275            let next = total + value * value;
276            if next.is_finite() {
277                Ok(next)
278            } else {
279                Err(EvaluationError::NonFiniteResult(
280                    "computing closest-vector norm".into(),
281                ))
282            }
283        })?;
284        Ok(Min(Some(squared_norm.sqrt())))
285    }
286
287    fn variant() -> Vec<(&'static str, &'static str)> {
288        vec![("target", T::NAME)]
289    }
290}
291
292crate::declare_variants! {
293    default ClosestVectorProblem<i64> => "2^(num_basis_vectors * log(num_basis_vectors))" create ClosestVectorProblemI64CreateSpec,
294    ClosestVectorProblem<f64> => "2^(num_basis_vectors * log(num_basis_vectors))" create ClosestVectorProblemF64CreateSpec,
295}
296
297#[cfg(feature = "example-db")]
298pub(crate) fn canonical_model_example_specs() -> Vec<crate::example_db::specs::ModelExampleSpec> {
299    vec![crate::example_db::specs::ModelExampleSpec {
300        id: "closest_vector_problem",
301        instance: Box::new(
302            ClosestVectorProblem::new(vec![vec![2, 0], vec![1, 2]], vec![3_i64, 2])
303                .expect("canonical closest-vector instance must be valid"),
304        ),
305        optimal_config: serde_json::json!(vec![1, 1]),
306        optimal_value: serde_json::json!(0.0),
307    }]
308}
309
310#[cfg(test)]
311#[path = "../../unit_tests/models/algebraic/closest_vector_problem.rs"]
312mod tests;