"""Shared no-autograd debiasing primitive for two-state local learners.""" from __future__ import annotations from dataclasses import dataclass from typing import Iterable import torch Tensor = torch.Tensor def _detached(tensors: Iterable[Tensor]) -> list[Tensor]: values = list(tensors) if any(value.requires_grad for value in values): raise ValueError("paper-facing local tensors must not require gradients") return values @torch.no_grad() def two_state_difference( first: Iterable[Tensor], second: Iterable[Tensor], denominator: float, ) -> list[Tensor]: """Form a two-state teaching measurement with common-mode cancellation.""" first = _detached(first) second = _detached(second) if len(first) != len(second): raise ValueError("state collections have different lengths") if denominator == 0.0: raise ValueError("state-difference denominator must be nonzero") return [ (value_second - value_first) / denominator for value_first, value_second in zip(first, second) ] @dataclass class LocalFilterState: intercept: Tensor slope: Tensor | None feature_center: Tensor feature_scale: Tensor class LocalAffineDebiaser: """Per-element affine filter trained only by normalized local LMS. The class is intentionally not a ``torch.nn.Module``. Coefficients are ordinary detached tensors, and every operation executes under ``torch.no_grad``. """ def __init__( self, templates: Iterable[Tensor], *, feature_centers: Iterable[Tensor | float], feature_scales: Iterable[Tensor | float], affine: bool = True, ) -> None: templates = _detached(templates) centers = list(feature_centers) scales = list(feature_scales) if not (len(templates) == len(centers) == len(scales)): raise ValueError("template and feature metadata lengths disagree") self.affine = affine self.states = [] with torch.no_grad(): for template, center, scale in zip(templates, centers, scales): center_tensor = torch.as_tensor( center, dtype=template.dtype, device=template.device) scale_tensor = torch.as_tensor( scale, dtype=template.dtype, device=template.device) if torch.any(scale_tensor <= 0): raise ValueError("feature scales must be positive") self.states.append(LocalFilterState( intercept=torch.zeros_like(template), slope=torch.zeros_like(template) if affine else None, feature_center=center_tensor.clone(), feature_scale=scale_tensor.clone(), )) self.neutral_observations = 0 @staticmethod def _feature(state: LocalFilterState, local_feature: Tensor) -> Tensor: return (local_feature - state.feature_center) / state.feature_scale @torch.no_grad() def predict(self, local_features: Iterable[Tensor]) -> list[Tensor]: features = _detached(local_features) if len(features) != len(self.states): raise ValueError("feature collection length changed") predictions = [] for state, feature in zip(self.states, features): normalized = self._feature(state, feature) prediction = state.intercept if self.affine: prediction = prediction + state.slope * normalized predictions.append(prediction.clone()) return predictions @torch.no_grad() def update_neutral( self, local_features: Iterable[Tensor], neutral_measurements: Iterable[Tensor], learning_rate: float, ) -> list[Tensor]: features = _detached(local_features) measurements = _detached(neutral_measurements) if not (len(features) == len(measurements) == len(self.states)): raise ValueError("neutral tuple lengths disagree") residuals = [] for state, feature, measurement in zip( self.states, features, measurements ): normalized = self._feature(state, feature) prediction = state.intercept normalization = torch.ones_like(normalized) if self.affine: prediction = prediction + state.slope * normalized normalization = normalization + normalized.square() residual = measurement - prediction state.intercept.add_(learning_rate * residual / normalization) if self.affine: state.slope.add_( learning_rate * residual * normalized / normalization) residuals.append(residual.clone()) self.neutral_observations += 1 return residuals @torch.no_grad() def residual( self, local_features: Iterable[Tensor], teaching_measurements: Iterable[Tensor], ) -> list[Tensor]: measurements = _detached(teaching_measurements) predictions = self.predict(local_features) if len(measurements) != len(predictions): raise ValueError("teaching tuple lengths disagree") return [ measurement - prediction for measurement, prediction in zip(measurements, predictions) ] @torch.no_grad() def replay_updates( self, local_features: Iterable[Tensor], teaching_measurements: Iterable[Tensor], eligibilities: Iterable[Tensor], learning_rate: float, ) -> list[Tensor]: residuals = self.residual(local_features, teaching_measurements) eligibilities = _detached(eligibilities) if len(residuals) != len(eligibilities): raise ValueError("eligibility tuple length changed") return [ learning_rate * residual * eligibility for residual, eligibility in zip(residuals, eligibilities) ] @torch.no_grad() def clone(self) -> "LocalAffineDebiaser": clone = LocalAffineDebiaser( [state.intercept for state in self.states], feature_centers=[state.feature_center for state in self.states], feature_scales=[state.feature_scale for state in self.states], affine=self.affine, ) for source, target in zip(self.states, clone.states): target.intercept.copy_(source.intercept) if self.affine: target.slope.copy_(source.slope) clone.neutral_observations = self.neutral_observations return clone @dataclass class LocalLeastSquaresState: intercept: Tensor slope: Tensor feature_center: Tensor feature_scale: Tensor mean_feature: Tensor mean_measurement: Tensor feature_sum_squares: Tensor cross_sum: Tensor class LocalLeastSquaresDebiaser: """Per-element online affine regression from neutral observations. Every element keeps its own scalar sufficient statistics. No observation, coefficient, or update is shared across elements, and no autograd graph is created. This is the direct online analogue of fitting a local baseline relation before taking its innovation. """ def __init__( self, templates: Iterable[Tensor], *, feature_centers: Iterable[Tensor | float], feature_scales: Iterable[Tensor | float], variance_floor: float = 1e-12, ) -> None: templates = _detached(templates) centers = list(feature_centers) scales = list(feature_scales) if not (len(templates) == len(centers) == len(scales)): raise ValueError("template and feature metadata lengths disagree") if variance_floor <= 0.0: raise ValueError("variance floor must be positive") self.variance_floor = variance_floor self.states = [] with torch.no_grad(): for template, center, scale in zip(templates, centers, scales): center_tensor = torch.as_tensor( center, dtype=template.dtype, device=template.device) scale_tensor = torch.as_tensor( scale, dtype=template.dtype, device=template.device) if torch.any(scale_tensor <= 0): raise ValueError("feature scales must be positive") self.states.append(LocalLeastSquaresState( intercept=torch.zeros_like(template), slope=torch.zeros_like(template), feature_center=center_tensor.clone(), feature_scale=scale_tensor.clone(), mean_feature=torch.zeros_like(template), mean_measurement=torch.zeros_like(template), feature_sum_squares=torch.zeros_like(template), cross_sum=torch.zeros_like(template), )) self.neutral_observations = 0 @staticmethod def _feature( state: LocalLeastSquaresState, local_feature: Tensor ) -> Tensor: return (local_feature - state.feature_center) / state.feature_scale @torch.no_grad() def predict(self, local_features: Iterable[Tensor]) -> list[Tensor]: features = _detached(local_features) if len(features) != len(self.states): raise ValueError("feature collection length changed") return [ ( state.intercept + state.slope * self._feature(state, feature) ).clone() for state, feature in zip(self.states, features) ] @torch.no_grad() def update_neutral( self, local_features: Iterable[Tensor], neutral_measurements: Iterable[Tensor], learning_rate: float, ) -> list[Tensor]: if learning_rate != 1.0: raise ValueError( "online least squares requires predictor_rate=1") features = _detached(local_features) measurements = _detached(neutral_measurements) if not (len(features) == len(measurements) == len(self.states)): raise ValueError("neutral tuple lengths disagree") count = self.neutral_observations + 1 residuals = [] for state, feature, measurement in zip( self.states, features, measurements ): normalized = self._feature(state, feature) residuals.append( measurement - state.intercept - state.slope * normalized) delta_feature = normalized - state.mean_feature delta_measurement = measurement - state.mean_measurement state.mean_feature.add_(delta_feature / count) state.mean_measurement.add_(delta_measurement / count) state.feature_sum_squares.add_( delta_feature * (normalized - state.mean_feature)) state.cross_sum.add_( delta_feature * (measurement - state.mean_measurement)) identifiable = state.feature_sum_squares > self.variance_floor state.slope.copy_(torch.where( identifiable, state.cross_sum / state.feature_sum_squares.clamp_min(self.variance_floor), torch.zeros_like(state.slope), )) state.intercept.copy_( state.mean_measurement - state.slope * state.mean_feature) self.neutral_observations = count return residuals @torch.no_grad() def residual( self, local_features: Iterable[Tensor], teaching_measurements: Iterable[Tensor], ) -> list[Tensor]: measurements = _detached(teaching_measurements) predictions = self.predict(local_features) if len(measurements) != len(predictions): raise ValueError("teaching tuple lengths disagree") return [ measurement - prediction for measurement, prediction in zip(measurements, predictions) ] class BatchedLocalAffineDebiaser: """Per-cell LMS whose coefficients are shared across observations. Inputs have shape ``[observations, *local_shape]``. The leading axis is reduced only inside each cell's predictor update; cells remain independent. """ def __init__( self, templates: Iterable[Tensor], *, feature_centers: Iterable[Tensor | float], feature_scales: Iterable[Tensor | float], affine: bool = True, ) -> None: templates = _detached(templates) if any(template.ndim < 1 for template in templates): raise ValueError("batched templates require an observation axis") self.filter = LocalAffineDebiaser( [template[0] for template in templates], feature_centers=feature_centers, feature_scales=feature_scales, affine=affine, ) self.affine = affine self.neutral_observations = 0 @torch.no_grad() def predict(self, local_features: Iterable[Tensor]) -> list[Tensor]: features = _detached(local_features) if len(features) != len(self.filter.states): raise ValueError("feature collection length changed") predictions = [] for state, feature in zip(self.filter.states, features): normalized = self.filter._feature(state, feature) prediction = state.intercept.unsqueeze(0).expand_as(feature) if self.affine: prediction = prediction + state.slope.unsqueeze(0) * normalized predictions.append(prediction.clone()) return predictions @torch.no_grad() def update_neutral( self, local_features: Iterable[Tensor], neutral_measurements: Iterable[Tensor], learning_rate: float, ) -> None: features = _detached(local_features) measurements = _detached(neutral_measurements) if not (len(features) == len(measurements) == len(self.filter.states)): raise ValueError("neutral tuple lengths disagree") batch_sizes = {feature.shape[0] for feature in features} if len(batch_sizes) != 1: raise ValueError("local populations have different observation counts") if any( feature.shape != measurement.shape for feature, measurement in zip(features, measurements) ): raise ValueError("feature and measurement shapes disagree") observations = next(iter(batch_sizes)) for index in range(observations): self.filter.update_neutral( [feature[index] for feature in features], [measurement[index] for measurement in measurements], learning_rate, ) self.neutral_observations += observations @torch.no_grad() def residual( self, local_features: Iterable[Tensor], teaching_measurements: Iterable[Tensor], ) -> list[Tensor]: measurements = _detached(teaching_measurements) predictions = self.predict(local_features) if len(measurements) != len(predictions): raise ValueError("teaching tuple lengths disagree") return [ measurement - prediction for measurement, prediction in zip(measurements, predictions) ]