gam_terms/analytic_penalties/
nested_prefix.rs1use super::*;
2
3#[derive(Debug, Clone)]
36pub struct NestedPrefixPenalty {
37 pub target: PsiSlice,
38 pub target_tier: PenaltyTier,
39 pub prefix_sizes: Vec<usize>,
41 pub shell_weights: Vec<f64>,
44 pub eps: f64,
47 pub rho_indices: Vec<usize>,
49 pub weight_schedule: Option<ScalarWeightSchedule>,
50}
51
52impl NestedPrefixPenalty {
53 #[must_use = "build error must be handled"]
63 pub fn new(
64 target: PsiSlice,
65 target_tier: PenaltyTier,
66 prefix_sizes: Vec<usize>,
67 shell_weights: Vec<f64>,
68 eps: f64,
69 ) -> Result<Self, String> {
70 if prefix_sizes.is_empty() {
71 return Err("NestedPrefixPenalty requires at least one prefix".into());
72 }
73 if shell_weights.len() != prefix_sizes.len() {
74 return Err(format!(
75 "NestedPrefixPenalty requires shell_weights.len() == prefix_sizes.len(); \
76 got {} weights for {} prefixes",
77 shell_weights.len(),
78 prefix_sizes.len()
79 ));
80 }
81 for w in &shell_weights {
82 if !w.is_finite() || *w < 0.0 {
83 return Err(format!(
84 "NestedPrefixPenalty shell weights must be finite and ≥ 0; got {w}"
85 ));
86 }
87 }
88 for i in 0..prefix_sizes.len() {
89 if prefix_sizes[i] == 0 {
90 return Err("NestedPrefixPenalty prefixes must be > 0".into());
91 }
92 if i > 0 && prefix_sizes[i] <= prefix_sizes[i - 1] {
93 return Err(format!(
94 "NestedPrefixPenalty prefixes must be strictly increasing; got {:?}",
95 prefix_sizes
96 ));
97 }
98 }
99 if let Some(d) = target.latent_dim {
100 let max_prefix = *prefix_sizes.last().expect("non-empty");
101 if max_prefix > d {
102 return Err(format!(
103 "NestedPrefixPenalty largest prefix {max_prefix} exceeds latent_dim {d}"
104 ));
105 }
106 }
107 if !(eps.is_finite() && eps > 0.0) {
108 return Err(format!(
109 "NestedPrefixPenalty requires eps > 0 (1/sqrt(x²+ε²) singularity at 0); got {eps}"
110 ));
111 }
112 let rho_indices = (0..prefix_sizes.len()).collect();
113 Ok(Self {
114 target,
115 target_tier,
116 prefix_sizes,
117 shell_weights,
118 eps,
119 rho_indices,
120 weight_schedule: None,
121 })
122 }
123
124 #[must_use]
127 pub fn with_weight_schedule(mut self, schedule: ScalarWeightSchedule) -> Self {
128 self.weight_schedule = Some(schedule);
129 self
130 }
131
132 fn latent_dim(&self) -> usize {
134 self.target
135 .latent_dim
136 .unwrap_or_else(|| *self.prefix_sizes.last().expect("non-empty"))
137 }
138
139 fn lambdas(&self, rho: ArrayView1<'_, f64>) -> Vec<f64> {
141 self.prefix_sizes
142 .iter()
143 .enumerate()
144 .map(|(k, _)| {
145 validated_learnable_weight(self.shell_weights[k], rho[self.rho_indices[k]])
146 })
147 .collect()
148 }
149
150 fn per_axis_weights(&self, lambdas: &[f64]) -> Vec<f64> {
153 let f = self.latent_dim();
154 let mut w = vec![0.0_f64; f];
155 for (k, &m_k) in self.prefix_sizes.iter().enumerate() {
159 let lam = lambdas[k];
160 if lam == 0.0 {
161 continue;
162 }
163 let end = m_k.min(f);
164 for entry in w.iter_mut().take(end) {
165 *entry += lam;
166 }
167 }
168 w
169 }
170}
171
172impl AnalyticPenalty for NestedPrefixPenalty {
173 fn tier(&self) -> PenaltyTier {
174 self.target_tier
175 }
176
177 fn validate_rho(&self, rho: ArrayView1<'_, f64>) -> Result<(), String> {
178 if rho.len() != self.rho_count() {
179 return Err(format!(
180 "nested-prefix rho length {} != shell count {}",
181 rho.len(),
182 self.rho_count()
183 ));
184 }
185 for shell in 0..self.prefix_sizes.len() {
186 resolve_learnable_weight(self.shell_weights[shell], rho[self.rho_indices[shell]])?;
187 }
188 Ok(())
189 }
190
191 fn rho_coordinate_domains(&self) -> Result<Vec<(f64, f64)>, String> {
192 self.shell_weights
193 .iter()
194 .map(|&weight| {
195 learnable_weight_coordinate_domain(weight)?
196 .ok_or_else(|| "nested-prefix shell weight must be positive".to_string())
197 })
198 .collect()
199 }
200
201 fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
202 let f = self.latent_dim();
203 assert!(
204 target.len().is_multiple_of(f),
205 "target length must be n_rows · F"
206 );
207 let n_rows = target.len() / f;
208 let lambdas = self.lambdas(rho);
209 let eps2 = self.eps * self.eps;
210 let mut s_axis = vec![0.0_f64; f];
212 for n in 0..n_rows {
213 let row = &target.as_slice().expect("contiguous")[n * f..(n + 1) * f];
214 for (i, &x) in row.iter().enumerate() {
215 s_axis[i] += (x * x + eps2).sqrt();
216 }
217 }
218 let mut total = 0.0;
220 for (k, &m_k) in self.prefix_sizes.iter().enumerate() {
221 let end = m_k.min(f);
222 let mut acc = 0.0;
223 for &v in s_axis.iter().take(end) {
224 acc += v;
225 }
226 total += lambdas[k] * acc;
227 }
228 total
229 }
230
231 fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
232 let f = self.latent_dim();
233 let n_rows = target.len() / f;
234 let lambdas = self.lambdas(rho);
235 let w_per_axis = self.per_axis_weights(&lambdas);
236 let eps2 = self.eps * self.eps;
237 let src = target.as_slice().expect("contiguous");
238 let mut g = Array1::<f64>::zeros(target.len());
239 let g_slice = g.as_slice_mut().expect("contiguous");
240 for n in 0..n_rows {
241 for i in 0..f {
242 let x = src[n * f + i];
243 let w = w_per_axis[i];
244 if w == 0.0 {
245 continue;
246 }
247 g_slice[n * f + i] = w * x / (x * x + eps2).sqrt();
248 }
249 }
250 g
251 }
252
253 fn hessian_diag(
254 &self,
255 target: ArrayView1<'_, f64>,
256 rho: ArrayView1<'_, f64>,
257 ) -> Option<Array1<f64>> {
258 let f = self.latent_dim();
259 let n_rows = target.len() / f;
260 let lambdas = self.lambdas(rho);
261 let w_per_axis = self.per_axis_weights(&lambdas);
262 let eps2 = self.eps * self.eps;
263 let src = target.as_slice().expect("contiguous");
264 let mut d = Array1::<f64>::zeros(target.len());
265 let d_slice = d.as_slice_mut().expect("contiguous");
266 for n in 0..n_rows {
267 for i in 0..f {
268 let w = w_per_axis[i];
269 if w == 0.0 {
270 continue;
271 }
272 let x = src[n * f + i];
273 let r = (x * x + eps2).sqrt();
274 d_slice[n * f + i] = w * eps2 / (r * r * r);
275 }
276 }
277 Some(d)
278 }
279
280 fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
281 let f = self.latent_dim();
282 let n_rows = target.len() / f;
283 let lambdas = self.lambdas(rho);
284 let eps2 = self.eps * self.eps;
285 let mut s_axis = vec![0.0_f64; f];
288 let src = target.as_slice().expect("contiguous");
289 for n in 0..n_rows {
290 for i in 0..f {
291 let x = src[n * f + i];
292 s_axis[i] += (x * x + eps2).sqrt();
293 }
294 }
295 let n_rho = self.rho_count();
296 let mut out = Array1::<f64>::zeros(n_rho);
297 for (k, &m_k) in self.prefix_sizes.iter().enumerate() {
298 let end = m_k.min(f);
299 let mut shell_sum = 0.0;
300 for &v in s_axis.iter().take(end) {
301 shell_sum += v;
302 }
303 out[self.rho_indices[k]] = lambdas[k] * shell_sum;
305 }
306 out
307 }
308
309 fn rho_count(&self) -> usize {
310 self.prefix_sizes.len()
311 }
312
313 fn name(&self) -> &str {
314 "nested_prefix"
315 }
316
317}