gam_gpu/
dictionary_score.rs1pub const DEFAULT_DICTIONARY_SCORE_MIN_ELEMS: usize = 1 << 20;
12
13pub const DEFAULT_DICTIONARY_SCORE_TILE_ELEMS: usize =
17 gam_runtime::resource::LIBRARY_ROW_CHUNK_TARGET_BYTES / std::mem::size_of::<f32>();
18
19#[derive(Clone, Copy, Debug, Eq, PartialEq)]
22pub struct DictionaryScoreRoutePlan {
23 pub n_rows: usize,
25 pub n_items: usize,
27 pub feature_dim: usize,
29 pub device_min_score_elems: usize,
31 pub max_tile_score_elems: usize,
33 pub tile_items: usize,
35 pub tile_count: usize,
37 pub device_admitted: bool,
39 pub peak_score_bytes: usize,
41 pub dot_flops_lower_bound: u128,
44}
45
46impl DictionaryScoreRoutePlan {
47 #[must_use]
51 pub fn with_limits(
52 n_rows: usize,
53 n_items: usize,
54 feature_dim: usize,
55 device_min_score_elems: usize,
56 max_tile_score_elems: usize,
57 ) -> Self {
58 let total_score_elems = n_rows.saturating_mul(n_items);
59 let nondegenerate = n_rows > 0 && n_items > 0 && feature_dim > 0;
60 let tile_items = if !nondegenerate {
61 0
62 } else {
63 (max_tile_score_elems / n_rows).clamp(1, n_items)
64 };
65 let tile_count = if tile_items == 0 {
66 0
67 } else {
68 n_items.div_ceil(tile_items)
69 };
70 let peak_tile_items = tile_items.min(n_items);
71 let peak_score_elems = n_rows.saturating_mul(peak_tile_items);
72 let dot_flops_lower_bound = 2u128
73 .saturating_mul(n_rows as u128)
74 .saturating_mul(n_items as u128)
75 .saturating_mul(feature_dim as u128);
76
77 Self {
78 n_rows,
79 n_items,
80 feature_dim,
81 device_min_score_elems,
82 max_tile_score_elems,
83 tile_items,
84 tile_count,
85 device_admitted: nondegenerate && total_score_elems >= device_min_score_elems,
86 peak_score_bytes: peak_score_elems.saturating_mul(std::mem::size_of::<f32>()),
87 dot_flops_lower_bound,
88 }
89 }
90
91 #[must_use]
93 pub fn default_for_shape(n_rows: usize, n_items: usize, feature_dim: usize) -> Self {
94 Self::with_limits(
95 n_rows,
96 n_items,
97 feature_dim,
98 DEFAULT_DICTIONARY_SCORE_MIN_ELEMS,
99 DEFAULT_DICTIONARY_SCORE_TILE_ELEMS,
100 )
101 }
102
103}
104
105#[cfg(test)]
106mod tests {
107 use super::*;
108
109 #[test]
110 fn target_k32k_shape_is_admitted_and_memory_bounded() {
111 let plan = DictionaryScoreRoutePlan::default_for_shape(256, 32_768, 64);
112 assert!(plan.device_admitted);
113 assert_eq!(plan.tile_items, 8_192);
114 assert_eq!(plan.tile_count, 4);
115 assert_eq!(
119 plan.peak_score_bytes,
120 gam_runtime::resource::LIBRARY_ROW_CHUNK_TARGET_BYTES
121 );
122 assert_eq!(
123 plan.dot_flops_lower_bound,
124 2u128 * 256u128 * 32_768u128 * 64u128
125 );
126 }
127
128 #[test]
129 fn peak_score_memory_does_not_grow_with_dictionary_width() {
130 let small = DictionaryScoreRoutePlan::default_for_shape(512, 4_096, 48);
131 let large = DictionaryScoreRoutePlan::default_for_shape(512, 131_072, 48);
132 assert_eq!(small.tile_items, large.tile_items);
133 assert_eq!(small.peak_score_bytes, large.peak_score_bytes);
134 assert!(large.tile_count > small.tile_count);
135 }
136
137 #[test]
138 fn tiny_tile_budget_still_makes_forward_progress() {
139 let plan = DictionaryScoreRoutePlan::with_limits(512, 1000, 32, 1, 7);
140 assert_eq!(plan.tile_items, 1);
141 assert_eq!(plan.tile_count, 1000);
142 assert_eq!(plan.peak_score_bytes, 512 * std::mem::size_of::<f32>());
143 }
144}