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