transformers
ed9c07e7 - Free per-source tensors in fuse ops + drain MPS pool between stages

Commit
139 days ago
Free per-source tensors in fuse ops + drain MPS pool between stages MoE fuse ops (`MergeModulelist`, `Concatenate`, `ErnieFuseAndSplitTextVisionExperts`) held per-expert source tensors alive in `input_dict` until the caller dropped the dict, doubling peak memory during the stack/cat. Pop or clear sources eagerly so the accelerator caching allocator can reuse buffers immediately. Also call `torch.mps.empty_cache()` after each weight in the load loop when the target is MPS — CUDA's allocator reclaims under pressure; MPS's does not, so freed buffers accumulate and push the load into swap. Gated on `device_map` targeting MPS so CPU/CUDA paths are unaffected. On Mixtral-8x7B (87 GB bf16) loaded via `device_map="mps"`, peak MPS heap drops from ~135 GB (load aborts) to ~88 GB = model size, and the load completes.
Author
Parents
Loading