Conversation
LigerFusedLinearDistillationBase (JSD, cosine similarity) always allocated a full grad_weight buffer and differentiated the student weight (and bias) on every chunk, even when they do not require grad, as with a frozen lm_head under LoRA/PEFT. Build argnums from requires_grad, as linkedin#1441 did for the preference losses: frozen parameters get no gradient buffer and no GEMM, input gradients are unchanged, and a call where nothing requires grad runs forward only.
1 of 3 tasks
kolehma8
reviewed
Sep 24, 2026
| # arg 6: teacher_bias | ||
| argnums = [] | ||
| if input_requires_grad: | ||
| argnums.append(0) |
Collaborator
There was a problem hiding this comment.
create enums for the numbers so they are easier to follow
Collaborator
There was a problem hiding this comment.
constants would also do, but enum is bit cleaner in my opinion
| ) = torch.func.grad_and_value(loss_func_to_call, argnums=argnums, has_aux=True)(*args) | ||
|
|
||
| grad_map = dict(zip(argnums, grads)) | ||
| chunk_grad_input = grad_map.get(0, None) |
Collaborator
There was a problem hiding this comment.
same here, lets use enums for the 0 etc.
| chunk_grad_bias = grad_map.get(5, None) | ||
|
|
||
| # Accumulate gradients | ||
| if grad_weight is not None and chunk_grad_weight is not None: |
Collaborator
There was a problem hiding this comment.
assert chunk_grad_weight is not None?
Collaborator
|
Thank you @djl11 for your contribution. I had few stylistic comments around the argument order constants but otherwise the PR looks good. Once you address those I am happy to merge the PR. |
This branch has not been deployed
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.
Summary
LigerFusedLinearDistillationBase, which backsLigerFusedLinearJSDLossandLigerFusedLinearCosineSimilarityLoss, always allocates a fullgrad_weightbuffer and differentiates the student weight (and bias) on every chunk, even when they do not require grad. That is the usual case when distilling into LoRA/PEFT adapters with a frozenlm_head: the weight gradient is computed on every chunk and then discarded by autograd.This applies to the distillation base the change #1441 made to the preference bases. The
argnumspassed totorch.func.grad_and_valueare built fromrequires_grad, so a frozen weight or bias gets no gradient buffer and no per-chunk GEMM, and a call where nothing requires grad runs forward only. Trainable parameters take the same path as before.Details
Forward + backward in bf16,
LigerFusedLinearJSDLosswith its defaults (compiled,chunk_size=1024), B·T = 4096, H = 4096. Each figure is the median of 10 timed iterations after 3 warm-ups, then the median over separate processes, with main and this branch run alternately.Loss and input gradients are unchanged; the new tests compare frozen and trainable runs for both losses.
#1489 touches the same function (it skips the hard loss when it is neither weighted nor returned). The two are independent, but whichever lands second needs a small rebase.
Testing Done
make testto ensure correctness: I ran the test files for the changed base rather than the full suite.test/chunked_loss/test_jsd_loss.pyandtest/chunked_loss/test_cosine_loss.py: 120 passed on this branch (104 on main, plus 16 new). They also pass on CPU.make checkstyleto ensure code style (ruff 0.15.22checkandformat --checkon the changed files)make test-convergenceto ensure convergence: not run; the convergence tests do not exercise the chunked losses.New tests, for both losses: a frozen vs a trainable weight and bias, compiled and eager; a trainable weight with a frozen bias; nothing requiring grad; and a check that a frozen weight is left out of
argnums, so its gradient is skipped rather than computed and discarded.