|
32 | 32 | logger = getLogger(__name__) |
33 | 33 |
|
34 | 34 |
|
| 35 | +def _remove_deepspeed_hooks(model: nn.Module) -> None: |
| 36 | + """Remove DeepSpeed-injected forward (pre/post) hooks from every submodule. |
| 37 | +
|
| 38 | + DeepSpeed registers hooks whose callables are local closures defined inside |
| 39 | + ``DeepSpeedEngine`` (e.g. ``_module_forward_post_hook``). They are not |
| 40 | + picklable, so they must be removed before ``torch.save(model)``. |
| 41 | + """ |
| 42 | + hook_dict_names = ( |
| 43 | + "_forward_hooks", |
| 44 | + "_forward_pre_hooks", |
| 45 | + "_forward_hooks_with_kwargs", |
| 46 | + "_forward_pre_hooks_with_kwargs", |
| 47 | + ) |
| 48 | + for module in model.modules(): |
| 49 | + forward_hooks = getattr(module, "_forward_hooks", None) |
| 50 | + pre_hooks = getattr(module, "_forward_pre_hooks", None) |
| 51 | + stale_ids = set() |
| 52 | + for hook_dict in (forward_hooks, pre_hooks): |
| 53 | + if not hook_dict: |
| 54 | + continue |
| 55 | + for handle_id, hook in list(hook_dict.items()): |
| 56 | + if "DeepSpeedEngine" in getattr(hook, "__qualname__", ""): |
| 57 | + stale_ids.add(handle_id) |
| 58 | + if not stale_ids: |
| 59 | + continue |
| 60 | + for hook_dict_name in hook_dict_names: |
| 61 | + hook_dict = getattr(module, hook_dict_name, None) |
| 62 | + if not hook_dict: |
| 63 | + continue |
| 64 | + for handle_id in stale_ids: |
| 65 | + hook_dict.pop(handle_id, None) |
| 66 | + |
| 67 | + |
35 | 68 | @dataclass |
36 | 69 | class GlobalPTQDistributed(PostQuantizationProcess): |
37 | 70 | """Global PTQ via Trainer-based KL distillation. |
@@ -468,4 +501,22 @@ def run( |
468 | 501 | param.requires_grad = False |
469 | 502 | quantized_model.eval() |
470 | 503 |
|
| 504 | + # HF Trainer enables gradient checkpointing by attaching a |
| 505 | + # non-picklable forward hook (``make_inputs_require_grads``, a |
| 506 | + # local closure) via ``enable_input_require_grads``. Leaving it on |
| 507 | + # the model breaks ``torch.save(model)`` in |
| 508 | + # ``save_quantized_model_pt`` with an AttributeError. Remove it |
| 509 | + # here (no-op when no hook is registered). |
| 510 | + if hasattr(quantized_model, "gradient_checkpointing_disable"): |
| 511 | + quantized_model.gradient_checkpointing_disable() |
| 512 | + if hasattr(quantized_model, "disable_input_require_grads"): |
| 513 | + quantized_model.disable_input_require_grads() |
| 514 | + |
| 515 | + # DeepSpeed wraps the module and injects non-picklable forward |
| 516 | + # (pre/post) hook closures (e.g. ``_module_forward_post_hook`` from |
| 517 | + # ``DeepSpeedEngine``) onto every submodule. These persist after |
| 518 | + # training and likewise break ``torch.save(model)``. Strip any hook |
| 519 | + # whose closure originates from DeepSpeed. |
| 520 | + _remove_deepspeed_hooks(quantized_model) |
| 521 | + |
471 | 522 | logger.info("GlobalPTQDistributed complete.") |
0 commit comments