use super::types::{DivergenceReport, KernelTrace};
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct BrickProfiler {
pub run_id: String,
pub traces: Vec<KernelTrace>,
pub total_time_us: f64,
pub diverged: bool,
pub divergence_diagnosis: String,
}
impl BrickProfiler {
pub fn new(run_id: &str) -> Self {
Self {
run_id: run_id.to_string(),
traces: Vec::new(),
total_time_us: 0.0,
diverged: false,
divergence_diagnosis: String::new(),
}
}
pub fn add_trace(&mut self, trace: KernelTrace) {
self.total_time_us += trace.time_us;
self.traces.push(trace);
}
pub fn is_diverged(&self) -> bool {
self.diverged
}
pub fn compare(&self, reference: &BrickProfiler) -> DivergenceReport {
let ref_index: std::collections::HashMap<(&str, usize, u32), &KernelTrace> = reference
.traces
.iter()
.map(|t| ((t.kernel_name.as_str(), t.layer_idx, t.position), t))
.collect();
let mut kernels_compared = 0;
for actual_trace in &self.traces {
let key = (
actual_trace.kernel_name.as_str(),
actual_trace.layer_idx,
actual_trace.position,
);
if let Some(expected_trace) = ref_index.get(&key) {
kernels_compared += 1;
if actual_trace.output_checksum != expected_trace.output_checksum {
return DivergenceReport::diverged(
(*expected_trace).clone(),
actual_trace.clone(),
kernels_compared,
);
}
}
}
DivergenceReport::matched(kernels_compared)
}
pub fn compare_and_mark(&mut self, reference: &BrickProfiler) -> DivergenceReport {
let report = self.compare(reference);
self.diverged = !report.matched;
self.divergence_diagnosis = report.diagnosis.clone();
report
}
pub fn traces_for_kernel(&self, kernel_name: &str) -> Vec<&KernelTrace> {
self.traces
.iter()
.filter(|t| t.kernel_name == kernel_name)
.collect()
}
pub fn traces_for_layer(&self, layer_idx: usize) -> Vec<&KernelTrace> {
self.traces
.iter()
.filter(|t| t.layer_idx == layer_idx)
.collect()
}
pub fn clear(&mut self) {
self.traces.clear();
self.total_time_us = 0.0;
self.diverged = false;
self.divergence_diagnosis.clear();
}
pub fn to_json(&self) -> Result<String, serde_json::Error> {
serde_json::to_string_pretty(self)
}
pub fn from_json(json: &str) -> Result<Self, serde_json::Error> {
serde_json::from_str(json)
}
}