Skip to content

Commit 8c63cbe

Browse files
committed
Bug-fix GlobalPTQDistributed to save a quantized model
1 parent bec2767 commit 8c63cbe

1 file changed

Lines changed: 51 additions & 0 deletions

File tree

onecomp/post_process/global_ptq_distributed.py

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,39 @@
3232
logger = getLogger(__name__)
3333

3434

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+
3568
@dataclass
3669
class GlobalPTQDistributed(PostQuantizationProcess):
3770
"""Global PTQ via Trainer-based KL distillation.
@@ -468,4 +501,22 @@ def run(
468501
param.requires_grad = False
469502
quantized_model.eval()
470503

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+
471522
logger.info("GlobalPTQDistributed complete.")

0 commit comments

Comments
 (0)