use plotters::prelude::*;
use plotters_statistical::colormap::{GradientColorMap, Normalization};
use plotters_statistical::series::heatmap::HeatmapAnnotation;
use plotters_statistical::Heatmap;
fn main() -> Result<(), Box<dyn std::error::Error>> {
let root = SVGBackend::new("heatmap.svg", (560, 520)).into_drawing_area();
root.fill(&WHITE)?;
let counts = vec![
vec![52.0, 3.0, 1.0],
vec![4.0, 43.0, 6.0],
vec![0.0, 5.0, 61.0],
];
let mut chart = ChartBuilder::on(&root)
.caption("Confusion matrix", ("sans-serif", 22))
.margin(20)
.set_label_area_size(LabelAreaPosition::Left, 40)
.set_label_area_size(LabelAreaPosition::Bottom, 40)
.build_cartesian_2d(-0.5f64..2.5f64, -0.5f64..2.5f64)?;
chart
.configure_mesh()
.disable_mesh()
.x_desc("predicted")
.y_desc("true")
.draw()?;
chart.draw_series(std::iter::once(
Heatmap::new(&counts)
.colormap(GradientColorMap::blues())
.normalization(Normalization::Linear {
min: 0.0,
max: 61.0,
})
.annotate(HeatmapAnnotation {
precision: 0,
font_size: 16,
text_color: None,
})
.cell_gap(1),
))?;
root.present()?;
println!("wrote heatmap.svg");
Ok(())
}