// src/services/observation/gaia_xp/basis.rs // // Gaia XP 基函数配置解析 + Hermite 函数递归求值 // // 配置文件(include_str! 编译时打包,来自 GaiaXPy 包): // {bp|rp}C03_{model}_bases.csv —— 单行 9 字段,含 inv_coef(55×55) + transf(55×55) 扁平数组 // {bp|rp}C03_{model}_dispersion.csv —— wl→pwl 映射(97/95 点) // {bp|rp}C03_{model}_response.csv —— wl→response 透过率(1581 点) // // Hermite 函数(probabilists',对齐 GaiaXPy __psi): // ψ_0(x) = π^(-1/4) · exp(-x²/2) // ψ_1(x) = π^(-1/4) · exp(-x²/2) · √2 · x // ψ_n(x) = √(2/n)·x·ψ_{n-1}(x) - √((n-1)/n)·ψ_{n-2}(x) use anyhow::{anyhow, Context, Result}; /// 单波段配置(bases + dispersion + response) pub struct BandConfig { pub inv_coef: Vec>, // [55][55] pub transf: Vec>, // [55][55] pub disp_x: Vec, // dispersion 输入 wl pub disp_y: Vec, // dispersion 输出 pwl pub resp_x: Vec, // response 输入 wl pub resp_y: Vec, // response 输出 透过率 pub scale: f64, pub offset: f64, } /// 加载波段配置(bp/rp) pub fn load_band_config(band: &str) -> Result { let (bases_csv, disp_csv, resp_csv) = match band { "bp" => ( include_str!("config/bpC03_v375wi_bases.csv"), include_str!("config/bpC03_v375wi_dispersion.csv"), include_str!("config/bpC03_v375wi_response.csv"), ), "rp" => ( include_str!("config/rpC03_v142r_bases.csv"), include_str!("config/rpC03_v142r_dispersion.csv"), include_str!("config/rpC03_v142r_response.csv"), ), _ => return Err(anyhow!("未知波段: {}", band)), }; let bases = parse_bases_csv(bases_csv)?; let (disp_x, disp_y) = parse_xy_csv(disp_csv)?; let (resp_x, resp_y) = parse_xy_csv(resp_csv)?; // scale/offset 从 bases 的 pwl/norm 范围计算 let scale = (bases.norm_range_max - bases.norm_range_min) / (bases.pwl_range_max - bases.pwl_range_min); let offset = bases.norm_range_min - bases.pwl_range_min * scale; Ok(BandConfig { inv_coef: bases.inv_coef, transf: bases.transf, disp_x, disp_y, resp_x, resp_y, scale, offset, }) } struct ParsedBases { inv_coef: Vec>, // [55][55] transf: Vec>, // [55][55] pwl_range_min: f64, pwl_range_max: f64, norm_range_min: f64, norm_range_max: f64, } /// 解析 bases CSV /// /// 实测格式(GaiaXPy bpC03_v375wi_bases.csv): /// 行0 = 表头(9 列名) /// 行1 = 数据,9 个 CSV 字段,其中第 6/8 字段是引号包裹的括号数组 "(v1,v2,...,v3025)" /// (csv 引号使括号内的逗号不作为字段分隔符) /// /// 字段顺序:nBases, pwlRangeMin, pwlRangeMax, normRangeMin, normRangeMax, /// nInverseBasesCoefficients, inverseBasesCoefficients, nTransformedBases, transformationMatrix fn parse_bases_csv(csv: &str) -> Result { // 用简单的状态机解析 CSV(处理引号内的逗号),取第 2 个非空行(数据行) let rows = parse_csv_rows(csv); let data_row = rows .into_iter() .filter(|r| !r.is_empty()) .nth(1) // 跳过表头 .ok_or_else(|| anyhow!("bases CSV 无数据行"))?; if data_row.len() < 9 { return Err(anyhow!("bases CSV 字段数不足: {}", data_row.len())); } let n_bases: usize = data_row[0].trim().parse().context("nBases")?; let pwl_min: f64 = data_row[1].trim().parse().context("pwlRangeMin")?; let pwl_max: f64 = data_row[2].trim().parse().context("pwlRangeMax")?; let norm_min: f64 = data_row[3].trim().parse().context("normRangeMin")?; let norm_max: f64 = data_row[4].trim().parse().context("normRangeMax")?; // field[6] = inverseBasesCoefficients "(v1,v2,...)",field[8] = transformationMatrix let inv_flat = parse_paren_array(&data_row[6]).context("解析 inverseBasesCoefficients")?; let transf_flat = parse_paren_array(&data_row[8]).context("解析 transformationMatrix")?; if inv_flat.len() != n_bases * n_bases { return Err(anyhow!( "inverseBasesCoefficients 长度 {} != {}×{}", inv_flat.len(), n_bases, n_bases )); } if transf_flat.len() != n_bases * n_bases { return Err(anyhow!( "transformationMatrix 长度 {} != {}×{}", transf_flat.len(), n_bases, n_bases )); } let inv_coef = reshape_square(&inv_flat, n_bases); let transf = reshape_square(&transf_flat, n_bases); Ok(ParsedBases { inv_coef, transf, pwl_range_min: pwl_min, pwl_range_max: pwl_max, norm_range_min: norm_min, norm_range_max: norm_max, }) } /// 简易 CSV 行解析(处理双引号内的逗号与换行) fn parse_csv_rows(csv: &str) -> Vec> { let mut rows = Vec::new(); let mut row = Vec::new(); let mut field = String::new(); let mut in_quotes = false; for ch in csv.chars() { match ch { '"' => in_quotes = !in_quotes, ',' if !in_quotes => { row.push(std::mem::take(&mut field)); } '\n' if !in_quotes => { row.push(std::mem::take(&mut field)); if !row.is_empty() { rows.push(std::mem::take(&mut row)); } } '\r' if !in_quotes => {} _ => field.push(ch), } } if !field.is_empty() || !row.is_empty() { row.push(field); rows.push(row); } rows } /// 解析括号数组 "(v1,v2,...,vN)" → Vec fn parse_paren_array(s: &str) -> Result> { let s = s.trim(); let inner = s .strip_prefix('(') .and_then(|s| s.strip_suffix(')')) .ok_or_else(|| anyhow!("期望括号数组,得到: {}", &s[..s.len().min(40)]))?; inner .split(',') .map(|t| { t.trim() .parse::() .map_err(|e| anyhow!("数值解析失败 '{}': {}", t, e)) }) .collect() } /// 扁平数组 → n×n 矩阵(行优先) fn reshape_square(flat: &[f64], n: usize) -> Vec> { let mut m = vec![vec![0.0; n]; n]; for i in 0..n { for j in 0..n { m[i][j] = flat[i * n + j]; } } m } /// 解析 dispersion/response CSV /// /// 实测格式(GaiaXPy bpC03_v375wi_dispersion.csv): /// 行0 = 所有 X 值(逗号分隔,97 或 1581 个) /// 行1 = 所有 Y 值(逗号分隔,同数量) /// 无表头 fn parse_xy_csv(csv: &str) -> Result<(Vec, Vec)> { let rows = parse_csv_rows(csv); let x_row = rows.first().ok_or_else(|| anyhow!("XY CSV 无数据行"))?; let y_row = rows .get(1) .ok_or_else(|| anyhow!("XY CSV 缺第二行(Y 值)"))?; if x_row.len() != y_row.len() { return Err(anyhow!( "X/Y 行长度不匹配: {} vs {}", x_row.len(), y_row.len() )); } let x: Vec = x_row .iter() .map(|s| s.trim().parse::()) .collect::, _>>() .context("解析 X 值")?; let y: Vec = y_row .iter() .map(|s| s.trim().parse::()) .collect::, _>>() .context("解析 Y 值")?; if x.is_empty() { return Err(anyhow!("XY CSV 无有效数据")); } Ok((x, y)) } /// Hermite 函数 ψ_n(x),n=0..(max_n-1),3-term 递归 /// /// 对齐 GaiaXPy _evaluate_hermite_function / populate_design_matrix.__psi: /// ψ_0(x) = π^(-1/4) · exp(-x²/2) /// ψ_1(x) = π^(-1/4) · exp(-x²/2) · √2 · x /// ψ_n(x) = √(2/n)·x·ψ_{n-1}(x) - √((n-1)/n)·ψ_{n-2}(x) pub fn hermite_functions(x: f64, max_n: usize) -> Vec { let mut psi = vec![0.0; max_n]; if max_n == 0 { return psi; } let sqrt_4_pi = std::f64::consts::PI.powf(-0.25); // π^(-1/4) let g = (-x * x / 2.0).exp(); psi[0] = sqrt_4_pi * g; if max_n == 1 { return psi; } psi[1] = sqrt_4_pi * g * 2f64.sqrt() * x; for n in 2..max_n { let n_f = n as f64; psi[n] = (2.0 / n_f).sqrt() * x * psi[n - 1] - ((n_f - 1.0) / n_f).sqrt() * psi[n - 2]; } psi } #[cfg(test)] mod tests { use super::*; #[test] fn test_hermite_basic() { let psi = hermite_functions(0.0, 55); assert_eq!(psi.len(), 55); // ψ_0(0) = π^(-1/4) ≈ 0.7511 assert!((psi[0] - std::f64::consts::PI.powf(-0.25)).abs() < 1e-10); // ψ_1(0) = 0(含 x 因子) assert!(psi[1].abs() < 1e-15); } #[test] fn test_hermite_recursion() { // 在 x=1.0 处,所有值应有限 let psi = hermite_functions(1.0, 55); for v in &psi { assert!(v.is_finite()); } } #[test] fn test_reshape_square() { let flat = vec![1.0, 2.0, 3.0, 4.0]; let m = reshape_square(&flat, 2); assert_eq!(m[0], vec![1.0, 2.0]); assert_eq!(m[1], vec![3.0, 4.0]); } #[test] #[ignore = "需要编译时打包的 CSV"] fn test_load_bp_config() { let cfg = load_band_config("bp").unwrap(); assert_eq!(cfg.inv_coef.len(), 55); assert_eq!(cfg.inv_coef[0].len(), 55); assert_eq!(cfg.transf.len(), 55); assert!(!cfg.disp_x.is_empty()); assert!(!cfg.resp_x.is_empty()); } }