The Statistics API: One Atom, Many Schedules

Navigation:

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."
    )