1use crate::ols::OLS;
2use greeners_core::linalg::LinalgInverse as _;
3use greeners_core::{CovarianceType, GreenersError};
4use ndarray::{Array1, Array2};
5use std::fmt;
6
7#[derive(Clone)]
9pub struct SurEquation {
10 pub y: Array1<f64>,
11 pub x: Array2<f64>,
12 pub name: String,
13}
14
15#[derive(Debug)]
16pub struct SurResult {
17 pub equations: Vec<SurEquationResult>,
18 pub sigma_cross: Array2<f64>,
19 pub system_r2: f64,
20}
21
22#[derive(Debug)]
23pub struct SurEquationResult {
24 pub name: String,
25 pub params: Array1<f64>,
26 pub std_errors: Array1<f64>,
27 pub t_values: Array1<f64>,
28 pub p_values: Array1<f64>,
29 pub r_squared: f64,
30}
31
32impl fmt::Display for SurResult {
33 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
34 writeln!(f, "\n{:=^78}", " Seemingly Unrelated Regressions (SUR) ")?;
35 writeln!(f, "Zellner's Efficient Estimator")?;
36
37 writeln!(f, "\n{:-^78}", " Cross-Equation Error Correlation (Sigma) ")?;
38 for row in self.sigma_cross.rows() {
39 write!(f, "[ ")?;
40 for val in row {
41 write!(f, "{:>10.4} ", val)?;
42 }
43 writeln!(f, "]")?;
44 }
45
46 for eq in &self.equations {
47 writeln!(f, "\n{:-^78}", format!(" Equation: {} ", eq.name))?;
48 writeln!(
49 f,
50 "{:<10} | {:>10} | {:>10} | {:>8} | {:>8}",
51 "Variable", "Coef", "Std Err", "t", "P>|t|"
52 )?;
53 writeln!(f, "{:-^78}", "")?;
54
55 for i in 0..eq.params.len() {
56 writeln!(
57 f,
58 "x{:<9} | {:>10.4} | {:>10.4} | {:>8.3} | {:>8.3}",
59 i, eq.params[i], eq.std_errors[i], eq.t_values[i], eq.p_values[i]
60 )?;
61 }
62 writeln!(f, "R-squared: {:.4}", eq.r_squared)?;
63 }
64 writeln!(f, "{:=^78}", "")
65 }
66}
67
68pub struct SUR;
69
70impl SUR {
71 pub fn fit(equations: &[SurEquation]) -> Result<SurResult, GreenersError> {
74 let n_obs = equations[0].y.len();
75 let n_eq = equations.len();
76
77 let mut residuals_ols = Array2::<f64>::zeros((n_obs, n_eq));
79
80 for (i, eq) in equations.iter().enumerate() {
81 if eq.y.len() != n_obs {
82 return Err(GreenersError::ShapeMismatch(
83 "All equations must have same N observations".into(),
84 ));
85 }
86
87 let ols = OLS::fit(&eq.y, &eq.x, CovarianceType::NonRobust)?;
89
90 let pred = eq.x.dot(&ols.params);
92 let u = &eq.y - &pred;
93 residuals_ols.column_mut(i).assign(&u);
94 }
95
96 let sigma = residuals_ols.t().dot(&residuals_ols) / (n_obs as f64);
99
100 let sigma_inv = sigma.inv().map_err(|_| GreenersError::SingularMatrix)?;
102
103 let mut k_total = 0;
107 let mut k_per_eq = Vec::new();
108 for eq in equations {
109 let k = eq.x.ncols();
110 k_per_eq.push(k);
111 k_total += k;
112 }
113
114 let mut lhs = Array2::<f64>::zeros((k_total, k_total));
115 let mut rhs = Array1::<f64>::zeros(k_total);
116
117 let mut start_i = 0;
118 for i in 0..n_eq {
119 let ki = k_per_eq[i];
120 let xi = &equations[i].x;
121
122 let mut start_j = 0;
123 for j in 0..n_eq {
124 let kj = k_per_eq[j];
125 let xj = &equations[j].x;
126
127 let s_ij = sigma_inv[[i, j]];
129
130 let block = xi.t().dot(xj) * s_ij;
132 lhs.slice_mut(ndarray::s![start_i..start_i + ki, start_j..start_j + kj])
133 .assign(&block);
134
135 let yj = &equations[j].y;
138 let vec_part = xi.t().dot(yj) * s_ij;
139 let mut target_slice = rhs.slice_mut(ndarray::s![start_i..start_i + ki]);
140 target_slice += &vec_part;
141
142 start_j += kj;
143 }
144 start_i += ki;
145 }
146
147 let lhs_inv = lhs.inv().map_err(|_| GreenersError::SingularMatrix)?;
149 let beta_sur = lhs_inv.dot(&rhs);
150
151 let mut final_results = Vec::new();
153 let mut cursor = 0;
154 let normal = statrs::distribution::Normal::standard();
155 use statrs::distribution::ContinuousCDF;
156
157 for (i, eq) in equations.iter().enumerate() {
158 let k = k_per_eq[i];
159 let params = beta_sur.slice(ndarray::s![cursor..cursor + k]).to_owned();
160
161 let cov_params = lhs_inv
163 .slice(ndarray::s![cursor..cursor + k, cursor..cursor + k])
164 .to_owned();
165 let std_errors = cov_params.diag().mapv(f64::sqrt);
166
167 let t_values = ¶ms / &std_errors;
168 let p_values = t_values.mapv(|t| 2.0 * (1.0 - normal.cdf(t.abs())));
169
170 let pred = eq.x.dot(¶ms);
172 let res = &eq.y - &pred;
173 let sst = (&eq.y
174 - eq.y.mean().ok_or_else(|| {
175 GreenersError::InvalidOperation("Empty dependent variable".to_string())
176 })?)
177 .mapv(|v| v.powi(2))
178 .sum();
179 let ssr = res.mapv(|v| v.powi(2)).sum();
180 let r2 = 1.0 - (ssr / sst);
181
182 final_results.push(SurEquationResult {
183 name: eq.name.clone(),
184 params,
185 std_errors,
186 t_values,
187 p_values,
188 r_squared: r2,
189 });
190
191 cursor += k;
192 }
193
194 Ok(SurResult {
195 equations: final_results,
196 sigma_cross: sigma,
197 system_r2: 0.0,
198 })
199 }
200}