From e536616888d51b453ed354a6f1e243fecb02ea08 Mon Sep 17 00:00:00 2001 From: Younes Belkada <49240599+younesbelkada@users.noreply.github.com> Date: Sat, 1 Apr 2023 14:54:46 +0200 Subject: [PATCH] [`core`] Fix offload issue (#248) * fix offload dir * remove offload index * safety checker * forward contrib credits from previous PR --------- Co-authored-by: cosimoiaia --- src/peft/peft_model.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 0afd047..f9573bb 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -164,6 +164,15 @@ class PeftModel(PushToHubMixin, torch.nn.Module): if getattr(model, "hf_device_map", None) is not None: device_map = kwargs.get("device_map", "auto") max_memory = kwargs.get("max_memory", None) + offload_dir = kwargs.get("offload_dir", None) + offload_index = kwargs.get("offload_index", None) + + dispatch_model_kwargs = {} + # Safety checker for previous `accelerate` versions + # `offload_index` was introduced in https://github.com/huggingface/accelerate/pull/873/ + if "offload_index" in inspect.signature(dispatch_model).parameters: + dispatch_model_kwargs["offload_index"] = offload_index + no_split_module_classes = model._no_split_modules if device_map != "sequential": max_memory = get_balanced_memory( @@ -176,7 +185,13 @@ class PeftModel(PushToHubMixin, torch.nn.Module): device_map = infer_auto_device_map( model, max_memory=max_memory, no_split_module_classes=no_split_module_classes ) - model = dispatch_model(model, device_map=device_map) + + model = dispatch_model( + model, + device_map=device_map, + offload_dir=offload_dir, + **dispatch_model_kwargs, + ) hook = AlignDevicesHook(io_same_device=True) if model.peft_config.peft_type == PeftType.LORA: add_hook_to_module(model.base_model.model, hook)