Reweighted Estimation: Strategy Pattern & IRLS Driver

Navigation:

Implementation only, for the working-response/weight formulas see the math theory.

The estimation/ package mirrors spectrum/: a base contract, one file per loss family, and a registry.

The GLM families

class GLMFamily(ReweightingStrategy):
    """A GLM defined by a link, a variance V(mu) and a unit deviance."""

    is_glm = True

    def __init__(
        self,
        link: Link,
        variance: Callable[[torch.Tensor], torch.Tensor],
        unit_deviance: Callable[[torch.Tensor, torch.Tensor], torch.Tensor],
        name: str,
    ):
        self.link = link
        self.variance = variance
        self.unit_deviance = unit_deviance
        self.name = name

    def initial_eta(self, y: torch.Tensor) -> torch.Tensor:
        return self.link.link(self.link.starting_mu(y))

    def inverse_link(self, eta: torch.Tensor) -> torch.Tensor:
        return self.link.inverse(eta)

    def working_response_and_weights(self, y, eta):
        mu = self.link.inverse(eta)
        g_prime = self.link.deriv(mu)
        working_response = eta + (y - mu) * g_prime
        weights = 1.0 / (g_prime * g_prime * self.variance(mu)).clamp_min(_TINY)
        return working_response, weights

    def mean_objective(self, y, eta):
        mu = self.link.inverse(eta)
        return self.unit_deviance(y, mu).mean()

Robust M-estimators and expectiles

class HuberLoss(ReweightingStrategy):
    """Huber: weight 1 within delta robust scales, else delta/|u| (bounded influence)."""

    name = "huber"

    def __init__(self, delta: float = 1.345):
        if delta <= 0:
            raise ValueError(f"huber delta must be positive; got {delta}")
        self.delta = float(delta)

    def working_response_and_weights(self, y, eta):
        residual = y - eta
        scale = _robust_scale(residual)
        standardized = (residual / scale).abs()
        weights = torch.where(standardized <= self.delta, torch.ones_like(standardized),
                              self.delta / standardized.clamp_min(_TINY))
        return y, weights
class ExpectileLoss(ReweightingStrategy):
    """Asymmetric least squares (Newey-Powell): weight tau above the fit, else 1-tau."""

    name = "expectile"

    def __init__(self, tau: float = 0.5):
        if not 0.0 < tau < 1.0:
            raise ValueError(f"expectile level tau must be in (0, 1); got {tau}")
        self.tau = float(tau)

    def working_response_and_weights(self, y, eta):
        residual = y - eta
        weights = torch.where(residual >= 0, torch.as_tensor(self.tau, dtype=y.dtype, device=y.device),
                              torch.as_tensor(1.0 - self.tau, dtype=y.dtype, device=y.device))
        return y, weights

    def mean_objective(self, y, eta):
        residual = y - eta
        asymmetric = torch.where(residual >= 0, self.tau, 1.0 - self.tau)
        return (asymmetric * residual * residual).mean()

The IRLS schedule

reweighted_penalized_fit normalises the weights to unit mean (preserving the nS penalty scale), solves one atom, and, for GLMs, step-halves on the penalized objective. For the Gaussian identity it converges in a single step.

def reweighted_penalized_fit(
    model,
    x_data: torch.Tensor,
    y_data: torch.Tensor,
    strategy: ReweightingStrategy,
    max_iter: int = 25,
    tol: float = 1e-6,
    max_halvings: int = 8,
) -> torch.Tensor:
    """Fit a strategy by iteratively reweighted penalized least squares; return the coefficients.

    Args:
        model: a TAM exposing _solve_pwls_step and _build_design_matrix.
        x_data: feature tensor (n_groups, n_samples, n_features).
        y_data: response tensor (n_groups, n_samples, 1).
        strategy: the reweighting rule producing (z, w) per iteration.
        max_iter: maximum outer iterations.
        tol: relative change in eta for convergence.
        max_halvings: maximum step-halvings per iteration (GLM divergence guard).
    """
    y_data = y_data.to(torch.get_default_dtype())
    eta = strategy.initial_eta(y_data)
    theta = None

    penalty_matrix = model._build_penalty_matrix()

    def penalized_objective(theta_value: torch.Tensor, eta_value: torch.Tensor) -> torch.Tensor:
        # mean deviance + roughness penalty theta.T * S * theta. Step-halving must judge the
        # penalized objective, else a step that lowers deviance but inflates roughness is accepted.
        base = strategy.mean_objective(y_data, eta_value)
        penalty_quadratic = theta_value.mT @ penalty_matrix.to(theta_value.dtype) @ theta_value
        return base + penalty_quadratic.mean()

    for _ in range(max_iter):
        working_response, weights = strategy.working_response_and_weights(y_data, eta)
        weights = weights / weights.mean().clamp_min(_TINY)
        theta_candidate = model._solve_pwls_step(x_data, working_response, weights=weights)
        eta_candidate = _linear_predictor(model, x_data, theta_candidate)

        if theta is not None and strategy.is_glm:
            theta_full_step = theta_candidate
            objective_previous = penalized_objective(theta, eta)
            objective_candidate = penalized_objective(theta_candidate, eta_candidate)
            step = 1.0
            halving = 0
            while (not torch.isfinite(objective_candidate) or objective_candidate > objective_previous) \
                    and halving < max_halvings:
                step *= 0.5
                theta_candidate = theta + step * (theta_full_step - theta)
                eta_candidate = _linear_predictor(model, x_data, theta_candidate)
                objective_candidate = penalized_objective(theta_candidate, eta_candidate)
                halving += 1

        relative_change = (eta_candidate - eta).norm() / eta.norm().clamp_min(_TINY)
        theta = theta_candidate
        eta = eta_candidate
        if relative_change < tol:
            break

    return theta