diff options
Diffstat (limited to 'experiments/transformer_feedback_smoke.py')
| -rw-r--r-- | experiments/transformer_feedback_smoke.py | 77 |
1 files changed, 77 insertions, 0 deletions
diff --git a/experiments/transformer_feedback_smoke.py b/experiments/transformer_feedback_smoke.py index 1e27b38..2cd22d9 100644 --- a/experiments/transformer_feedback_smoke.py +++ b/experiments/transformer_feedback_smoke.py @@ -319,6 +319,82 @@ def audit_pepita(): } +def audit_forward_forward(): + config = LocalTransformerConfig( + vocab_size=7, context_length=5, depth=2, width=8, heads=2, + mlp_ratio=2, seed=151) + bp = LocalDecoderTransformer(config, "bp", dtype=torch.float64) + ff = LocalDecoderTransformer(config, "ff", dtype=torch.float64) + generator = torch.Generator().manual_seed(152) + tokens = torch.randint(0, 7, (3, 5), generator=generator) + positive = torch.randint(0, 7, (3, 5), generator=generator) + negative = (positive + torch.randint( + 1, 7, positive.shape, generator=generator)) % 7 + with torch.no_grad(): + forward_error = float(torch.max(torch.abs( + bp(tokens)["logits"] - ff(tokens)["logits"]))) + all_parameters = tuple(ff.parameters()) + audited = [] + for layer_index in range(ff.ff_num_layers): + for parameter in all_parameters: + parameter.grad = None + target_ids = { + id(parameter) + for parameter in ff.ff_layer_parameters(layer_index)} + loss, metrics = ff.ff_local_loss( + layer_index, tokens, positive, negative) + loss.backward() + unexpected = [ + index for index, parameter in enumerate(all_parameters) + if parameter.grad is not None and id(parameter) not in target_ids] + missing = [ + index for index, parameter in enumerate(all_parameters) + if parameter.grad is None and id(parameter) in target_ids] + explicit_positive = ff._ff_local_outputs( + layer_index, tokens, positive, + ff.ff_forward(tokens, positive)) + explicit_negative = ff._ff_local_outputs( + layer_index, tokens, negative, + ff.ff_forward(tokens, negative)) + explicit = ( + torch.nn.functional.softplus( + -explicit_positive.square().mean(dim=-1) + 2.0) + + torch.nn.functional.softplus( + explicit_negative.square().mean(dim=-1) - 2.0) + ).mean() + loss_error = abs(float(loss.detach() - explicit.detach())) + if unexpected or missing or loss_error >= 1e-12: + raise AssertionError({ + "layer": layer_index, + "unexpected": unexpected, + "missing": missing, + "loss_error": loss_error, + }) + audited.append({ + "layer": layer_index, + "loss": float(loss.detach()), + "loss_error": loss_error, + **metrics, + }) + scores = ff.ff_candidate_scores(tokens) + if (forward_error >= 1e-12 or scores.shape != (3, 5, 7) + or not torch.isfinite(scores).all()): + raise AssertionError({ + "forward_error": forward_error, + "score_shape": list(scores.shape), + }) + return { + "matched_forward_max_absolute_error": forward_error, + "audited_greedy_layers": len(audited), + "max_local_objective_absolute_error": max( + row["loss_error"] for row in audited), + "non_target_parameter_gradients": 0, + "candidate_score_shape": list(scores.shape), + "candidate_presentations_per_evaluation": + config.vocab_size, + } + + def main(): torch.set_num_threads(1) result = { @@ -329,6 +405,7 @@ def main(): "sdil_innovation_identity": audit_sdil_innovation_identity(), "strict_dfa": audit_strict_dfa(), "pepita": audit_pepita(), + "forward_forward": audit_forward_forward(), } print(json.dumps(result, indent=2, sort_keys=True)) |
