diff options
Diffstat (limited to 'sdil/rain_ep_adapter.py')
| -rw-r--r-- | sdil/rain_ep_adapter.py | 215 |
1 files changed, 215 insertions, 0 deletions
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py index ecdb5e5..9be1128 100644 --- a/sdil/rain_ep_adapter.py +++ b/sdil/rain_ep_adapter.py @@ -215,6 +215,221 @@ def attach_to_rain_estimator(estimator, corrector: RainGradientCorrector): return estimator +class DillavouUpdateCorrector: + """Fixed per-parameter update bias from Dillavou et al. Eq. (7). + + The corruption is added *after* the two-state EP estimator has formed its + local parameter update. It is therefore independent of the sign or + centering of beta. ``constant`` and ``innovation`` intentionally coincide + when ``drift_ratio`` is zero: under the paper's strictly constant bias + model, the Harnett-style predictor reduces to a local intercept estimate. + + ``drift_ratio`` is an optional state-dependent extension. It is kept + separate from the exact Dillavou model so reports cannot conflate them. + """ + + MODES = RainGradientCorrector.MODES + + def __init__( + self, + *, + mode: str, + bias_ratio: float, + predictor_rate: float = 1.0, + neutral_cadence: int = 1, + drift_ratio: float = 0.0, + seed: int = 1729, + ) -> None: + if mode not in self.MODES: + raise ValueError(f"unrecognized correction mode {mode}") + if bias_ratio < 0.0 or drift_ratio < 0.0: + raise ValueError("Dillavou bias ratios must be nonnegative") + if not 0.0 < predictor_rate <= 1.0: + raise ValueError("predictor rate must lie in (0, 1]") + if neutral_cadence < 0: + raise ValueError("neutral cadence must be nonnegative") + self.mode = mode + self.bias_ratio = bias_ratio + self.predictor_rate = predictor_rate + self.neutral_cadence = neutral_cadence + self.drift_ratio = drift_ratio + self.seed = seed + self.steps = 0 + self.debiaser: LocalAffineDebiaser | None = None + self._offsets: list[Tensor] | None = None + self._slopes: list[Tensor] | None = None + self._centers: list[Tensor] | None = None + self._scales: list[Tensor] | None = None + self._noise_generators: list[torch.Generator] | None = None + self.last_diagnostics: dict[str, float | int | str] = {} + + @staticmethod + @torch.no_grad() + def _unit_pattern(template: Tensor, phase: float) -> Tensor: + index = torch.arange( + template.numel(), dtype=template.dtype, device=template.device + ).reshape(template.shape) + pattern = torch.sin((index + phase) * 1.61803398875) + pattern.add_(0.5 * torch.cos((index + phase) * 0.754877666)) + return pattern / pattern.square().mean().sqrt().clamp_min(1e-30) + + @torch.no_grad() + def _initialize( + self, clean: list[Tensor], parameter_states: list[Tensor] + ) -> None: + if len(clean) != len(parameter_states): + raise ValueError("gradient and parameter collections disagree") + self._offsets = [] + self._slopes = [] + self._centers = [] + self._scales = [] + feature_scales = [] + for index, (gradient, parameter) in enumerate( + zip(clean, parameter_states) + ): + gradient_scale = gradient.square().mean().sqrt().clamp_min(1e-12) + parameter_scale = parameter.square().mean().sqrt().clamp_min(1e-6) + offset_pattern = self._unit_pattern( + gradient, float(self.seed + 97 * index)) + slope_pattern = self._unit_pattern( + gradient, float(self.seed + 193 * index + 41)) + self._offsets.append( + self.bias_ratio * gradient_scale * offset_pattern) + self._slopes.append( + self.drift_ratio * gradient_scale * slope_pattern) + self._centers.append(parameter.clone()) + self._scales.append(parameter_scale.clone()) + feature_scales.append(torch.ones_like(parameter_scale)) + self.debiaser = LocalAffineDebiaser( + clean, + feature_centers=[0.0] * len(clean), + feature_scales=feature_scales, + affine=self.mode == "innovation", + ) + + @torch.no_grad() + def _measure( + self, parameter_states: list[Tensor] + ) -> tuple[list[Tensor], list[Tensor]]: + features = [] + biases = [] + for parameter, center, scale, offset, slope in zip( + parameter_states, + self._centers, + self._scales, + self._offsets, + self._slopes, + ): + feature = torch.tanh((parameter - center) / scale) + features.append(feature) + biases.append(offset + slope * feature) + return features, biases + + @torch.no_grad() + def apply( + self, clean: Iterable[Tensor], parameter_states: Iterable[Tensor] + ) -> list[Tensor]: + clean = list(clean) + parameter_states = list(parameter_states) + if any(value.requires_grad for value in clean + parameter_states): + raise ValueError("Dillavou adapter received a requires-grad tensor") + if self.mode == "clean": + return [value.clone() for value in clean] + if self._offsets is None: + self._initialize(clean, parameter_states) + features, bias = self._measure(parameter_states) + measured = [ + value + corruption for value, corruption in zip(clean, bias) + ] + if self.mode == "raw": + corrected = measured + elif self.mode == "oracle": + corrected = [value.clone() for value in clean] + elif self.mode == "same_rms_noise": + if self._noise_generators is None: + self._noise_generators = [] + for index, template in enumerate(clean): + generator = torch.Generator(device=template.device) + generator.manual_seed(self.seed + 1009 * index) + self._noise_generators.append(generator) + corrected = [] + for value, corruption, generator in zip( + clean, bias, self._noise_generators + ): + noise = torch.randn( + value.shape, + dtype=value.dtype, + device=value.device, + generator=generator, + ) + noise.mul_(corruption.square().mean().sqrt()) + corrected.append(value + noise) + else: + if ( + self.neutral_cadence > 0 + and self.steps % self.neutral_cadence == 0 + ): + self.debiaser.update_neutral( + features, bias, self.predictor_rate) + corrected = self.debiaser.residual(features, measured) + residual_bias = [ + value - target for value, target in zip(corrected, clean) + ] + total_elements = sum(value.numel() for value in clean) + clean_square = sum(float(value.square().sum()) for value in clean) + bias_square = sum(float(value.square().sum()) for value in bias) + residual_square = sum( + float(value.square().sum()) for value in residual_bias) + clean_rms = (clean_square / total_elements) ** 0.5 + bias_rms = (bias_square / total_elements) ** 0.5 + residual_rms = (residual_square / total_elements) ** 0.5 + self.last_diagnostics = { + "step": self.steps, + "bias_model": ( + "dillavou_constant_update" + if self.drift_ratio == 0.0 + else "dillavou_plus_local_state_drift" + ), + "clean_update_rms": clean_rms, + "bias_update_rms": bias_rms, + "residual_update_rms": residual_rms, + "bias_to_clean_update_rms": ( + bias_rms / max(clean_rms, 1e-30)), + "residual_to_clean_update_rms": ( + residual_rms / max(clean_rms, 1e-30)), + "neutral_observations": self.debiaser.neutral_observations, + } + self.steps += 1 + return corrected + + +def attach_dillavou_to_rain_estimator( + estimator, corrector: DillavouUpdateCorrector +): + """Add the hardware update offset after Rain's local EP estimate.""" + + @torch.no_grad() + def corrected_standard_param_grads(self, layers_first, layers_second): + for layer in self._layers: + layer.state = layers_first[layer.name] + grads_first = [updater.grad() for updater in self._param_updaters] + for layer in self._layers: + layer.state = layers_second[layer.name] + grads_second = [updater.grad() for updater in self._param_updaters] + denominator = self._second_nudging - self._first_nudging + clean = [ + (second - first) / denominator + for first, second in zip(grads_first, grads_second) + ] + parameter_states = [parameter.state for parameter in self._params] + return corrector.apply(clean, parameter_states) + + estimator._standard_param_grads = MethodType( + corrected_standard_param_grads, estimator) + estimator.sdil_corrector = corrector + return estimator + + @torch.no_grad() def observe_rain_neutral(estimator, corrector: RainGradientCorrector) -> None: """Expose one label-free Rain equilibrium to the local predictor.""" |
