use indexmap::IndexSet;
use ndarray::{Array2, ArrayView1};
use crate::embed::tools::degrees::*;
use crate::embed::tools::edge::{IN, OUT};
use crate::io::embeddedbson::EmbeddedBsonReload;
type Distance<F> = fn(&[F], &[F]) -> f64;
#[derive(Debug)]
pub enum EmbeddingMode {
Hope,
NodeSketch,
}
pub const TAG_OUT: u8 = 0;
pub const TAG_IN: u8 = 1;
pub const TAG_IN_OUT: u8 = 1;
pub trait EmbeddedT<F> {
fn is_symetric(&self) -> bool;
fn get_dimension(&self) -> usize;
fn get_noderank_distance(&self, node_rank1: usize, node_rank2: usize) -> f64;
fn get_vec_distance(&self, from: &[F], to: &[F]) -> f64;
fn get_nb_nodes(&self) -> usize;
fn get_embedded_node(&self, node_rank: usize, _tag: u8) -> ArrayView1<F>;
fn get_distance(&self) -> fn(&[F], &[F]) -> f64;
}
pub struct Embedded<F> {
data: Array2<F>,
distance: fn(&[F], &[F]) -> f64,
}
impl<F> Embedded<F> {
pub(crate) fn new(arr: Array2<F>, distance: Distance<F>) -> Self {
Embedded {
data: arr,
distance: distance,
}
}
pub fn get_embedded(&self) -> &Array2<F> {
&self.data
}
pub fn get_distance_ref(&self) -> &fn(&[F], &[F]) -> f64 {
&self.distance
}
pub fn get_embedded_node(&self, node_rank: usize, _tag: u8) -> ArrayView1<F> {
self.data.row(node_rank)
}
}
impl<F> EmbeddedT<F> for Embedded<F> {
fn is_symetric(&self) -> bool {
true
}
fn get_dimension(&self) -> usize {
self.data.dim().1
}
fn get_vec_distance(&self, data1: &[F], data2: &[F]) -> f64 {
assert_eq!(data1.len(), self.get_dimension());
(self.distance)(data1, data2)
}
fn get_noderank_distance(&self, node1: usize, node2: usize) -> f64 {
(self.distance)(
self.data.row(node1).as_slice().unwrap(),
self.data.row(node2).as_slice().unwrap(),
)
}
fn get_nb_nodes(&self) -> usize {
self.data.dim().0
}
fn get_embedded_node(&self, node_rank: usize, _tag: u8) -> ArrayView1<F> {
self.data.row(node_rank)
}
fn get_distance(&self) -> fn(&[F], &[F]) -> f64 {
self.distance
}
}
pub struct EmbeddedAsym<F> {
source: Array2<F>,
target: Array2<F>,
degrees: Option<Vec<Degree>>,
distance: Distance<F>,
}
impl<F> EmbeddedAsym<F> {
pub(crate) fn new(
source: Array2<F>,
target: Array2<F>,
degrees: Option<Vec<Degree>>,
distance: Distance<F>,
) -> Self {
assert_eq!(source.dim().0, target.dim().0);
assert_eq!(source.dim().1, target.dim().1);
EmbeddedAsym {
source,
target,
degrees,
distance,
}
}
pub fn get_embedded_source(&self) -> &Array2<F> {
&self.source
}
pub fn get_embedded_target(&self) -> &Array2<F> {
&self.target
}
}
impl<F> EmbeddedT<F> for EmbeddedAsym<F> {
fn is_symetric(&self) -> bool {
false
}
fn get_dimension(&self) -> usize {
self.source.dim().1
}
fn get_vec_distance(&self, data1: &[F], data2: &[F]) -> f64 {
(self.distance)(data1, data2)
}
fn get_noderank_distance(&self, node_rank1: usize, node_rank2: usize) -> f64 {
let mut distances = Vec::<f64>::with_capacity(3);
let dist_s = (self.distance)(
self.source.row(node_rank1).as_slice().unwrap(),
self.source.row(node_rank2).as_slice().unwrap(),
);
distances.push(dist_s);
let dist_t = (self.distance)(
self.target.row(node_rank1).as_slice().unwrap(),
self.target.row(node_rank2).as_slice().unwrap(),
);
distances.push(dist_t);
let dist_t = (self.distance)(
self.source.row(node_rank1).as_slice().unwrap(),
self.target.row(node_rank2).as_slice().unwrap(),
);
distances.push(dist_t);
if !distances.is_empty() {
distances.iter().sum::<f64>() / distances.len() as f64
} else {
match &self.degrees {
Some(degrees) => {
log::error!(
"cannot get distance between node rank1 : {} degree = {:?}, node rank2 : {}, degree = {:?}",
node_rank1,
degrees[node_rank1],
node_rank2,
degrees[node_rank2]
);
}
None => {
log::error!(
"cannot get distance between node rank1 : {} , node rank2 : {}",
node_rank1,
node_rank2
);
}
}
log::error!("get_noderank_distance asymetric no distance computed");
1.
}
}
fn get_nb_nodes(&self) -> usize {
self.source.dim().0
}
fn get_embedded_node(&self, node_rank: usize, tag: u8) -> ArrayView1<F> {
match tag {
OUT => {
self.source.row(node_rank)
}
IN => {
self.target.row(node_rank)
}
_ => {
log::error!(" for asymetric embedding tag in get_embedded_node must be 0 or 1");
std::panic!("for asymetric embedding tag in get_embedded_node must be 0 or 1");
}
}
}
fn get_distance(&self) -> fn(&[F], &[F]) -> f64 {
self.distance
}
}
pub trait EmbedderT<F> {
type Output: EmbeddedT<F>;
fn embed(&mut self) -> Result<Self::Output, anyhow::Error>;
}
pub struct Embedding<F, NodeId: std::hash::Hash + std::cmp::Eq, EmbeddedData: EmbeddedT<F>> {
nodeindexation: IndexSet<NodeId>,
embedded: EmbeddedData,
mark: std::marker::PhantomData<F>,
}
impl<NodeId, EmbeddedData, F> Embedding<F, NodeId, EmbeddedData>
where
EmbeddedData: EmbeddedT<F>,
NodeId: std::hash::Hash + std::cmp::Eq,
{
pub fn new(
nodeindexation: IndexSet<NodeId>,
embedder: &mut dyn EmbedderT<F, Output = EmbeddedData>,
) -> Result<Self, anyhow::Error> {
let embedded_res = embedder.embed();
if embedded_res.is_err() {
log::error!("embedding failed");
Err(embedded_res.err().unwrap())
} else {
Ok(Embedding {
nodeindexation,
embedded: embedded_res.unwrap(),
mark: std::marker::PhantomData,
})
}
}
pub fn get_node_indexation(&self) -> &IndexSet<NodeId> {
&self.nodeindexation
}
pub fn get_embedded_data(&self) -> &EmbeddedData {
&self.embedded
}
pub fn get_node_distance(&self, node1: NodeId, node2: NodeId) -> f64 {
let rank1 = self.nodeindexation.get_index_of(&node1).unwrap();
let rank2 = self.nodeindexation.get_index_of(&node2).unwrap();
self.embedded.get_noderank_distance(rank1, rank2)
}
pub fn get_node_rank(&self, node_id: NodeId) -> Option<usize> {
self.nodeindexation.get_index_of(&node_id)
}
pub fn get_node_id(&self, rank: usize) -> Option<&NodeId> {
self.nodeindexation.get_index(rank)
}
}
pub fn from_bson_with_jaccard<F, NodeId>(
bson_reload: EmbeddedBsonReload<F, NodeId>,
) -> Result<Embedding<F, NodeId, Embedded<F>>, anyhow::Error>
where
F: Eq,
NodeId: std::hash::Hash + std::cmp::Eq,
{
let embedded_data = Embedded::new(
bson_reload.out_embedded,
crate::embed::tools::jaccard::jaccard_distance::<F>,
);
if bson_reload.node_indexation.is_none() {
return Err(anyhow::anyhow!("no node indexation in bson dump"));
}
let embedding = Embedding::<F, NodeId, Embedded<F>> {
nodeindexation: bson_reload.node_indexation.unwrap(),
embedded: embedded_data,
mark: std::marker::PhantomData,
};
Ok(embedding)
}