Skip to main content

greeners_ols/
three_sls.rs

1use greeners_core::linalg::LinalgInverse as _;
2use greeners_core::GreenersError;
3use ndarray::{Array1, Array2, Axis};
4use statrs::distribution::ContinuousCDF;
5use std::fmt;
6
7/// Structure to define a single system equation
8#[derive(Clone)]
9pub struct Equation {
10    pub y: Array1<f64>,
11    pub x: Array2<f64>, //Includes endogenous and exogenous
12    pub name: String,
13    pub var_names: Vec<String>,
14}
15
16/// 3SLS System Result
17#[derive(Debug)]
18pub struct ThreeSLSResult {
19    pub equations: Vec<EquationResult>,
20    pub sigma_cross: Array2<f64>, //Covariance matrix of errors between equations
21    pub system_r2: f64,           //McElroy's R2 (Optional but chic)
22}
23
24#[derive(Debug)]
25pub struct EquationResult {
26    pub name: String,
27    pub params: Array1<f64>,
28    pub std_errors: Array1<f64>,
29    pub t_values: Array1<f64>,
30    pub p_values: Array1<f64>,
31    pub r_squared: f64,
32    pub var_names: Vec<String>,
33}
34
35impl fmt::Display for ThreeSLSResult {
36    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
37        writeln!(f, "\n{:=^78}", " Three-Stage Least Squares (3SLS) System ")?;
38        writeln!(f, "Number of Equations: {}", self.equations.len())?;
39
40        //Show the Cross-Equation Correlation Correlation matrix
41        writeln!(f, "\n{:-^78}", " Residual Covariance Matrix (Sigma) ")?;
42        for row in self.sigma_cross.rows() {
43            write!(f, "[ ")?;
44            for val in row {
45                write!(f, "{:>10.4} ", val)?;
46            }
47            writeln!(f, "]")?;
48        }
49
50        for eq in &self.equations {
51            writeln!(f, "\n{:-^78}", format!(" Equation: {} ", eq.name))?;
52            writeln!(
53                f,
54                "{:<10} | {:>10} | {:>10} | {:>8} | {:>8}",
55                "Variable", "Coef", "Std Err", "t", "P>|t|"
56            )?;
57            writeln!(f, "{:-^78}", "")?;
58
59            for i in 0..eq.params.len() {
60                let label = eq
61                    .var_names
62                    .get(i)
63                    .cloned()
64                    .unwrap_or_else(|| format!("x{i}"));
65                writeln!(
66                    f,
67                    "{:<10} | {:>10.4} | {:>10.4} | {:>8.3} | {:>8.3}",
68                    label, eq.params[i], eq.std_errors[i], eq.t_values[i], eq.p_values[i]
69                )?;
70            }
71            writeln!(f, "R-squared: {:.4}", eq.r_squared)?;
72        }
73        writeln!(f, "{:=^78}", "")
74    }
75}
76
77pub struct ThreeSLS;
78
79/// Check if the instrument matrix already has a constant column.
80/// If not, add a 1s column at the beginning. That makes it
81/// projections of a first stage constant are accurate, allowing
82/// which 3SLS equations include intercept when the frontend so specifies.
83fn ensure_constant_instruments(z: &Array2<f64>) -> Array2<f64> {
84    let n = z.nrows();
85    if n == 0 {
86        return z.clone();
87    }
88
89    let has_const = z
90        .axis_iter(Axis(1))
91        .any(|col| col.iter().all(|&v| (v - 1.0).abs() < 1e-12));
92
93    if has_const {
94        z.clone()
95    } else {
96        let mut z_out = Array2::<f64>::ones((n, z.ncols() + 1));
97        z_out.slice_mut(ndarray::s![.., 1..]).assign(z);
98        z_out
99    }
100}
101
102impl ThreeSLS {
103    /// Estimates a system of simultaneous equations via 3SLS.
104    ///
105    /// # Arguments
106    /// * `equations` - Vector of structures `Equation` (each with y and X).
107    /// * `z_instruments` - Matriz global de instrumentos (união de todas as exógenas).
108    pub fn fit(
109        equations: &[Equation],
110        z_instruments: &Array2<f64>,
111    ) -> Result<ThreeSLSResult, GreenersError> {
112        let n_obs = z_instruments.nrows();
113        let n_eq = equations.len();
114
115        // --- STAGE 1: Reduced Form & Projection ---
116        // Projetar cada X no espaço de Z para obter X_hat = Z(Z'Z)^-1 Z'X
117        // X_hat é a versão "limpa" das endógenas.
118
119        //Ensures that the instrument matrix includes a constant, so that
120        //intercept projections are exact when the structural equation has
121        //a column of 1s.
122        let z_instruments = ensure_constant_instruments(z_instruments);
123
124        // Pré-calcular P_z = Z (Z'Z)^-1 Z'
125        // Para eficiência, calculamos apenas a parte (Z'Z)^-1 Z' e multiplicamos depois
126        let z_t = z_instruments.t();
127        let ztz = z_t.dot(&z_instruments);
128        let ztz_inv = ztz.inv().map_err(|_| GreenersError::SingularMatrix)?;
129        let projection_matrix_part = z_instruments.dot(&ztz_inv).dot(&z_t); //N x N (Beware of memory here if N is huge)
130
131        let mut x_hat_list = Vec::new();
132        let mut residuals_2sls = Array2::<f64>::zeros((n_obs, n_eq));
133
134        // --- STAGE 2: 2SLS Equation-by-Equation ---
135        for (i, eq) in equations.iter().enumerate() {
136            // X_hat = P_z * X
137            let x_hat = projection_matrix_part.dot(&eq.x);
138
139            // Beta_2sls = (X_hat' X)^-1 X_hat' y
140            // Nota: Em 2SLS clássico, usamos X_hat' X_hat ou X_hat' X, é equivalente.
141            let xt_x = x_hat.t().dot(&eq.x);
142            let xt_x_inv = xt_x.inv().map_err(|_| GreenersError::SingularMatrix)?;
143            let xt_y = x_hat.t().dot(&eq.y);
144            let beta_2sls = xt_x_inv.dot(&xt_y);
145
146            //residuals u = y - X * beta (We use the original X for residuals!)
147            let pred = eq.x.dot(&beta_2sls);
148            let u = &eq.y - &pred;
149
150            //Save to next step
151            residuals_2sls.column_mut(i).assign(&u);
152            x_hat_list.push(x_hat);
153        }
154
155        //Calculate Error Covariance Matrix (Sigma)
156        // Sigma_ij = (u_i' u_j) / N
157        let sigma = residuals_2sls.t().dot(&residuals_2sls) / (n_obs as f64);
158        let sigma_inv = sigma.inv().map_err(|_| GreenersError::SingularMatrix)?;
159
160        // --- STAGE 3: GLS Estimation on the System ---
161        // Resolver o sistema gigante: [X_hat' (Sigma^-1 ox I) X_hat] Beta = X_hat' (Sigma^-1 ox I) y
162
163        //1. Count total parameters
164        let mut k_total = 0;
165        let mut k_per_eq = Vec::new();
166        for eq in equations {
167            let k = eq.x.ncols();
168            k_per_eq.push(k);
169            k_total += k;
170        }
171
172        //2. Build LHS Matrix (System Hessian) and RHS Vector
173        //We use block construction to avoid explicit Kronecker.
174        let mut lhs_system = Array2::<f64>::zeros((k_total, k_total));
175        let mut rhs_system = Array1::<f64>::zeros(k_total);
176
177        let mut start_i = 0;
178        for i in 0..n_eq {
179            let ki = k_per_eq[i];
180            let x_hat_i = &x_hat_list[i];
181
182            let mut start_j = 0;
183            for j in 0..n_eq {
184                let kj = k_per_eq[j];
185                let x_hat_j = &x_hat_list[j];
186
187                // Elemento Sigma^{ij} (escalar)
188                let s_ij = sigma_inv[[i, j]];
189
190                // Bloco LHS = s_ij * (X_hat_i' * X_hat_j)
191                let block = x_hat_i.t().dot(x_hat_j) * s_ij;
192
193                // Inserir na matriz grandona
194                lhs_system
195                    .slice_mut(ndarray::s![start_i..start_i + ki, start_j..start_j + kj])
196                    .assign(&block);
197
198                //Part of RHS (only when loop j runs, accumulates for i)
199                // RHS_i = sum_j (s_ij * X_hat_i' * y_j)
200                let y_j = &equations[j].y;
201                let vec_part = x_hat_i.t().dot(y_j) * s_ij;
202
203                //Add to RHS vector in position i
204                let mut target_slice = rhs_system.slice_mut(ndarray::s![start_i..start_i + ki]);
205                target_slice += &vec_part;
206
207                start_j += kj;
208            }
209            start_i += ki;
210        }
211
212        // 3. Resolver Beta 3SLS
213        let lhs_inv = lhs_system
214            .inv()
215            .map_err(|_| GreenersError::SingularMatrix)?;
216        let beta_3sls_all = lhs_inv.dot(&rhs_system);
217
218        //--- POST-STIMATION: Separate results and Statistics ---
219        let mut final_results = Vec::new();
220        let mut cursor = 0;
221
222        for (i, eq) in equations.iter().enumerate() {
223            let k = k_per_eq[i];
224            let params = beta_3sls_all
225                .slice(ndarray::s![cursor..cursor + k])
226                .to_owned();
227
228            //Variance Asymptotic coefficients of this equation
229            //It is the corresponding diagonal block of the system inverse
230            let cov_params = lhs_inv
231                .slice(ndarray::s![cursor..cursor + k, cursor..cursor + k])
232                .to_owned();
233            let std_errors = cov_params.diag().mapv(f64::sqrt);
234
235            //Statistics T and P
236            let t_values = &params / &std_errors;
237            let p_values = t_values
238                .mapv(|t| 2.0 * (1.0 - statrs::distribution::Normal::standard().cdf(t.abs())));
239
240            //R2 (Using 3SLS final residuals)
241            let pred = eq.x.dot(&params);
242            let res = &eq.y - &pred;
243            let sst = (&eq.y
244                - eq.y.mean().ok_or_else(|| {
245                    GreenersError::InvalidOperation("Empty dependent variable".to_string())
246                })?)
247            .mapv(|v| v.powi(2))
248            .sum();
249            let ssr = res.mapv(|v| v.powi(2)).sum();
250            let r2 = 1.0 - (ssr / sst);
251
252            final_results.push(EquationResult {
253                name: eq.name.clone(),
254                params,
255                std_errors,
256                t_values,
257                p_values,
258                r_squared: r2,
259                var_names: eq.var_names.clone(),
260            });
261
262            cursor += k;
263        }
264
265        Ok(ThreeSLSResult {
266            equations: final_results,
267            sigma_cross: sigma,
268            system_r2: 0.0, // Placeholder
269        })
270    }
271}