Initial load to offload device instead of device_map
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user