Reweighted Estimation: Strategy Pattern & IRLS Driver¶
Navigation:
Theory introduction: See the Intro
Related mathematical theory: Reweighted estimation
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