From 5d7adb2604f9f1e083523e71880533ac750c1cc4 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 16:34:53 -0500 Subject: feat: add shared no-grad two-state debiaser --- sdil/two_state_debias.py | 181 +++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 181 insertions(+) create mode 100644 sdil/two_state_debias.py (limited to 'sdil') diff --git a/sdil/two_state_debias.py b/sdil/two_state_debias.py new file mode 100644 index 0000000..fad980a --- /dev/null +++ b/sdil/two_state_debias.py @@ -0,0 +1,181 @@ +"""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 + -- cgit v1.2.3