use ndarray::{Array1, Array3};
use crate::canonical::{KnnClassifier, KnnRegressor, KnnWeight};
use crate::graph_builder::{
assemble_model, int_attribute, ints_attribute, make_i64_tensor, make_node, make_tensor,
make_typed_value_info, make_value_info, Dimension, INT64,
};
use crate::proto::{GraphProto, ModelProto, NodeProto, TensorProto};
use crate::{IR_VERSION, OPSET_VERSION};
fn neighbor_graph(samples: &ndarray::Array2<f64>, k: usize) -> (Vec<NodeProto>, Vec<TensorProto>) {
let sample_count = samples.nrows();
let features = samples.ncols();
let samples = Array3::from_shape_fn((1, sample_count, features), |(_, row, column)| {
samples[(row, column)] as f32
});
(
vec![
make_node("Unsqueeze", ["X", "axis_one"], ["expanded"], Vec::new()),
make_node("Sub", ["expanded", "samples"], ["delta"], Vec::new()),
make_node("Mul", ["delta", "delta"], ["squared"], Vec::new()),
make_node(
"ReduceSum",
["squared", "axis_features"],
["squared_distances"],
vec![int_attribute("keepdims", 0)],
),
make_node("Sqrt", ["squared_distances"], ["distances"], Vec::new()),
make_node("Neg", ["distances"], ["negative_distances"], Vec::new()),
make_node(
"TopK",
["negative_distances", "k"],
["negative_neighbor_distances", "indices"],
vec![
int_attribute("axis", 1),
int_attribute("largest", 1),
int_attribute("sorted", 1),
],
),
make_node(
"Neg",
["negative_neighbor_distances"],
["neighbor_distances"],
Vec::new(),
),
],
vec![
make_tensor("samples", &samples.into_dyn()),
make_i64_tensor("axis_one", &[1], vec![1]),
make_i64_tensor("axis_features", &[1], vec![2]),
make_i64_tensor("k", &[1], vec![k as i64]),
],
)
}
fn weight_nodes(weight: KnnWeight) -> (Vec<NodeProto>, Vec<TensorProto>, &'static str) {
if weight == KnnWeight::Uniform {
return (Vec::new(), Vec::new(), "");
}
let zero = Array1::from_vec(vec![0.0_f32]);
let one = Array1::from_vec(vec![1.0_f32]);
(
vec![
make_node(
"Equal",
["neighbor_distances", "zero"],
["zero_mask"],
Vec::new(),
),
make_node(
"Cast",
["zero_mask"],
["zero_weights"],
vec![int_attribute("to", 1)],
),
make_node(
"ReduceSum",
["zero_weights", "axis_neighbors"],
["zero_count"],
vec![int_attribute("keepdims", 1)],
),
make_node("Greater", ["zero_count", "zero"], ["has_zero"], Vec::new()),
make_node(
"Where",
["zero_mask", "one", "neighbor_distances"],
["safe_distances"],
Vec::new(),
),
make_node(
"Div",
["one", "safe_distances"],
["inverse_weights"],
Vec::new(),
),
make_node(
"Where",
["has_zero", "zero_weights", "inverse_weights"],
["weights"],
Vec::new(),
),
],
vec![
make_tensor("zero", &zero.into_dyn()),
make_tensor("one", &one.into_dyn()),
make_i64_tensor("axis_neighbors", &[1], vec![1]),
],
"weights",
)
}
#[must_use]
pub fn export_knn_regressor(model: &KnnRegressor) -> ModelProto {
let (mut nodes, mut initializer) = neighbor_graph(&model.samples, model.k);
let targets = model.targets.mapv(|value| value as f32);
initializer.push(make_tensor("targets", &targets.into_dyn()));
nodes.push(make_node(
"Gather",
["targets", "indices"],
["neighbor_targets"],
vec![int_attribute("axis", 0)],
));
if model.weight == KnnWeight::Uniform {
nodes.push(make_node(
"ReduceMean",
["neighbor_targets"],
["prediction"],
vec![
ints_attribute("axes", vec![1]),
int_attribute("keepdims", 1),
],
));
} else {
let (weighting, extra, _) = weight_nodes(model.weight);
nodes.extend(weighting);
initializer.extend(extra);
nodes.extend([
make_node(
"Mul",
["neighbor_targets", "weights"],
["weighted_targets"],
Vec::new(),
),
make_node(
"ReduceSum",
["weighted_targets", "axis_neighbors"],
["target_sum"],
vec![int_attribute("keepdims", 1)],
),
make_node(
"ReduceSum",
["weights", "axis_neighbors"],
["weight_sum"],
vec![int_attribute("keepdims", 1)],
),
make_node(
"Div",
["target_sum", "weight_sum"],
["prediction"],
Vec::new(),
),
]);
}
assemble_knn(
model.samples.ncols(),
nodes,
initializer,
make_value_info(
"prediction",
&[Dimension::Symbolic("batch".into()), Dimension::Fixed(1)],
),
)
}
#[must_use]
pub fn export_knn_classifier(model: &KnnClassifier) -> ModelProto {
let (mut nodes, mut initializer) = neighbor_graph(&model.samples, model.k);
let classes = model.class_labels.len();
initializer.extend([
make_i64_tensor(
"targets",
&[model.target_indices.len()],
model.target_indices.clone(),
),
make_i64_tensor("labels", &[classes], model.class_labels.clone()),
make_i64_tensor("depth", &[], vec![classes as i64]),
make_tensor(
"one_hot_values",
&Array1::from_vec(vec![0.0_f32, 1.0]).into_dyn(),
),
make_i64_tensor("axis_neighbors", &[1], vec![1]),
]);
nodes.extend([
make_node(
"Gather",
["targets", "indices"],
["neighbor_classes"],
vec![int_attribute("axis", 0)],
),
make_node(
"OneHot",
["neighbor_classes", "depth", "one_hot_values"],
["votes"],
vec![int_attribute("axis", -1)],
),
]);
if model.weight == KnnWeight::Distance {
let (weighting, extra, _) = weight_nodes(model.weight);
nodes.extend(weighting);
initializer.extend(extra);
nodes.extend([
make_node(
"Unsqueeze",
["weights", "axis_two"],
["expanded_weights"],
Vec::new(),
),
make_node(
"Mul",
["votes", "expanded_weights"],
["weighted_votes"],
Vec::new(),
),
make_node(
"ReduceSum",
["weighted_votes", "axis_neighbors"],
["class_scores"],
vec![int_attribute("keepdims", 0)],
),
]);
initializer.push(make_i64_tensor("axis_two", &[1], vec![2]));
} else {
nodes.push(make_node(
"ReduceSum",
["votes", "axis_neighbors"],
["class_scores"],
vec![int_attribute("keepdims", 0)],
));
}
nodes.extend([
make_node(
"ArgMax",
["class_scores"],
["class_index"],
vec![int_attribute("axis", 1), int_attribute("keepdims", 0)],
),
make_node(
"Gather",
["labels", "class_index"],
["label"],
vec![int_attribute("axis", 0)],
),
]);
assemble_knn(
model.samples.ncols(),
nodes,
initializer,
make_typed_value_info("label", &[Dimension::Symbolic("batch".into())], INT64),
)
}
fn assemble_knn(
features: usize,
nodes: Vec<NodeProto>,
initializer: Vec<TensorProto>,
output: crate::proto::ValueInfoProto,
) -> ModelProto {
assemble_model(
GraphProto {
node: nodes,
name: "k_nearest_neighbors".into(),
initializer,
doc_string: String::new(),
input: vec![make_value_info(
"X",
&[
Dimension::Symbolic("batch".into()),
Dimension::Fixed(features),
],
)],
output: vec![output],
value_info: Vec::new(),
},
OPSET_VERSION,
IR_VERSION,
)
}