Fix feature_matching loss weight applied twice (400x instead of 20x) - #358
Open
taemincho wants to merge 1 commit into
Open
Fix feature_matching loss weight applied twice (400x instead of 20x)#358taemincho wants to merge 1 commit into
taemincho wants to merge 1 commit into
Conversation
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.
Bug
In
RAVE.training_step, entries added toloss_genare already weighted atconstruction time, e.g.:
But the final accumulation loop re-applies the same weight via
self.weights.get(k, 1.):Any key in
loss_genwhose name also exists inself.weights getsdouble-weighted. With the default config, feature_matching's weight (20)
is effectively squared to 400x. adversarial (weight 1.0) is unaffected
numerically, and the spectral distance terms escape only because their
dict keys (
multiband_spectral_distance,fullband_spectral_distance)don't match
self.weights's keys (multiband_audio_distance,audio_distance).This regression was introduced in 62a168a ("merge with last version +
normalization + bug fixes + input/output transforms"), which replaced the
previous,
correct loss = sum(loss_gen.values(), 0)with the current loop,without removing the pre-existing per-key weighting.
Fix
Since every entry in loss_gen is already weighted before insertion, the
final loop should just sum the values without re-weighting:
This restores the original (pre-62a168a) behavior.