diff options
Diffstat (limited to 'sdil/rain_ep_adapter.py')
| -rw-r--r-- | sdil/rain_ep_adapter.py | 186 |
1 files changed, 185 insertions, 1 deletions
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py index 27314bd..73cd70d 100644 --- a/sdil/rain_ep_adapter.py +++ b/sdil/rain_ep_adapter.py @@ -13,7 +13,10 @@ from typing import Iterable import torch -from sdil.two_state_debias import LocalAffineDebiaser +from sdil.two_state_debias import ( + BatchedLocalAffineDebiaser, + LocalAffineDebiaser, +) Tensor = torch.Tensor @@ -217,3 +220,184 @@ def observe_rain_neutral(estimator, corrector: RainGradientCorrector) -> None: """Expose one label-free Rain equilibrium to the local predictor.""" local_states = [updater.grad() for updater in estimator._param_updaters] corrector.observe_neutral(local_states) + + +class RainLayerStateCorrector: + """Correct a structured per-neuron bias before Rain's local EP update.""" + + MODES = RainGradientCorrector.MODES + + def __init__( + self, + *, + mode: str, + bias_ratio: float, + predictor_rate: float = 0.2, + calibration_steps: int = 1, + seed: int = 1729, + ) -> None: + if mode not in self.MODES: + raise ValueError(f"unrecognized correction mode {mode}") + if bias_ratio < 0.0: + raise ValueError("bias ratio must be nonnegative") + if calibration_steps < 0: + raise ValueError("calibration steps must be nonnegative") + self.mode = mode + self.bias_ratio = bias_ratio + self.predictor_rate = predictor_rate + self.calibration_steps = calibration_steps + self.seed = seed + self.steps = 0 + self.last_diagnostics: dict[str, float | int] = {} + self.debiaser: BatchedLocalAffineDebiaser | None = None + self._metadata: list[tuple[Tensor, Tensor, Tensor]] | None = None + self._noise_generators: list[torch.Generator] | None = None + + @torch.no_grad() + def _initialize( + self, first_states: list[Tensor], clean_differences: list[Tensor] + ) -> None: + self._metadata = [] + for index, (first, clean) in enumerate(zip( + first_states, clean_differences + )): + scale = first.square().mean(dim=0).sqrt().clamp_min(1e-6) + flat_index = torch.arange( + scale.numel(), dtype=first.dtype, device=first.device + ).reshape(scale.shape) + phase = flat_index + float(self.seed + 97 * index) + offset = torch.where( + torch.remainder(phase, 2.0) < 1.0, + torch.full_like(scale, -0.5), + torch.full_like(scale, 0.5), + ) + slope = 0.5 + torch.remainder(phase * 0.61803398875, 1.0) + basis = torch.tanh(first / scale) + source = offset.unsqueeze(0) + slope.unsqueeze(0) * basis + source_rms = source.square().mean().sqrt().clamp_min(1e-30) + clean_rms = clean.square().mean().sqrt() + gain = self.bias_ratio * clean_rms / source_rms + self._metadata.append( + (scale, gain * offset, gain * slope) + ) + self.debiaser = BatchedLocalAffineDebiaser( + first_states, + feature_centers=[0.0] * len(first_states), + feature_scales=[1.0] * len(first_states), + affine=self.mode == "innovation", + ) + + @torch.no_grad() + def _measure( + self, first_states: list[Tensor] + ) -> tuple[list[Tensor], list[Tensor]]: + bases = [] + biases = [] + for first, (scale, offset, slope) in zip( + first_states, self._metadata + ): + basis = torch.tanh(first / scale) + bases.append(basis) + biases.append( + offset.unsqueeze(0) + slope.unsqueeze(0) * basis) + return bases, biases + + @torch.no_grad() + def apply( + self, + layers_first: dict[str, Tensor], + layers_second: dict[str, Tensor], + layer_names: list[str], + ) -> dict[str, Tensor]: + if self.mode == "clean": + return dict(layers_second) + first = [layers_first[name] for name in layer_names] + second = [layers_second[name] for name in layer_names] + if any(value.requires_grad for value in first + second): + raise ValueError("Rain layer adapter received a requires-grad tensor") + clean = [after - before for before, after in zip(first, second)] + if self._metadata is None: + self._initialize(first, clean) + bases, bias = self._measure(first) + 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.steps < self.calibration_steps: + self.debiaser.update_neutral( + bases, bias, self.predictor_rate) + corrected = self.debiaser.residual(bases, 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, + "clean_state_difference_rms": clean_rms, + "bias_state_difference_rms": bias_rms, + "residual_state_difference_rms": residual_rms, + "bias_to_clean_state_difference_rms": ( + bias_rms / max(clean_rms, 1e-30)), + "residual_to_clean_state_difference_rms": ( + residual_rms / max(clean_rms, 1e-30)), + "neutral_observations": ( + 0 if self.debiaser is None else self.debiaser.neutral_observations), + } + used_second = dict(layers_second) + for name, before, difference in zip(layer_names, first, corrected): + used_second[name] = before + difference + self.steps += 1 + return used_second + + +def attach_layer_to_rain_estimator( + estimator, corrector: RainLayerStateCorrector +): + """Patch Rain immediately before its hand-written local parameter rule.""" + layer_names = [layer.name for layer in estimator._layers[1:]] + + @torch.no_grad() + def corrected_standard_param_grads(self, layers_first, layers_second): + used_second = corrector.apply( + layers_first, layers_second, layer_names) + 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 = used_second[layer.name] + grads_second = [updater.grad() for updater in self._param_updaters] + denominator = self._second_nudging - self._first_nudging + return [ + (second - first) / denominator + for first, second in zip(grads_first, grads_second) + ] + + estimator._standard_param_grads = MethodType( + corrected_standard_param_grads, estimator) + estimator.sdil_corrector = corrector + return estimator |
