From 0e3f3c31df1671596b9dd2758772c5ce227895f2 Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Tue, 28 Jun 2022 21:01:50 +0500 Subject: [PATCH] small fix in injection module to replace transformer --- deepspeed/module_inject/replace_module.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/deepspeed/module_inject/replace_module.py b/deepspeed/module_inject/replace_module.py index 72599e9e43a1..0f686a062acb 100755 --- a/deepspeed/module_inject/replace_module.py +++ b/deepspeed/module_inject/replace_module.py @@ -710,7 +710,7 @@ def replace_module(model, orig_class, replace_fn, _replace_policy): A modified ``model``. """ policy = {} - if orig_class is not None: + if orig_class is not None and _replace_policy is not None: policy.update({orig_class: (replace_fn, _replace_policy)}) else: for plcy in replace_policies: @@ -740,6 +740,11 @@ def _replace_module(model, policies, layer_id=0): Returns: Modified ``model``. """ + if model.__class__ in policies: + replaced_module = policies[model.__class__][0](model, + policies[model.__class__][-1], + layer_id) + return replaced_module, layer_id for name, child in model.named_children(): if child.__class__ in policies: replaced_module = policies[child.__class__][0](child,