use crate::{gradcheck, Backend, Config, Report};
#[derive(Debug, Clone)]
pub struct SweepReport {
pub reports: Vec<Report>,
}
impl SweepReport {
pub fn failing(&self) -> Vec<&Report> {
self.reports.iter().filter(|r| !r.passed()).collect()
}
pub fn all_passed(&self) -> bool {
self.reports.iter().all(|r| r.passed())
}
pub fn len(&self) -> usize {
self.reports.len()
}
pub fn is_empty(&self) -> bool {
self.reports.is_empty()
}
pub fn assert_all_pass(&self) {
let bad = self.failing();
if !bad.is_empty() {
let detail = bad
.iter()
.map(|r| r.to_string())
.collect::<Vec<_>>()
.join("\n ");
panic!(
"gradcheck sweep: {} of {} shapes mismatched:\n {detail}",
bad.len(),
self.reports.len()
);
}
}
}
pub fn shape_sweep<B, F, D>(
name: &str,
shapes: &[Vec<usize>],
data_for: D,
f: F,
cfg: &Config,
) -> SweepReport
where
B: Backend,
F: Fn(B::Tensor) -> B::Tensor + Copy,
D: Fn(usize) -> Vec<f64>,
{
let mut reports = Vec::with_capacity(shapes.len());
for shape in shapes {
let n: usize = shape.iter().product();
let data = data_for(n);
let label = format!("{name}{shape:?}");
reports.push(gradcheck::<B, _>(&label, &data, shape, f, cfg));
}
SweepReport { reports }
}