The Statistics API: One Atom, Many Schedules¶
Navigation:
Theory introduction: See the Intro
Related mathematical theory: One Atom, Many Statistics
This page shows only the implementation wiring. For the equations (working response, weights, links, quantiles) see the math theory.
Core interface: _base.py::BaseTAM._solve_pwls_step · Schedule layer: statistics/estimation/
The Atom and its bit-identical default¶
Per-observation weights enter through a numerically-stable sqrt(W) row-scaling of both Phi and Y; sample_weights=None never multiplies, so the default loss="l2" path is bit-identical to the original least-squares solver.
def _compute_weighted_covariances(
phi: torch.Tensor,
y_data: torch.Tensor,
loss_L_star_L: torch.Tensor,
sample_weights: Optional[torch.Tensor] = None
) -> tuple:
r"""
Computes the loss-weighted covariance matrices (LHS and RHS of normal equations).
Calculates:
- cov_X = Phi.H @ W @ (L*L) @ Phi (Weighted Feature Covariance)
- cov_XY = Phi.H @ W @ (L*L) @ Y (Weighted Feature-Target Covariance)
loss_L_star_L weights the output channel (shape (d_out, d_out)); sample_weights is an optional
per-observation weight w (shape (..., n_samples, 1)) applied by a stable sqrt(W) row-scaling of both
Phi and Y - the atom of any reweighted / GLM fit. sample_weights=None reproduces the plain
least-squares path bit-for-bit.
Args:
phi: The design matrix. Shape: (..., n_samples, n_coeffs).
y_data: The target tensor. Shape: (..., n_samples, d_out).
loss_L_star_L: The loss-weighting matrix. Shape: (d_out, d_out).
sample_weights: Optional per-observation weights. Shape: (..., n_samples, 1). Non-negative.
Returns:
Tuple[torch.Tensor, torch.Tensor]: (cov_X, cov_XY)
"""
# Ensure all tensors are on the same device
phi = phi.to(TORCH_DEVICE)
y_data = y_data.to(TORCH_DEVICE)
loss_L_star_L = loss_L_star_L.to(TORCH_DEVICE)
# Cast y_data to match phi's dtype
y_data_aligned = y_data.to(phi.dtype)
if sample_weights is not None:
root_w = sample_weights.to(device=TORCH_DEVICE, dtype=phi.dtype).clamp_min(0).sqrt()
phi = phi * root_w
y_data_aligned = y_data_aligned * root_w
# Compute Weighted Y
# y_weighted shape: (..., n_samples, d_out)
y_weighted = y_data_aligned @ loss_L_star_L
# Compute RHS: cov_XY = Phi^H @ Y_weighted
cov_XY = phi.mT @ y_weighted
# Compute LHS: cov_X
# Note: Current implementation treats single-target and multi-target differently
if loss_L_star_L.shape[0] == 1:
# Single-target: Apply scalar weight via square root
L_sqrt = loss_L_star_L[0, 0].sqrt()
phi_weighted = phi * L_sqrt
cov_X = phi_weighted.mT @ phi_weighted
else:
# Multi-target: Currently defaulting to unweighted feature covariance
# (Phi^H @ Phi) for stability in hierarchical cases.
cov_X = phi.mT @ phi
return cov_X, cov_XY
Every schedule reaches the solver through this single method, the atom (weights=None is the ordinary penalized solve; the schedules call it repeatedly with updated (z, W)).
def _solve_pwls_step(
self,
x_data: torch.Tensor,
z_target: torch.Tensor,
weights: Optional[torch.Tensor] = None
) -> torch.Tensor:
"""Solve one penalized weighted least-squares system: the atom of every TAM fit.
Solves (Phi.T * W * Phi + n*S) * theta = Phi.T * W * z, where z_target is the (working) response
and weights is an optional per-observation weight w of shape (n_groups, n_samples, 1).
weights=None is the plain penalized solve (W = I). The reweighting loop calls this repeatedly with
updated (z, w); a plain Gaussian fit calls it once.
Args:
x_data: feature tensor (n_groups, n_samples, n_features).
z_target: (working) response (n_groups, n_samples, 1).
weights: optional per-observation weights (n_groups, n_samples, 1), non-negative.
"""
penalty_M_star_M = self._build_penalty_matrix()
loss_L_star_L = self._build_loss_matrix()
return smart_solve(
x_data=x_data,
y_data=z_target,
effects_list=self.effects_list_,
penalty_matrix=penalty_M_star_M,
loss_matrix=loss_L_star_L,
num_samples=x_data.shape[1],
sample_weights=weights,
)
The router¶
StaticTAM.fit routes by its inputs: a string formula to standard IRLS, a {param: formula} dict to the location-scale schedule, and mixture_components=K to the EM schedule. No wrapper classes.
def fit(self, data_train: pd.DataFrame, **schedule_kwargs) -> "StaticTAM":
"""Fit, routing to the schedule implied by the constructor inputs.
Distributional (dict formula) -> distributional; mixture (mixture_components) -> EM; otherwise the
standard BaseTAM.fit (one exact solve for l2, else IRLS).
"""
if getattr(self, "_mode_", "plain") == "distributional":
return _distributional.fit(self, data_train, **schedule_kwargs)
if getattr(self, "_mixture_components_", None) is not None:
return _mixture.fit(self, data_train)
return super().fit(data_train)
The strategy contract and factory¶
A ReweightingStrategy is the statistical analogue of a BaseEffect: a small, swappable object that reshapes (z, W). build_strategy maps a user-facing loss name to a concrete strategy.
class ReweightingStrategy(ABC):
"""Supplies the working response z and weights w for the next P-WLS solve."""
is_glm: bool = False
name: str = "strategy"
def initial_eta(self, y: torch.Tensor) -> torch.Tensor:
"""Starting linear predictor. Default (identity link): the response itself."""
return y.clone()
@abstractmethod
def working_response_and_weights(
self, y: torch.Tensor, eta: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Return (z, w) given the current linear predictor eta."""
def inverse_link(self, eta: torch.Tensor) -> torch.Tensor:
"""Map the linear predictor to the mean scale (identity by default)."""
return eta
def mean_objective(self, y: torch.Tensor, eta: torch.Tensor) -> torch.Tensor:
"""Scalar objective for convergence / step-halving. Must stay a mean so it
scales with the penalized roughness term."""
residual = y - eta
return (residual * residual).mean()
def build_strategy(loss: str, tau: float = 0.5, nu: float = 4.0, delta: float = 1.345) -> ReweightingStrategy:
"""Map a loss name to a ReweightingStrategy."""
key = loss.lower().replace("-", "_")
if key in _STRATEGY_ALIASES:
return _STRATEGY_ALIASES[key]()
if key == "expectile":
return ExpectileLoss(tau)
if key == "huber":
return HuberLoss(delta)
if key in ("student_t", "studentt", "student"):
return StudentTLoss(nu)
raise ValueError(
f"Unknown loss '{loss}'. Supported: l2/gaussian, gamma, poisson, binomial, expectile, huber, student_t."
)