fix(autograd): CrossEntropyLoss tape detach, softmax/variance backward on rank>=3 - #877
Merged
Merged
Conversation
…d on rank>=3 Three autograd correctness bugs that caused silently-frozen or crashing training: - CrossEntropyLoss (#862): the index-target path built its result host-side via tensorDataFactory + fromData, detaching the tape so no gradient reached the predictions; the soft-target path only recorded when targets lived on the recording context. Both now compute the NLL with differentiable ops dispatched through the predictions' ops (a constant one-hot for the index path), keeping the tape connected. - softmax/logSoftmax backward (#863): a negative dim (e.g. softmax(dim=-1)) was passed straight to broadcastToInput, which reserves -1 as a no-unsqueeze sentinel, so the reduced axis was never re-expanded and backward crashed for rank>=3. softmaxGrad/logSoftmaxGrad now normalize the dim. - varianceBackward (#864): the reduced mean and upstream kept the reduced shape, so the subtract/multiply could not broadcast for rank>=2, and N used the total volume instead of the axis size. Both are now expanded over the reduced axis and N is the axis size. Adds finite-difference backward tests for softmax/logSoftmax/variance on rank-3 tensors and analytic-gradient tests for index- and soft-target CrossEntropyLoss. Closes #862, #863, #864
|
📖 Documentation Preview The documentation has been built successfully for this PR. Generated Files:
Artifacts:
This comment will be updated automatically when the PR is updated. |
aharakal
approved these changes
Jul 24, 2026
This was referenced Jul 24, 2026
Closed
MacOS
pushed a commit
to MacOS/SKaiNET
that referenced
this pull request
Jul 27, 2026
Bump version 0.36.0 -> 0.37.0 (gradle.properties, docs/antora.yml, README quickstart). Promote CHANGELOG [Unreleased] to [0.37.0]: Lstm layer (SKaiNET-developers#824), real Dropout masking (SKaiNET-developers#867), LR schedules and mutable optimizer lr (SKaiNET-developers#866), optional-bias and open Linear (SKaiNET-developers#870, SKaiNET-developers#875), androidNative IO targets (SKaiNET-developers#836, SKaiNET-developers#842, SKaiNET-developers#845), the SDPA default-scale fix (SKaiNET-developers#880), three autograd fixes (SKaiNET-developers#877), the argMax DAG output spec (SKaiNET-developers#878), tokenizer BPE inference and N-D gather (SKaiNET-developers#879), plus the CI/docs supply-chain hardening and toolchain bumps. Refresh README "What's New" and add a Contributors (0.37.0) section. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
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.
Fixes three autograd correctness bugs that cause silently frozen or crashing training — the "wrong but no error" class.
#862— CrossEntropyLoss detaches the tapeThe index-target path built its result host-side (
tensorDataFactory.init+fromData), readinglogProbsvalues out of the tensor. The resulting loss had no recorded parents, sobackwardstopped there and the predictions got a null/zero gradient — training froze at constant loss with no error. The soft-target path had a subtler variant:targets * logProbsdispatches throughtargets.ops, so it only recorded when the targets happened to live on the recording context.Fix: compute the NLL with differentiable ops dispatched through the predictions' ops. The index path builds a constant one-hot selector (reading discrete indices host-side is fine — they carry no gradient) and computes
-sum(oneHot * logProbs); the soft path multiplies vialogProbs.ops. Loss values are unchanged; the tape stays connected.#863— softmax/logSoftmax backward crashes on negative dim, rank ≥ 3softmax(dim = -1)passed-1straight tobroadcastToInput, which reserves-1as a "no-unsqueeze" sentinel — so the reduced axis was never re-expanded and the subsequent elementwise op threwShapes [2,2,4,4] and [2,2,4] cannot be broadcasted. Any attention-style model using the naturalsoftmax(dim = -1)spelling crashed in backward.Fix:
softmaxGrad/logSoftmaxGradnormalize a negative dim (dim + rank) before use. Softmax always reduces exactly one real axis, so this is unambiguous and leaves the sentinel semantics ofsumGrad/meanGradintact.#864— varianceBackward mis-broadcasts on rank ≥ 2mean(x, dim)dropped the reduced axis (no keepdim), sox - meancould not broadcast;Nused the total volume instead of the axis size; and the reduced-shapeupstreamwas never expanded back. A LayerNorm built as(x - mean) / sqrt(variance(x, dim) + eps)crashed in backward.Fix: re-expand the mean and upstream over the reduced axis (unsqueeze /
broadcastToInput) and use the axis size forN.Tests
OpsAutodiffBackwardTest: finite-difference backward checks forsoftmax(dim=-1),logSoftmax(dim=-1), andvariance(dim=2)on rank-3 tensors (non-uniform upstream so a wrong gradient is detectable).CrossEntropyBackwardTest(new): index-target CE gradient matches the exact analytic(softmax − oneHot)/N; soft-target CE records a non-zero gradient even when the targets live on a separate eager context.jvmTestforskainet-lang-core,skainet-compile-dag, andskainet-backend-cpugreen.Found while porting an educational GPT trainer to SKaiNET; each bug had a precise reproduction there (training froze / crashed until worked around). This removes the need for those workarounds.