From 820689f6dd5e67bc1a908db3a40e9f59e544176f Mon Sep 17 00:00:00 2001 From: Silver <65376327+silveroxides@users.noreply.github.com> Date: Mon, 14 Apr 2025 12:03:39 +0200 Subject: [PATCH] Initial load to offload device instead of device_map --- nodes.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/nodes.py b/nodes.py index 57d5834..dc3972a 100644 --- a/nodes.py +++ b/nodes.py @@ -123,7 +123,7 @@ class DownloadAndLoadFlorence2Model: print(f"Florence2 using {attention} for attention") with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement - model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device, torch_dtype=dtype,trust_remote_code=True) + model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, torch_dtype=dtype,trust_remote_code=True).to(offload_device) processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True) if lora is not None: @@ -197,12 +197,13 @@ class Florence2ModelLoader: def loadmodel(self, model, precision, attention, lora=None): device = mm.get_torch_device() + offload_device = mm.unet_offload_device() dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] model_path = Florence2ModelLoader.model_paths.get(model) print(f"Loading model from {model_path}") print(f"Florence2 using {attention} for attention") with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement - model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device, torch_dtype=dtype,trust_remote_code=True) + model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, torch_dtype=dtype,trust_remote_code=True).to(offload_device) processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True) if lora is not None: