summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/transformer_feedback_smoke.py116
-rw-r--r--sdil/transformer.py285
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