1use std::path::Path;
5use std::sync::{Arc, Mutex};
6
7use serde::{Deserialize, Serialize};
8use silicera::hardware::HardwareInfo;
9use silicera::hnep::{HnepProfile, SizeClassEntry};
10use silicera::measure::{MeasurementConfig, MeasurementSummary};
11use silicera::retrain::{parse_only_filter, target_matches_filter};
12use silicera::specialize::{DecisionTree, SpecializeConfig};
13use silicera::tournament::{Tournament, TournamentConfig, TournamentResult};
14use silicera::variant::{Variant, VariantId};
15use silicera::Result;
16
17use crate::bench::{
18 BranchBench, CacheTarget, ConcurrencyBench, FloatBench, IntegerBench, MemOpBench, MemoryBench,
19};
20
21#[derive(Debug, Clone, Serialize, Deserialize)]
23pub struct ProfileGenConfig {
24 pub measurement: MeasurementConfig,
26 pub label: String,
28 pub size_tree: bool,
30 pub size_class_tournaments: bool,
32 pub extended_size_targets: bool,
34 pub min_improvement: f64,
36 pub only: Option<String>,
38}
39
40impl Default for ProfileGenConfig {
41 fn default() -> Self {
42 Self {
43 measurement: MeasurementConfig {
44 warmup: 3,
45 iterations: 20,
46 ..Default::default()
47 },
48 label: "lab-train".into(),
49 size_tree: true,
50 size_class_tournaments: true,
51 extended_size_targets: true,
52 min_improvement: 0.03,
53 only: None,
54 }
55 }
56}
57
58struct IdVariant {
59 id: VariantId,
60 desc: &'static str,
61 body: Box<dyn Fn() + Send + Sync>,
62}
63
64impl Variant for IdVariant {
65 type Output = u64;
66 fn id(&self) -> &VariantId {
67 &self.id
68 }
69 fn description(&self) -> &str {
70 self.desc
71 }
72 fn run(&self) -> Self::Output {
73 (self.body)();
74 1
75 }
76}
77
78pub fn generate_profile(
80 info: &HardwareInfo,
81 cfg: &ProfileGenConfig,
82 out: &Path,
83) -> Result<(HnepProfile, Vec<TournamentResult>)> {
84 let fp = info.fingerprint.as_ref().ok_or_else(|| {
85 silicera::SiliceraError::UnsupportedCpu(
86 "cannot train: host is not a supported AMD Zen machine".into(),
87 )
88 })?;
89
90 let filter = cfg
91 .only
92 .as_deref()
93 .map(parse_only_filter)
94 .unwrap_or_default();
95
96 let mut results = Vec::new();
97 if target_matches_filter("memory-l2", &filter) {
98 results.push(tournament_memory(info, &cfg.measurement, cfg.min_improvement)?);
99 }
100 if target_matches_filter("integer", &filter) {
101 results.push(tournament_integer(&cfg.measurement, cfg.min_improvement)?);
102 }
103 if target_matches_filter("float", &filter) {
104 results.push(tournament_float(&cfg.measurement, cfg.min_improvement)?);
105 }
106 if target_matches_filter("branch", &filter) {
107 results.push(tournament_branch(&cfg.measurement, cfg.min_improvement)?);
108 }
109 if target_matches_filter("concurrency", &filter) {
110 results.push(tournament_concurrency(
111 info,
112 &cfg.measurement,
113 cfg.min_improvement,
114 )?);
115 }
116
117 let size_classes = if cfg.size_class_tournaments
118 && (filter.is_empty() || target_matches_filter("memscan-l1", &filter) || filter.iter().any(|f| f == "size_classes" || f.starts_with("memscan")))
119 {
120 let (entries, sc_results) =
121 measure_size_classes(info, &cfg.measurement, cfg.min_improvement, &filter)?;
122 results.extend(sc_results);
123 entries
124 } else {
125 Vec::new()
126 };
127
128 if cfg.extended_size_targets {
129 let (ext_entries, ext_results) =
130 measure_extended_size_targets(info, &cfg.measurement, cfg.min_improvement, &filter)?;
131 let _ = ext_entries;
134 results.extend(ext_results);
135 }
136
137 let tree = if cfg.size_tree && (filter.is_empty() || size_classes.len() == 4 || filter.iter().any(|f| f == "size_classes")) {
138 Some(if size_classes.is_empty() {
139 DecisionTree::from_topology(
140 &info.topology,
141 "l1_resident",
142 "l2_resident",
143 "l3_resident",
144 "dram_friendly",
145 "baseline",
146 )
147 } else {
148 DecisionTree::from_size_classes(&info.topology, &size_classes, "baseline")
149 })
150 } else {
151 let _ = SpecializeConfig::default();
152 None
153 };
154
155 let profile = if !filter.is_empty() && out.exists() {
157 merge_partial_profile(info, fp, cfg, out, results.clone(), size_classes, tree)?
158 } else {
159 HnepProfile::from_tournaments(
160 fp,
161 info.environment.clone(),
162 &results,
163 size_classes,
164 tree,
165 cfg.label.clone(),
166 )?
167 };
168 profile.write_to(out)?;
169 Ok((profile, results))
170}
171
172fn merge_partial_profile(
173 info: &HardwareInfo,
174 fp: &silicera::Fingerprint,
175 cfg: &ProfileGenConfig,
176 out: &Path,
177 new_results: Vec<TournamentResult>,
178 new_size_classes: Vec<SizeClassEntry>,
179 new_tree: Option<DecisionTree>,
180) -> Result<HnepProfile> {
181 let mut existing = HnepProfile::read_from(out)?;
182 for r in &new_results {
184 let winner_rec = r.records.iter().find(|x| x.id == r.winner);
185 let baseline_rec = r.records.iter().find(|x| x.id.0 == "baseline");
186 let entry = silicera::WorkloadEntry {
187 name: r.name.clone(),
188 winner: r.winner.0.clone(),
189 confidence: r.confidence,
190 rationale: r.rationale.clone(),
191 winner_median_ns: winner_rec.map(|w| w.summary.median_ns),
192 baseline_median_ns: baseline_rec.map(|b| b.summary.median_ns),
193 };
194 if let Some(slot) = existing.workloads.iter_mut().find(|w| w.name == r.name) {
195 *slot = entry;
196 } else {
197 existing.workloads.push(entry);
198 }
199 }
200 for sc in new_size_classes {
201 if let Some(slot) = existing
202 .size_classes
203 .iter_mut()
204 .find(|c| c.class == sc.class)
205 {
206 *slot = sc;
207 } else {
208 existing.size_classes.push(sc);
209 }
210 }
211 if let Some(tree) = new_tree {
212 existing.decision_tree = Some(tree);
213 } else if !existing.size_classes.is_empty() {
214 existing.decision_tree = Some(DecisionTree::from_size_classes(
215 &info.topology,
216 &existing.size_classes,
217 "baseline",
218 ));
219 }
220 existing.environment = info.environment.clone();
221 existing.header.fingerprint = fp.value.clone();
222 existing.header.label = cfg.label.clone();
223 existing.header.created_at = chrono::Utc::now().to_rfc3339();
224 existing.header.silicera_version = silicera::VERSION.into();
225 existing.header.version = silicera::HNEP_VERSION;
226 existing.recompute_digest()?;
227 Ok(existing)
228}
229
230#[derive(Debug, Clone, Serialize, Deserialize)]
232pub struct SizeClassWinnerStory {
233 pub winners: Vec<(String, String)>,
235 pub winners_differ: bool,
237 pub summary: String,
239}
240
241pub fn size_class_winner_story(classes: &[SizeClassEntry]) -> SizeClassWinnerStory {
243 let winners: Vec<(String, String)> = classes
244 .iter()
245 .map(|c| (c.class.clone(), c.winner.clone()))
246 .collect();
247 let winners_differ = winners
248 .first()
249 .map(|(_, w)| winners.iter().any(|(_, x)| x != w))
250 .unwrap_or(false);
251 let summary = if winners.is_empty() {
252 "No size-class measurements".into()
253 } else if winners_differ {
254 format!(
255 "Winners DO change by size: {}",
256 winners
257 .iter()
258 .map(|(c, w)| format!("{c}={w}"))
259 .collect::<Vec<_>>()
260 .join(", ")
261 )
262 } else {
263 format!(
264 "Winners do NOT change by size on this host/run (all → {})",
265 winners[0].1
266 )
267 };
268 SizeClassWinnerStory {
269 winners,
270 winners_differ,
271 summary,
272 }
273}
274
275pub fn measure_size_classes(
277 info: &HardwareInfo,
278 mcfg: &MeasurementConfig,
279 min_improvement: f64,
280 filter: &[String],
281) -> Result<(Vec<SizeClassEntry>, Vec<TournamentResult>)> {
282 let mut entries = Vec::new();
283 let mut results = Vec::new();
284 for target in CacheTarget::all() {
285 let name = format!("memscan-{}", target.class_id().to_lowercase());
286 if !target_matches_filter(&name, filter) {
287 continue;
288 }
289 let (entry, result) = tournament_memop_size_class(info, target, mcfg, min_improvement)?;
290 entries.push(entry);
291 results.push(result);
292 }
293 Ok((entries, results))
294}
295
296fn measure_extended_size_targets(
298 info: &HardwareInfo,
299 mcfg: &MeasurementConfig,
300 min_improvement: f64,
301 filter: &[String],
302) -> Result<(Vec<SizeClassEntry>, Vec<TournamentResult>)> {
303 let mut results = Vec::new();
304 for target in CacheTarget::all() {
305 let ws = target.size_bytes(&info.topology) as usize;
306 let fname = format!("float-{}", target.class_id().to_lowercase());
307 if target_matches_filter(&fname, filter) {
308 let n = (ws / 8).clamp(256, 1 << 20);
309 results.push(run_pair(
310 &fname,
311 mcfg,
312 min_improvement,
313 {
314 let a = FloatBench::new(n);
315 move || {
316 let _ = a.run_baseline();
317 }
318 },
319 {
320 let b = FloatBench::new(n);
321 move || {
322 let _ = b.run_candidate();
323 }
324 },
325 )?);
326 }
327 let iname = format!("intmix-{}", target.class_id().to_lowercase());
328 if target_matches_filter(&iname, filter) {
329 results.push(tournament_integer_named(&iname, mcfg, min_improvement)?);
331 }
332 }
333 Ok((Vec::new(), results))
334}
335
336fn tournament_memop_size_class(
337 info: &HardwareInfo,
338 target: CacheTarget,
339 mcfg: &MeasurementConfig,
340 min_improvement: f64,
341) -> Result<(SizeClassEntry, TournamentResult)> {
342 let name = format!("memscan-{}", target.class_id().to_lowercase());
343 let bench = Arc::new(Mutex::new(MemOpBench::for_target(&info.topology, target)));
344 let b_scan = Arc::clone(&bench);
345 let b_dense = Arc::clone(&bench);
346 let b_copy = Arc::clone(&bench);
347 let b_loop = Arc::clone(&bench);
348 let b_unroll = Arc::clone(&bench);
349
350 let variants = vec![
351 IdVariant {
352 id: VariantId::new("baseline"),
353 desc: "scan_stride",
354 body: Box::new(move || {
355 let g = b_scan.lock().unwrap();
356 let _ = g.run_scan_stride();
357 }),
358 },
359 IdVariant {
360 id: VariantId::new("scan_dense"),
361 desc: "scan_dense",
362 body: Box::new(move || {
363 let g = b_dense.lock().unwrap();
364 let _ = g.run_scan_dense();
365 }),
366 },
367 IdVariant {
368 id: VariantId::new("copy"),
369 desc: "memcpy",
370 body: Box::new(move || {
371 let mut g = b_copy.lock().unwrap();
372 let _ = g.run_copy();
373 }),
374 },
375 IdVariant {
376 id: VariantId::new("copy_loop"),
377 desc: "copy_loop",
378 body: Box::new(move || {
379 let mut g = b_loop.lock().unwrap();
380 let _ = g.run_copy_loop();
381 }),
382 },
383 IdVariant {
384 id: VariantId::new("copy_unrolled8"),
385 desc: "copy_unrolled8",
386 body: Box::new(move || {
387 let mut g = b_unroll.lock().unwrap();
388 let _ = g.run_copy_unrolled8();
389 }),
390 },
391 ];
392
393 let result = Tournament::new(TournamentConfig {
394 measurement: mcfg.clone(),
395 baseline_id: VariantId::new("baseline"),
396 min_improvement,
397 ..Default::default()
398 })
399 .run(&name, &variants)?;
400
401 let winner_rec = result.records.iter().find(|x| x.id == result.winner);
402 let baseline_rec = result.records.iter().find(|x| x.id.0 == "baseline");
403 let entry = SizeClassEntry {
404 class: target.class_id().into(),
405 threshold_bytes: target.threshold_bytes(&info.topology),
406 working_set_bytes: target.size_bytes(&info.topology),
407 winner: result.winner.0.clone(),
408 confidence: result.confidence,
409 rationale: result.rationale.clone(),
410 winner_median_ns: winner_rec.map(|w| w.summary.median_ns),
411 baseline_median_ns: baseline_rec.map(|b| b.summary.median_ns),
412 };
413 Ok((entry, result))
414}
415
416fn run_pair(
417 name: &str,
418 mcfg: &MeasurementConfig,
419 min_improvement: f64,
420 baseline: impl Fn() + Send + Sync + 'static,
421 candidate: impl Fn() + Send + Sync + 'static,
422) -> Result<TournamentResult> {
423 let variants = vec![
424 IdVariant {
425 id: VariantId::new("baseline"),
426 desc: "baseline",
427 body: Box::new(baseline),
428 },
429 IdVariant {
430 id: VariantId::new("candidate"),
431 desc: "candidate",
432 body: Box::new(candidate),
433 },
434 ];
435 Tournament::new(TournamentConfig {
436 measurement: mcfg.clone(),
437 baseline_id: VariantId::new("baseline"),
438 min_improvement,
439 ..Default::default()
440 })
441 .run(name, &variants)
442}
443
444fn tournament_memory(
445 info: &HardwareInfo,
446 mcfg: &MeasurementConfig,
447 min_improvement: f64,
448) -> Result<TournamentResult> {
449 let b1 = MemoryBench::for_target(&info.topology, CacheTarget::L2);
450 let b2 = MemoryBench::for_target(&info.topology, CacheTarget::L2);
451 run_pair(
452 "memory-l2",
453 mcfg,
454 min_improvement,
455 move || {
456 let _ = b1.run();
457 },
458 move || {
459 let _ = b2.run_prefetch_friendly();
460 },
461 )
462}
463
464fn tournament_integer(mcfg: &MeasurementConfig, min_improvement: f64) -> Result<TournamentResult> {
465 tournament_integer_named("integer", mcfg, min_improvement)
466}
467
468fn tournament_integer_named(
469 name: &str,
470 mcfg: &MeasurementConfig,
471 min_improvement: f64,
472) -> Result<TournamentResult> {
473 let a = IntegerBench { n: 42 };
474 let b = IntegerBench { n: 42 };
475 run_pair(
476 name,
477 mcfg,
478 min_improvement,
479 move || {
480 let _ = a.run_baseline();
481 },
482 move || {
483 let _ = b.run_candidate();
484 },
485 )
486}
487
488fn tournament_float(mcfg: &MeasurementConfig, min_improvement: f64) -> Result<TournamentResult> {
489 let a = FloatBench::new(4096);
490 let b = FloatBench::new(4096);
491 run_pair(
492 "float",
493 mcfg,
494 min_improvement,
495 move || {
496 let _ = a.run_baseline();
497 },
498 move || {
499 let _ = b.run_candidate();
500 },
501 )
502}
503
504fn tournament_branch(mcfg: &MeasurementConfig, min_improvement: f64) -> Result<TournamentResult> {
505 let a = BranchBench::new(8192);
506 let b = BranchBench::new(8192);
507 run_pair(
508 "branch",
509 mcfg,
510 min_improvement,
511 move || {
512 let _ = a.run_baseline();
513 },
514 move || {
515 let _ = b.run_candidate();
516 },
517 )
518}
519
520fn tournament_concurrency(
521 info: &HardwareInfo,
522 mcfg: &MeasurementConfig,
523 min_improvement: f64,
524) -> Result<TournamentResult> {
525 let threads = 4.min(info.topology.thread_count().max(1));
526 let a = ConcurrencyBench {
527 iters: 20_000,
528 threads,
529 };
530 let b = ConcurrencyBench {
531 iters: 20_000,
532 threads,
533 };
534 run_pair(
535 "concurrency",
536 mcfg,
537 min_improvement,
538 move || {
539 let _ = a.run_baseline();
540 },
541 move || {
542 let _ = b.run_candidate();
543 },
544 )
545}
546
547#[derive(Debug, Clone, Serialize, Deserialize)]
549pub struct DispatchSizeSample {
550 pub size_bytes: u64,
552 pub direct_median_ns: f64,
554 pub tree_median_ns: f64,
556 pub delta_ns: f64,
558}
559
560#[derive(Debug, Clone, Serialize, Deserialize)]
562pub struct DispatchOverheadReport {
563 pub batch_direct: MeasurementSummary,
565 pub batch_tree: MeasurementSummary,
567 pub batch_delta_ns: f64,
569 pub per_size: Vec<DispatchSizeSample>,
571 pub delta_median_ns: f64,
573 pub delta_mean_ns: f64,
575 pub delta_min_ns: f64,
577 pub delta_max_ns: f64,
579 pub note: String,
581}
582
583pub fn measure_dispatch_overhead(
588 tree: &DecisionTree,
589 mcfg: &MeasurementConfig,
590) -> Result<(MeasurementSummary, MeasurementSummary)> {
591 let report = measure_dispatch_overhead_report(tree, mcfg)?;
592 Ok((report.batch_direct, report.batch_tree))
593}
594
595pub fn measure_dispatch_overhead_report(
597 tree: &DecisionTree,
598 mcfg: &MeasurementConfig,
599) -> Result<DispatchOverheadReport> {
600 use silicera::measure::MeasurementEngine;
601 let eng = MeasurementEngine::new(mcfg.clone());
602 let sizes = [1024u64, 16_384, 100_000, 1_000_000, 2_000_000, 64_000_000];
603 const INNER: u64 = 4_096;
604
605 let batch_direct = eng.measure(|| {
606 let mut x = 0u64;
607 for &s in &sizes {
608 for i in 0..INNER {
609 x = std::hint::black_box(
610 x.wrapping_add(s)
611 .wrapping_mul(3)
612 .wrapping_add(i)
613 .wrapping_add(1),
614 );
615 }
616 }
617 std::hint::black_box(x);
618 })?;
619 let batch_tree = eng.measure(|| {
620 let mut x = 0u64;
621 for &s in &sizes {
622 for i in 0..INNER {
623 let v = tree.evaluate(std::hint::black_box(s.wrapping_add(i % 17)));
624 let mix = v.as_bytes().iter().fold(0u64, |a, &b| a.wrapping_add(b as u64));
625 x = std::hint::black_box(
626 x.wrapping_add(mix)
627 .wrapping_add(s)
628 .wrapping_mul(3)
629 .wrapping_add(i)
630 .wrapping_add(1),
631 );
632 }
633 }
634 std::hint::black_box(x);
635 })?;
636
637 let mut per_size = Vec::new();
638 for &s in &sizes {
639 let direct = eng.measure(|| {
640 let mut x = 0u64;
641 for i in 0..INNER {
642 x = std::hint::black_box(x.wrapping_add(s).wrapping_add(i).wrapping_mul(3));
643 }
644 std::hint::black_box(x);
645 })?;
646 let via = eng.measure(|| {
647 let mut x = 0u64;
648 for i in 0..INNER {
649 let v = tree.evaluate(std::hint::black_box(s.wrapping_add(i % 17)));
650 let mix = v.as_bytes().iter().fold(0u64, |a, &b| a.wrapping_add(b as u64));
651 x = std::hint::black_box(x.wrapping_add(mix).wrapping_add(s).wrapping_add(i));
652 }
653 std::hint::black_box(x);
654 })?;
655 per_size.push(DispatchSizeSample {
656 size_bytes: s,
657 direct_median_ns: direct.median_ns,
658 tree_median_ns: via.median_ns,
659 delta_ns: via.median_ns - direct.median_ns,
660 });
661 }
662
663 let mut deltas: Vec<f64> = per_size.iter().map(|p| p.delta_ns).collect();
664 deltas.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Less));
665 let delta_median_ns = if deltas.is_empty() {
666 0.0
667 } else if deltas.len() % 2 == 0 {
668 (deltas[deltas.len() / 2 - 1] + deltas[deltas.len() / 2]) / 2.0
669 } else {
670 deltas[deltas.len() / 2]
671 };
672 let delta_mean_ns = if deltas.is_empty() {
673 0.0
674 } else {
675 deltas.iter().sum::<f64>() / deltas.len() as f64
676 };
677 let delta_min_ns = deltas.first().copied().unwrap_or(0.0);
678 let delta_max_ns = deltas.last().copied().unwrap_or(0.0);
679
680 Ok(DispatchOverheadReport {
681 batch_delta_ns: batch_tree.median_ns - batch_direct.median_ns,
682 batch_direct,
683 batch_tree,
684 per_size,
685 delta_median_ns,
686 delta_mean_ns,
687 delta_min_ns,
688 delta_max_ns,
689 note: "Heavy stand-in (4096 iters × sizes): tree.evaluate + string mix vs direct add. \
690 Not a full variant call. Host-specific; quiet system recommended."
691 .into(),
692 })
693}