Commit 0315ebf
mem: re-dequantize frozen expert weights in backward instead of saving them
Route every expert projection through _FrozenLinearRecomputeBackward, a
small autograd Function whose forward computes dequantize + F.linear
directly and whose backward re-dequantizes the frozen weight to form
grad_output @ W:
- Bit-exact by construction, in and out of grad mode, on every device:
recomputation changes what is saved, never what is computed. (An earlier
iteration routed through bnb.matmul_4bit instead; it was dropped because
the fused gemm_4bit kernel it dispatches to for small token batches
differs from dequantize+linear by accumulation order.)
- The dequantized [out, in] expert weight never enters autograd's
saved-tensor storage, so training activation memory stays independent of
the number of experts held between forward and backward — and the
mechanism is scheme-agnostic (nf4/fp4/int8/fp8/passthrough), no
per-scheme kernel needed.
- Cost: one extra dequantize per projection in backward; subsumed by full
gradient checkpointing when that is enabled.
Tests: test_experts4bit_forward_is_bit_exact_dequantize_linear pins forward
== plain dequantize+linear at rtol=0/atol=0 (grad and no_grad);
test_experts4bit_backward_saves_no_dequantized_weight uses
saved_tensors_hooks to assert nothing weight-shaped (either orientation) is
saved while a plain dequantize+linear control does save it, and that
gradients match the control exactly.
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>1 parent fb1b247 commit 0315ebf
2 files changed
Lines changed: 128 additions & 5 deletions
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
3 | 3 | | |
4 | 4 | | |
5 | 5 | | |
| 6 | + | |
6 | 7 | | |
7 | 8 | | |
8 | 9 | | |
| |||
42 | 43 | | |
43 | 44 | | |
44 | 45 | | |
| 46 | + | |
| 47 | + | |
| 48 | + | |
| 49 | + | |
| 50 | + | |
| 51 | + | |
| 52 | + | |
| 53 | + | |
| 54 | + | |
| 55 | + | |
| 56 | + | |
| 57 | + | |
| 58 | + | |
| 59 | + | |
| 60 | + | |
| 61 | + | |
| 62 | + | |
| 63 | + | |
| 64 | + | |
| 65 | + | |
| 66 | + | |
| 67 | + | |
| 68 | + | |
| 69 | + | |
| 70 | + | |
| 71 | + | |
| 72 | + | |
| 73 | + | |
45 | 74 | | |
46 | 75 | | |
47 | 76 | | |
| |||
72 | 101 | | |
73 | 102 | | |
74 | 103 | | |
75 | | - | |
76 | | - | |
| 104 | + | |
| 105 | + | |
| 106 | + | |
| 107 | + | |
| 108 | + | |
77 | 109 | | |
78 | 110 | | |
79 | 111 | | |
| |||
299 | 331 | | |
300 | 332 | | |
301 | 333 | | |
302 | | - | |
303 | | - | |
304 | | - | |
| 334 | + | |
| 335 | + | |
| 336 | + | |
| 337 | + | |
| 338 | + | |
| 339 | + | |
| 340 | + | |
305 | 341 | | |
306 | 342 | | |
307 | 343 | | |
| |||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
348 | 348 | | |
349 | 349 | | |
350 | 350 | | |
| 351 | + | |
| 352 | + | |
| 353 | + | |
| 354 | + | |
| 355 | + | |
| 356 | + | |
| 357 | + | |
| 358 | + | |
| 359 | + | |
| 360 | + | |
| 361 | + | |
| 362 | + | |
| 363 | + | |
| 364 | + | |
| 365 | + | |
| 366 | + | |
| 367 | + | |
| 368 | + | |
| 369 | + | |
| 370 | + | |
| 371 | + | |
| 372 | + | |
| 373 | + | |
| 374 | + | |
| 375 | + | |
| 376 | + | |
| 377 | + | |
| 378 | + | |
| 379 | + | |
| 380 | + | |
| 381 | + | |
| 382 | + | |
| 383 | + | |
| 384 | + | |
| 385 | + | |
| 386 | + | |
| 387 | + | |
| 388 | + | |
| 389 | + | |
| 390 | + | |
| 391 | + | |
| 392 | + | |
| 393 | + | |
| 394 | + | |
| 395 | + | |
| 396 | + | |
| 397 | + | |
| 398 | + | |
| 399 | + | |
| 400 | + | |
| 401 | + | |
| 402 | + | |
| 403 | + | |
| 404 | + | |
| 405 | + | |
| 406 | + | |
| 407 | + | |
| 408 | + | |
| 409 | + | |
| 410 | + | |
| 411 | + | |
| 412 | + | |
| 413 | + | |
| 414 | + | |
| 415 | + | |
| 416 | + | |
| 417 | + | |
| 418 | + | |
| 419 | + | |
| 420 | + | |
| 421 | + | |
| 422 | + | |
| 423 | + | |
| 424 | + | |
| 425 | + | |
| 426 | + | |
| 427 | + | |
| 428 | + | |
| 429 | + | |
| 430 | + | |
| 431 | + | |
| 432 | + | |
| 433 | + | |
| 434 | + | |
| 435 | + | |
| 436 | + | |
| 437 | + | |
0 commit comments