Skip to content

LayeredWeightedPhysicsLoss reads layer_residuals off the model instead of receiving it #101

Description

@albanpuech

Low priority / cleanup. Pre-existing design issue, no known incorrect results today. Surfaced while reviewing #87.

GNS_heterogeneous.forward writes self.layer_residuals as a side effect, and LayeredWeightedPhysicsLoss (training/loss.py:363) reads model.layer_residuals afterwards. The data dependency is invisible in both signatures — the loss just receives model=self.model.

Concrete downsides:

  • Unenforced ordering. Requires forward() before loss_fn() on the same instance. If the dict is empty, L = 0 produces no loop iterations, total_loss stays a Python float, and it fails on total_loss.item() — an obscure error far from the cause.
  • Not reuse-safe. One dict per module instance: two forward passes before a loss call silently overwrite, so gradient accumulation or evaluating two batches before reducing would read the wrong batch's residuals and return a plausible number rather than erroring.
  • Hidden model/loss coupling. Only works with models exposing layer_residuals, i.e. GNS_heterogeneous; AttributeError on GRIT. A config pairing GRIT with this loss fails at runtime, not at validation.

Working as intended today: the stored tensors keep their autograd graph (no .detach()), so gradients flow correctly.

Suggested fix: return residuals from forward alongside the predictions (as #87 now does for embeddings) and pass them into the loss explicitly, so the dependency is in the signature and the model= kwarg is no longer needed for this.

Deferred from #87 because the fix touches the shared BaseLoss.forward signature across all losses.

Refs #87

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    enhancementNew feature or request

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions