use scirs2_core::ndarray::{Array1, Array2, ArrayView1, ArrayView2, Axis};
use sklears_core::{
error::{Result as SklResult, SklearsError},
traits::{Estimator, Fit, Transform, Untrained},
types::Float,
};
#[derive(Debug, Clone)]
pub struct RiemannianManifold<S = Untrained> {
n_neighbors: usize,
metric_type: String,
geodesic_method: String,
parallel_transport_method: String,
curvature_estimation_radius: Float,
state: S,
}
#[derive(Debug, Clone)]
pub struct RiemannianManifoldTrained {
pub metric_tensor: Array2<Float>,
pub geodesic_distances: Array2<Float>,
pub christoffel_symbols: Array2<Float>,
pub gaussian_curvature: Array1<Float>,
pub mean_curvature: Array1<Float>,
pub data_points: Array2<Float>,
pub n_neighbors: usize,
}
impl Default for RiemannianManifold<Untrained> {
fn default() -> Self {
Self::new()
}
}
impl RiemannianManifold<Untrained> {
pub fn builder() -> RiemannianManifoldBuilder {
RiemannianManifoldBuilder::new()
}
pub fn new() -> Self {
Self {
n_neighbors: 5,
metric_type: "euclidean".to_string(),
geodesic_method: "dijkstra".to_string(),
parallel_transport_method: "schild_ladder".to_string(),
curvature_estimation_radius: 1.0,
state: Untrained,
}
}
}
#[derive(Debug)]
pub struct RiemannianManifoldBuilder {
n_neighbors: usize,
metric_type: String,
geodesic_method: String,
parallel_transport_method: String,
curvature_estimation_radius: Float,
}
impl RiemannianManifoldBuilder {
fn new() -> Self {
Self {
n_neighbors: 5,
metric_type: "euclidean".to_string(),
geodesic_method: "dijkstra".to_string(),
parallel_transport_method: "schild_ladder".to_string(),
curvature_estimation_radius: 1.0,
}
}
pub fn n_neighbors(mut self, n_neighbors: usize) -> Self {
self.n_neighbors = n_neighbors;
self
}
pub fn metric_type(mut self, metric_type: &str) -> Self {
self.metric_type = metric_type.to_string();
self
}
pub fn geodesic_method(mut self, method: &str) -> Self {
self.geodesic_method = method.to_string();
self
}
pub fn parallel_transport_method(mut self, method: &str) -> Self {
self.parallel_transport_method = method.to_string();
self
}
pub fn curvature_estimation_radius(mut self, radius: Float) -> Self {
self.curvature_estimation_radius = radius;
self
}
pub fn build(self) -> RiemannianManifold<Untrained> {
RiemannianManifold {
n_neighbors: self.n_neighbors,
metric_type: self.metric_type,
geodesic_method: self.geodesic_method,
parallel_transport_method: self.parallel_transport_method,
curvature_estimation_radius: self.curvature_estimation_radius,
state: Untrained,
}
}
}
impl Estimator for RiemannianManifold<Untrained> {
type Config = ();
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&()
}
}
impl Fit<ArrayView2<'_, Float>, ()> for RiemannianManifold<Untrained> {
type Fitted = RiemannianManifold<RiemannianManifoldTrained>;
fn fit(self, x: &ArrayView2<'_, Float>, _y: &()) -> SklResult<Self::Fitted> {
let (n_samples, _n_features) = x.dim();
if n_samples < self.n_neighbors {
return Err(SklearsError::InvalidInput(format!(
"Number of samples ({}) must be >= n_neighbors ({})",
n_samples, self.n_neighbors
)));
}
let metric_tensor = compute_metric_tensor(x, self.n_neighbors)?;
let geodesic_distances = match self.geodesic_method.as_str() {
"dijkstra" => compute_geodesic_distances_dijkstra(x, self.n_neighbors)?,
"floyd_warshall" => compute_geodesic_distances_floyd_warshall(x, self.n_neighbors)?,
_ => {
return Err(SklearsError::InvalidInput(format!(
"Unknown geodesic method: {}",
self.geodesic_method
)))
}
};
let christoffel_symbols = compute_christoffel_symbols(&metric_tensor)?;
let (gaussian_curvature, mean_curvature) =
estimate_curvature(x, self.n_neighbors, self.curvature_estimation_radius)?;
let trained_state = RiemannianManifoldTrained {
metric_tensor,
geodesic_distances,
christoffel_symbols,
gaussian_curvature,
mean_curvature,
data_points: x.to_owned(),
n_neighbors: self.n_neighbors,
};
Ok(RiemannianManifold {
n_neighbors: self.n_neighbors,
metric_type: self.metric_type,
geodesic_method: self.geodesic_method,
parallel_transport_method: self.parallel_transport_method,
curvature_estimation_radius: self.curvature_estimation_radius,
state: trained_state,
})
}
}
impl Transform<ArrayView2<'_, Float>, Array2<Float>>
for RiemannianManifold<RiemannianManifoldTrained>
{
fn transform(&self, x: &ArrayView2<'_, Float>) -> SklResult<Array2<Float>> {
let n_new = x.nrows();
let n_train = self.state.data_points.nrows();
let mut result = Array2::zeros((n_new, n_train));
for i in 0..n_new {
for j in 0..n_train {
let dist = compute_geodesic_distance_between_points(
&x.row(i),
&self.state.data_points.row(j),
&self.state.metric_tensor,
)?;
result[(i, j)] = dist;
}
}
Ok(result)
}
}
impl RiemannianManifold<RiemannianManifoldTrained> {
pub fn parallel_transport(
&self,
vector: &ArrayView1<Float>,
start_point: &ArrayView1<Float>,
end_point: &ArrayView1<Float>,
) -> SklResult<Array1<Float>> {
match self.parallel_transport_method.as_str() {
"schild_ladder" => parallel_transport_schild_ladder(
vector,
start_point,
end_point,
&self.state.christoffel_symbols,
),
"pole_ladder" => parallel_transport_pole_ladder(
vector,
start_point,
end_point,
&self.state.christoffel_symbols,
),
_ => Err(SklearsError::InvalidInput(format!(
"Unknown parallel transport method: {}",
self.parallel_transport_method
))),
}
}
pub fn exponential_map(
&self,
base_point: &ArrayView1<Float>,
tangent_vector: &ArrayView1<Float>,
) -> SklResult<Array1<Float>> {
exponential_map(base_point, tangent_vector, &self.state.christoffel_symbols)
}
pub fn logarithmic_map(
&self,
base_point: &ArrayView1<Float>,
target_point: &ArrayView1<Float>,
) -> SklResult<Array1<Float>> {
logarithmic_map(base_point, target_point, &self.state.christoffel_symbols)
}
pub fn gaussian_curvature(&self) -> &Array1<Float> {
&self.state.gaussian_curvature
}
pub fn mean_curvature(&self) -> &Array1<Float> {
&self.state.mean_curvature
}
pub fn metric_tensor(&self) -> &Array2<Float> {
&self.state.metric_tensor
}
pub fn geodesic_distances(&self) -> &Array2<Float> {
&self.state.geodesic_distances
}
}
fn compute_metric_tensor(
x: &ArrayView2<'_, Float>,
n_neighbors: usize,
) -> SklResult<Array2<Float>> {
let (n_samples, n_features) = x.dim();
let mut metric = Array2::zeros((n_features, n_features));
for i in 0..n_samples {
let point = x.row(i);
let mut distances: Vec<(Float, usize)> = Vec::new();
for j in 0..n_samples {
if i != j {
let dist = (&x.row(j) - &point).mapv(|x| x * x).sum().sqrt();
distances.push((dist, j));
}
}
distances.sort_by(|a, b| a.0.partial_cmp(&b.0).expect("operation should succeed"));
let neighbors: Vec<usize> = distances
.iter()
.take(n_neighbors.min(n_samples - 1))
.map(|(_, idx)| *idx)
.collect();
let mut local_data = Array2::zeros((neighbors.len(), n_features));
for (k, &neighbor_idx) in neighbors.iter().enumerate() {
let diff = &x.row(neighbor_idx) - &point;
local_data.row_mut(k).assign(&diff);
}
let mean = local_data
.mean_axis(Axis(0))
.expect("operation should succeed");
let centered_data = &local_data - &mean;
let cov = centered_data.t().dot(¢ered_data) / (neighbors.len() as Float - 1.0);
metric = metric + cov;
}
metric /= n_samples as Float;
for i in 0..n_features {
metric[(i, i)] += 1e-6;
}
Ok(metric)
}
fn compute_geodesic_distances_dijkstra(
x: &ArrayView2<'_, Float>,
n_neighbors: usize,
) -> SklResult<Array2<Float>> {
let n_samples = x.nrows();
let mut distances = Array2::from_elem((n_samples, n_samples), Float::INFINITY);
for i in 0..n_samples {
distances[(i, i)] = 0.0;
}
for i in 0..n_samples {
let point = x.row(i);
let mut neighbor_distances: Vec<(Float, usize)> = Vec::new();
for j in 0..n_samples {
if i != j {
let dist = (&x.row(j) - &point).mapv(|x| x * x).sum().sqrt();
neighbor_distances.push((dist, j));
}
}
neighbor_distances.sort_by(|a, b| a.0.partial_cmp(&b.0).expect("operation should succeed"));
for &(dist, j) in neighbor_distances
.iter()
.take(n_neighbors.min(n_samples - 1))
{
distances[(i, j)] = dist;
distances[(j, i)] = dist; }
}
for start in 0..n_samples {
let mut dist = vec![Float::INFINITY; n_samples];
let mut visited = vec![false; n_samples];
dist[start] = 0.0;
for _ in 0..n_samples {
let mut min_dist = Float::INFINITY;
let mut min_vertex = 0;
for v in 0..n_samples {
if !visited[v] && dist[v] < min_dist {
min_dist = dist[v];
min_vertex = v;
}
}
visited[min_vertex] = true;
for neighbor in 0..n_samples {
if !visited[neighbor] && distances[(min_vertex, neighbor)] != Float::INFINITY {
let new_dist = dist[min_vertex] + distances[(min_vertex, neighbor)];
if new_dist < dist[neighbor] {
dist[neighbor] = new_dist;
}
}
}
}
for end in 0..n_samples {
distances[(start, end)] = dist[end];
}
}
Ok(distances)
}
fn compute_geodesic_distances_floyd_warshall(
x: &ArrayView2<'_, Float>,
n_neighbors: usize,
) -> SklResult<Array2<Float>> {
let n_samples = x.nrows();
let mut distances = Array2::from_elem((n_samples, n_samples), Float::INFINITY);
for i in 0..n_samples {
distances[(i, i)] = 0.0;
}
for i in 0..n_samples {
let point = x.row(i);
let mut neighbor_distances: Vec<(Float, usize)> = Vec::new();
for j in 0..n_samples {
if i != j {
let dist = (&x.row(j) - &point).mapv(|x| x * x).sum().sqrt();
neighbor_distances.push((dist, j));
}
}
neighbor_distances.sort_by(|a, b| a.0.partial_cmp(&b.0).expect("operation should succeed"));
for &(dist, j) in neighbor_distances
.iter()
.take(n_neighbors.min(n_samples - 1))
{
distances[(i, j)] = dist;
distances[(j, i)] = dist; }
}
for k in 0..n_samples {
for i in 0..n_samples {
for j in 0..n_samples {
if distances[(i, k)] != Float::INFINITY && distances[(k, j)] != Float::INFINITY {
let new_dist = distances[(i, k)] + distances[(k, j)];
if new_dist < distances[(i, j)] {
distances[(i, j)] = new_dist;
}
}
}
}
}
Ok(distances)
}
fn compute_christoffel_symbols(metric: &Array2<Float>) -> SklResult<Array2<Float>> {
let n = metric.nrows();
let christoffel = Array2::zeros((n * n, n));
Ok(christoffel)
}
fn estimate_curvature(
x: &ArrayView2<'_, Float>,
n_neighbors: usize,
radius: Float,
) -> SklResult<(Array1<Float>, Array1<Float>)> {
let n_samples = x.nrows();
let mut gaussian_curvature = Array1::zeros(n_samples);
let mut mean_curvature = Array1::zeros(n_samples);
for i in 0..n_samples {
let point = x.row(i);
let mut neighbors = Vec::new();
for j in 0..n_samples {
if i != j {
let dist = (&x.row(j) - &point).mapv(|x| x * x).sum().sqrt();
if dist <= radius {
neighbors.push(j);
}
}
}
if neighbors.len() >= n_neighbors {
let local_variance = estimate_local_variance(x, i, &neighbors);
gaussian_curvature[i] = 1.0 / (1.0 + local_variance);
mean_curvature[i] = local_variance.sqrt();
}
}
Ok((gaussian_curvature, mean_curvature))
}
fn estimate_local_variance(
x: &ArrayView2<'_, Float>,
center_idx: usize,
neighbors: &[usize],
) -> Float {
let center = x.row(center_idx);
let mut variance = 0.0;
for &neighbor_idx in neighbors {
let diff = &x.row(neighbor_idx) - ¢er;
variance += diff.mapv(|x| x * x).sum();
}
if !neighbors.is_empty() {
variance / neighbors.len() as Float
} else {
0.0
}
}
fn compute_geodesic_distance_between_points(
point1: &ArrayView1<Float>,
point2: &ArrayView1<Float>,
metric: &Array2<Float>,
) -> SklResult<Float> {
let diff = point2 - point1;
let dist_squared = diff.dot(&metric.dot(&diff));
Ok(dist_squared.sqrt())
}
fn parallel_transport_schild_ladder(
vector: &ArrayView1<Float>,
_start_point: &ArrayView1<Float>,
_end_point: &ArrayView1<Float>,
_christoffel: &Array2<Float>,
) -> SklResult<Array1<Float>> {
Ok(vector.to_owned())
}
fn parallel_transport_pole_ladder(
vector: &ArrayView1<Float>,
_start_point: &ArrayView1<Float>,
_end_point: &ArrayView1<Float>,
_christoffel: &Array2<Float>,
) -> SklResult<Array1<Float>> {
Ok(vector.to_owned())
}
fn exponential_map(
base_point: &ArrayView1<Float>,
tangent_vector: &ArrayView1<Float>,
_christoffel: &Array2<Float>,
) -> SklResult<Array1<Float>> {
Ok(base_point + tangent_vector)
}
fn logarithmic_map(
base_point: &ArrayView1<Float>,
target_point: &ArrayView1<Float>,
_christoffel: &Array2<Float>,
) -> SklResult<Array1<Float>> {
Ok(target_point - base_point)
}