Let a loss function take the batch inputs - #676
Open
georgeguimaraes wants to merge 1 commit into
Open
georgeguimaraes wants to merge 1 commit into
georgeguimaraes wants to merge 1 commit into
Conversation
A loss can now be an arity-3 function loss(y_true, y_pred, x) that also receives the batch inputs, for losses that need something from the batch besides the targets, such as a per-sample weight carried under its own input key. The train and eval step states keep the batch as :x so the loss metric, and any custom metric, can read it as well. Arity-2 losses and the built-in ones are unchanged.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
A loss can now be an arity-3 function
loss(y_true, y_pred, x)that also receives the batch inputs. The train and eval step states keep the batch as:x, so the"loss"metric and custom metrics can read it too, andvalidate/4passes it along. Arity-2 losses and the built-in ones are unchanged.We needed this in soothsayer to weight recent rows more in the loss, NeuralProphet's newer samples weight. The weight is per sample, so it belongs with the batch, but the loss only saw targets and predictions and we had to keep a training loop of our own outside
Axon.Loopto get it in. With this the weight rides under its own input key and the loss picks it up.