Fix review findings for PR #15
This commit is contained in:
parent
29ab0fab74
commit
5e8e4487ed
2 changed files with 21 additions and 11 deletions
|
|
@ -2483,7 +2483,7 @@ class FastLlamaModel:
|
|||
load_in_4bit = load_in_4bit,
|
||||
load_in_8bit = kwargs.get("load_in_8bit", False),
|
||||
offload_embedding = False,
|
||||
fast_inference = False,
|
||||
fast_inference = fast_inference,
|
||||
)
|
||||
elif not fast_inference:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
|
|
|
|||
|
|
@ -106,8 +106,17 @@ def _infer_device_map_from_loaded_model(model):
|
|||
def _assign(module, prefix):
|
||||
subtree_devs = {
|
||||
p.device for _, p in module.named_parameters(remove_duplicate = False)
|
||||
} | {b.device for _, b in module.named_buffers()}
|
||||
}
|
||||
if not subtree_devs:
|
||||
bufs = list(module.named_buffers())
|
||||
if bufs:
|
||||
buf_devs = {b.device for _, b in bufs}
|
||||
if len(buf_devs) == 1:
|
||||
device_map[prefix] = next(iter(buf_devs))
|
||||
else:
|
||||
for child_name, child in module.named_children():
|
||||
child_prefix = f"{prefix}.{child_name}" if prefix else child_name
|
||||
_assign(child, child_prefix)
|
||||
return
|
||||
if len(subtree_devs) == 1:
|
||||
device_map[prefix] = next(iter(subtree_devs))
|
||||
|
|
@ -119,13 +128,11 @@ def _infer_device_map_from_loaded_model(model):
|
|||
p.device for _, p in module.named_parameters(
|
||||
recurse = False, remove_duplicate = False,
|
||||
)
|
||||
} | {b.device for _, b in module.named_buffers(recurse = False)}
|
||||
}
|
||||
if local_devs and len(local_devs) == 1:
|
||||
device_map[prefix] = next(iter(local_devs))
|
||||
|
||||
_assign(model, "")
|
||||
if "" in device_map and len(device_map) > 1:
|
||||
device_map.pop("")
|
||||
return device_map
|
||||
|
||||
|
||||
|
|
@ -139,7 +146,14 @@ def _attach_bnb_multidevice_hooks(
|
|||
"""
|
||||
if fast_inference:
|
||||
return
|
||||
if not (load_in_4bit or load_in_8bit):
|
||||
is_bnb = (
|
||||
load_in_4bit
|
||||
or load_in_8bit
|
||||
or getattr(model, "is_loaded_in_4bit", False)
|
||||
or getattr(model, "is_loaded_in_8bit", False)
|
||||
or getattr(model, "quantization_method", None) == "bitsandbytes"
|
||||
)
|
||||
if not is_bnb:
|
||||
return
|
||||
if offload_embedding:
|
||||
return
|
||||
|
|
@ -147,11 +161,7 @@ def _attach_bnb_multidevice_hooks(
|
|||
return # already dispatched
|
||||
|
||||
try:
|
||||
all_devs = {
|
||||
p.device
|
||||
for p in model.parameters()
|
||||
if hasattr(p, "device")
|
||||
}
|
||||
all_devs = {p.device for p in model.parameters()}
|
||||
except Exception as exc:
|
||||
warnings.warn(
|
||||
"Unsloth: Failed to determine device placement from model parameters, "
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue