summaryrefslogtreecommitdiff
path: root/sdil/two_state_debias.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:34:53 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:34:53 -0500
commit5d7adb2604f9f1e083523e71880533ac750c1cc4 (patch)
tree062f33f3edbe0f42a5a3038f0c5156701fcae623 /sdil/two_state_debias.py
parent37f002930e11b7566a4397ecf1f385c66ee00b6c (diff)
feat: add shared no-grad two-state debiaser
Diffstat (limited to 'sdil/two_state_debias.py')
-rw-r--r--sdil/two_state_debias.py181
1 files changed, 181 insertions, 0 deletions
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
+