diff options
| -rw-r--r-- | experiments/transformer_feedback_smoke.py | 116 | ||||
| -rw-r--r-- | sdil/transformer.py | 285 |
2 files changed, 399 insertions, 2 deletions
diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py index 2cd22d9..ff34369 100644 --- a/experiments/transformer_feedback_smoke.py +++ b/experiments/transformer_feedback_smoke.py @@ -395,6 +395,120 @@ def audit_forward_forward(): } +def audit_dualprop(): + config = LocalTransformerConfig( + vocab_size=7, context_length=4, depth=2, width=8, heads=2, + mlp_ratio=2, seed=161) + bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) + dp = LocalDecoderTransformer(config, "dualprop", dtype=torch.float64) + generator = torch.Generator().manual_seed(162) + tokens = torch.randint(0, 7, (2, 4), generator=generator) + targets = torch.randint(0, 7, (2, 4), generator=generator) + with torch.no_grad(): + forward_error = float(torch.max(torch.abs( + bp(tokens)["logits"] - dp(tokens)["logits"]))) + plus, minus = dp.dualprop_states( + tokens, targets, alpha=0.0, beta=0.1, inference_passes=2) + metrics = dp.dualprop_gradients( + tokens, targets, alpha=0.0, beta=0.1, inference_passes=2) + actual = [ + parameter.grad.detach().clone() + for parameter in dp.blocks[0].parameters()] + missing = [ + name for name, parameter in dp.named_parameters() + if parameter.grad is None] + nonfinite = [ + name for name, parameter in dp.named_parameters() + if parameter.grad is not None + and not torch.isfinite(parameter.grad).all()] + states_finite = all( + torch.isfinite(state).all() for state in plus + minus) + + # Audit one complete block direction against its explicit detached local + # contrastive objective. + states = [negative.detach() for negative in minus[:-1]] + delta = ((plus[1] - minus[1]) / 0.1).detach() + parameters = tuple(dp.blocks[0].parameters()) + objective = -torch.sum( + dp.blocks[0](states[0]) * delta) / targets.numel() + expected = torch.autograd.grad(objective, parameters) + equation_error = max( + _relative_error(a, e) for a, e in zip(actual, expected)) + if (forward_error >= 1e-12 or equation_error >= 1e-12 + or missing or nonfinite or not states_finite): + raise AssertionError({ + "forward_error": forward_error, + "equation_error": equation_error, + "missing": missing, + "nonfinite": nonfinite, + "states_finite": states_finite, + }) + return { + "matched_forward_max_absolute_error": forward_error, + "local_contrastive_equation_max_relative_error": equation_error, + "inferred_state_tensors": len(plus) + len(minus), + "finite_inferred_states": states_finite, + "local_vjp_evaluations": metrics["local_vjp_evaluations"], + "uses_symmetric_block_vjps": True, + "uses_task_loss_backward": metrics["uses_task_loss_backward"], + } + + +def audit_equilibrium_propagation(): + config = LocalTransformerConfig( + vocab_size=7, context_length=4, depth=2, width=8, heads=2, + mlp_ratio=2, seed=171) + bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) + ep = LocalDecoderTransformer(config, "ep", dtype=torch.float64) + generator = torch.Generator().manual_seed(172) + tokens = torch.randint(0, 7, (2, 4), generator=generator) + targets = torch.randint(0, 7, (2, 4), generator=generator) + with torch.no_grad(): + forward_error = float(torch.max(torch.abs( + bp(tokens)["logits"] - ep(tokens)["logits"]))) + metrics, free, nudged = ep.ep_gradients( + tokens, targets, ep_beta=0.5, dt=0.5, + free_steps=3, nudge_steps=2, beta_sign=1.0) + missing = [ + name for name, parameter in ep.named_parameters() + if parameter.grad is None] + nonfinite = [ + name for name, parameter in ep.named_parameters() + if parameter.grad is not None + and not torch.isfinite(parameter.grad).all()] + states = free + nudged + lower = min(float(state.min()) for state in states) + upper = max(float(state.max()) for state in states) + + free_groups = ep._ep_phase_gradients(tokens, free) + nudged_groups = ep._ep_phase_gradients(tokens, nudged) + expected = [ + -(nudged_gradient - free_gradient) / 0.5 + for nudged_gradient, free_gradient in zip( + nudged_groups[1][1], free_groups[1][1])] + actual = [ + parameter.grad for parameter in free_groups[1][0]] + equation_error = max( + _relative_error(a, e) for a, e in zip(actual, expected)) + if (forward_error >= 1e-12 or equation_error >= 1e-12 + or missing or nonfinite or lower < 0.0 or upper > 1.0): + raise AssertionError({ + "forward_error": forward_error, + "equation_error": equation_error, + "missing": missing, + "nonfinite": nonfinite, + "state_range": [lower, upper], + }) + return { + "matched_forward_max_absolute_error": forward_error, + "phase_contrast_equation_max_relative_error": equation_error, + "state_range": [lower, upper], + "local_vjp_evaluations": metrics["local_vjp_evaluations"], + "uses_symmetric_block_vjps": True, + "uses_task_loss_backward": metrics["uses_task_loss_backward"], + } + + def main(): torch.set_num_threads(1) result = { @@ -406,6 +520,8 @@ def main(): "strict_dfa": audit_strict_dfa(), "pepita": audit_pepita(), "forward_forward": audit_forward_forward(), + "dual_propagation": audit_dualprop(), + "equilibrium_propagation": audit_equilibrium_propagation(), } print(json.dumps(result, indent=2, sort_keys=True)) diff --git a/sdil/transformer.py b/sdil/transformer.py index d47fd18..8934d8a 100644 --- a/sdil/transformer.py +++ b/sdil/transformer.py @@ -15,7 +15,8 @@ import torch.nn.functional as F _FEEDBACK_METHODS = {"fa", "clean_kp", "sdil"} -_SUPPORTED_METHODS = {"bp", "dfa", "pepita", "ff"} | _FEEDBACK_METHODS +_SUPPORTED_METHODS = { + "bp", "dfa", "pepita", "ff", "dualprop", "ep"} | _FEEDBACK_METHODS class _FeedbackLinearFunction(torch.autograd.Function): @@ -126,7 +127,8 @@ class FeedbackLinear(nn.Module): self.register_buffer("last_traffic_rms", torch.zeros((), dtype=dtype)) def forward(self, x): - if self.method in {"bp", "dfa", "pepita", "ff"}: + if self.method in { + "bp", "dfa", "pepita", "ff", "dualprop", "ep"}: return F.linear(x, self.weight, self.bias) method_code = {"fa": 0, "clean_kp": 1, "sdil": 2}[self.method] return _FeedbackLinearFunction.apply( @@ -696,3 +698,282 @@ class LocalDecoderTransformer(nn.Module): for output in layer_outputs] scores.append(sum(goodness[score_from_layer:])) return torch.stack(scores, dim=-1) + + def _embedding_prediction(self, tokens): + return ( + self.token_embedding(tokens) + + self.position_embedding[:tokens.shape[1]]) + + @staticmethod + def _local_vjp(function, state, field): + """Evaluate one explicitly local state VJP at a detached boundary.""" + with torch.enable_grad(): + local_state = state.detach().requires_grad_(True) + prediction = function(local_state) + objective = torch.sum(prediction * field.detach()) + result, = torch.autograd.grad(objective, local_state) + return result.detach() + + def _head_prediction(self, state): + return self.head(self.final_norm(state)) + + def dualprop_states( + self, tokens, targets, alpha=0.0, beta=0.1, + inference_passes=16, clean=None): + """Infer author-style plus/minus states over Transformer block edges.""" + if self.method != "dualprop": + raise ValueError( + "dualprop_states is only valid for method='dualprop'") + if not 0.0 <= alpha <= 1.0 or beta <= 0 or inference_passes < 1: + raise ValueError("invalid Dual Propagation settings") + if clean is None: + with torch.no_grad(): + clean = self.forward(tokens, return_cache=True) + hidden_clean = ( + [clean["embedded"]] + list(clean["block_outputs"])) + plus = [state.detach().clone() for state in hidden_clean] + minus = [state.detach().clone() for state in hidden_clean] + plus.append(clean["logits"].detach().clone()) + minus.append(clean["logits"].detach().clone()) + one_hot = F.one_hot( + targets, self.config.vocab_size).to(clean["logits"].dtype) + fixed_error = torch.softmax( + clean["logits"].detach(), dim=-1) - one_hot + + for _ in range(inference_passes): + for index in range(self.config.depth + 1): + states = [ + alpha * positive + (1.0 - alpha) * negative + for positive, negative in zip(plus[:-1], minus[:-1])] + if index == 0: + prediction = self._embedding_prediction(tokens) + else: + prediction = self.blocks[index - 1]( + states[index - 1].detach()) + if index < self.config.depth: + feedback = self._local_vjp( + self.blocks[index], states[index], + plus[index + 1] - minus[index + 1]) + else: + feedback = self._local_vjp( + self._head_prediction, states[index], + plus[-1] - minus[-1]) + plus[index] = ( + prediction + (1.0 - alpha) * feedback).detach() + minus[index] = ( + prediction - alpha * feedback).detach() + + states = [ + alpha * positive + (1.0 - alpha) * negative + for positive, negative in zip(plus[:-1], minus[:-1])] + prediction = self._head_prediction(states[-1].detach()) + output_field = beta * fixed_error + plus[-1] = ( + prediction - (1.0 - alpha) * output_field).detach() + minus[-1] = ( + prediction + alpha * output_field).detach() + return plus, minus + + def dualprop_gradients( + self, tokens, targets, alpha=0.0, beta=0.1, + inference_passes=16): + """Populate local Dual-Propagation contrastive gradients.""" + if beta <= 0: + raise ValueError("Dual Propagation beta must be positive") + with torch.no_grad(): + clean = self.forward(tokens, return_cache=True) + plus, minus = self.dualprop_states( + tokens, targets, alpha=alpha, beta=beta, + inference_passes=inference_passes, clean=clean) + states = [ + (alpha * positive + (1.0 - alpha) * negative).detach() + for positive, negative in zip(plus[:-1], minus[:-1])] + deltas = [ + ((positive - negative) / beta).detach() + for positive, negative in zip(plus, minus)] + for parameter in self.parameters(): + parameter.grad = None + observations = targets.numel() + + embedding_parameters = ( + self.token_embedding.weight, self.position_embedding) + embedding_objective = -torch.sum( + self._embedding_prediction(tokens) * deltas[0]) / observations + embedding_gradients = torch.autograd.grad( + embedding_objective, embedding_parameters) + self._assign_gradients( + embedding_parameters, embedding_gradients) + + for index, block in enumerate(self.blocks): + prediction = block(states[index]) + parameters = tuple(block.parameters()) + objective = -torch.sum( + prediction * deltas[index + 1]) / observations + gradients = torch.autograd.grad(objective, parameters) + self._assign_gradients(parameters, gradients) + + output_prediction = self._head_prediction(states[-1]) + output_parameters = ( + *tuple(self.final_norm.parameters()), + *tuple(self.head.parameters())) + output_objective = -torch.sum( + output_prediction * deltas[-1]) / observations + output_gradients = torch.autograd.grad( + output_objective, output_parameters) + self._assign_gradients(output_parameters, output_gradients) + return { + "clean_loss": float(F.cross_entropy( + clean["logits"].reshape(-1, self.config.vocab_size), + targets.reshape(-1))), + "state_difference_rms": float(torch.stack([ + (positive - negative).square().mean().sqrt() + for positive, negative in zip(plus, minus)]).mean()), + "inference_passes": int(inference_passes), + "local_vjp_evaluations": + int((self.config.depth + 1) * inference_passes), + "uses_task_loss_backward": False, + } + + @staticmethod + def ep_rho(state): + return state.clamp(0.0, 1.0) + + @staticmethod + def ep_rhop(state): + return ((state >= 0.0) & (state <= 1.0)).to(state.dtype) + + def ep_settle( + self, tokens, targets, beta=0.0, steps=20, dt=0.5, + initial_states=None): + """Settle synchronous hard-sigmoid states in a symmetric block graph.""" + if self.method != "ep": + raise ValueError("ep_settle is only valid for method='ep'") + if steps < 1 or not 0.0 < dt <= 1.0: + raise ValueError("invalid Equilibrium Propagation dynamics") + with torch.no_grad(): + clean = self.forward(tokens, return_cache=True) + if initial_states is None: + hidden_clean = ( + [clean["embedded"]] + list(clean["block_outputs"])) + states = [torch.zeros_like(state) for state in hidden_clean] + states.append(torch.zeros_like(clean["logits"])) + else: + states = [state.detach().clone() for state in initial_states] + one_hot = F.one_hot( + targets, self.config.vocab_size).to(clean["logits"].dtype) + + for _ in range(steps): + rho_hidden = [ + self.ep_rho(state).detach() for state in states[:-1]] + rho_output = self.ep_rho(states[-1]).detach() + updated = [] + for index, state in enumerate(states[:-1]): + if index == 0: + prediction = self._embedding_prediction(tokens).detach() + else: + prediction = self.blocks[index - 1]( + rho_hidden[index - 1]).detach() + if index < self.config.depth: + feedback = self._local_vjp( + self.blocks[index], rho_hidden[index], + rho_hidden[index + 1]) + else: + feedback = self._local_vjp( + self._head_prediction, rho_hidden[index], + rho_output) + drive = -self.ep_rho(state) + prediction + feedback + new_state = ( + state + dt * self.ep_rhop(state) * drive) + updated.append(self.ep_rho(new_state).detach()) + + output_prediction = self._head_prediction( + rho_hidden[-1]).detach() + output_state = states[-1] + output_drive = ( + -self.ep_rho(output_state) + output_prediction) + if beta: + output_drive = output_drive + 2.0 * beta * ( + one_hot - self.ep_rho(output_state)) + new_output = ( + output_state + + dt * self.ep_rhop(output_state) * output_drive) + updated.append(self.ep_rho(new_output).detach()) + states = updated + return states + + def _ep_phase_gradients(self, tokens, states): + rho_hidden = [ + self.ep_rho(state).detach() for state in states[:-1]] + rho_output = self.ep_rho(states[-1]).detach() + observations = tokens.numel() + groups = [] + + embedding_parameters = ( + self.token_embedding.weight, self.position_embedding) + embedding_correlation = torch.sum( + self._embedding_prediction(tokens) + * rho_hidden[0]) / observations + groups.append(( + embedding_parameters, + torch.autograd.grad( + embedding_correlation, embedding_parameters))) + + for index, block in enumerate(self.blocks): + parameters = tuple(block.parameters()) + correlation = torch.sum( + block(rho_hidden[index]) * rho_hidden[index + 1] + ) / observations + groups.append(( + parameters, torch.autograd.grad(correlation, parameters))) + + output_parameters = ( + *tuple(self.final_norm.parameters()), + *tuple(self.head.parameters())) + output_correlation = torch.sum( + self._head_prediction(rho_hidden[-1]) + * rho_output) / observations + groups.append(( + output_parameters, + torch.autograd.grad(output_correlation, output_parameters))) + return groups + + def ep_gradients( + self, tokens, targets, ep_beta=0.5, dt=0.5, + free_steps=20, nudge_steps=4, beta_sign=1.0): + """Populate the free-versus-nudged EP contrastive gradients.""" + signed_beta = float(beta_sign) * float(ep_beta) + if signed_beta == 0: + raise ValueError("EP contrast requires a nonzero nudge") + free = self.ep_settle( + tokens, targets, beta=0.0, steps=free_steps, dt=dt) + nudged = self.ep_settle( + tokens, targets, beta=signed_beta, steps=nudge_steps, dt=dt, + initial_states=free) + free_groups = self._ep_phase_gradients(tokens, free) + nudged_groups = self._ep_phase_gradients(tokens, nudged) + for parameter in self.parameters(): + parameter.grad = None + for (parameters, free_gradients), ( + nudged_parameters, nudged_gradients) in zip( + free_groups, nudged_groups): + if tuple(map(id, parameters)) != tuple( + map(id, nudged_parameters)): + raise AssertionError("EP phase parameter groups differ") + # The optimizer descends, so negate the EP ascent direction. + gradients = [ + -(nudged_gradient - free_gradient) / signed_beta + for nudged_gradient, free_gradient in zip( + nudged_gradients, free_gradients)] + self._assign_gradients(parameters, gradients) + one_hot = F.one_hot( + targets, self.config.vocab_size).to(free[-1].dtype) + return { + "free_mse": float(F.mse_loss(free[-1], one_hot)), + "free_steps": int(free_steps), + "nudge_steps": int(nudge_steps), + "signed_beta": signed_beta, + "local_vjp_evaluations": int( + (self.config.depth + 1) + * (free_steps + nudge_steps)), + "uses_task_loss_backward": False, + }, free, nudged |
