Skip to content

perf(distillation): skip dW/dB calculation when parameters are frozen - #1487

Open
djl11 wants to merge 2 commits into
linkedin:mainfrom
djl11:pr/distillation-skip-frozen-grads
Open

djl11 wants to merge 2 commits into
linkedin:mainfrom
djl11:pr/distillation-skip-frozen-grads

Conversation

@djl11

@djl11 djl11 commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

Summary

LigerFusedLinearDistillationBase, which backs LigerFusedLinearJSDLoss and LigerFusedLinearCosineSimilarityLoss, always allocates a full grad_weight buffer 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 frozen lm_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 argnums passed to torch.func.grad_and_value are built from requires_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, LigerFusedLinearJSDLoss with 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.

V student head main this PR
128,256 frozen 42.3 ms, 6.51 GiB 32.0 ms, 5.53 GiB (−24%)
151,936 frozen 70.8 ms, 7.98 GiB 59.1 ms, 6.82 GiB (−17%)
128,256 trainable 42.1 ms 42.0 ms
151,936 trainable 71.4 ms 71.3 ms

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

  • Hardware Type: NVIDIA H100 80GB HBM3 (torch 2.14.0+cu130, triton 3.8.0)
  • run make test to ensure correctness: I ran the test files for the changed base rather than the full suite. test/chunked_loss/test_jsd_loss.py and test/chunked_loss/test_cosine_loss.py: 120 passed on this branch (104 on main, plus 16 new). They also pass on CPU.
  • run make checkstyle to ensure code style (ruff 0.15.22 check and format --check on the changed files)
  • run make test-convergence to 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.

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.
# arg 6: teacher_bias
argnums = []
if input_requires_grad:
argnums.append(0)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

create enums for the numbers so they are easier to follow

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

assert chunk_grad_weight is not None?

@kolehma8

Copy link
Copy Markdown
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

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

Labels

None yet

2 participants