Checks
This commit is contained in:
parent
236f1029b1
commit
fc6ad684c9
4 changed files with 92 additions and 19 deletions
|
|
@ -46,7 +46,8 @@ huggingface = [
|
|||
"trl>=0.7.9,<0.9.0",
|
||||
"peft>=0.7.1,!=0.11.0",
|
||||
"protobuf<4.0.0",
|
||||
"huggingface_hub[hf_transfer]",
|
||||
"huggingface_hub",
|
||||
"hf-transfer",
|
||||
]
|
||||
cu118only = [
|
||||
"xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.22.post7%2Bcu118-cp39-cp39-manylinux2014_x86_64.whl ; python_version=='3.9'",
|
||||
|
|
@ -196,7 +197,8 @@ colab-new = [
|
|||
"wheel>=0.42.0",
|
||||
"numpy",
|
||||
"protobuf<4.0.0",
|
||||
"huggingface_hub[hf_transfer]",
|
||||
"huggingface_hub",
|
||||
"hf-transfer",
|
||||
]
|
||||
colab-no-deps = [
|
||||
"accelerate>=0.26.1",
|
||||
|
|
|
|||
|
|
@ -122,7 +122,8 @@ pass
|
|||
# =============================================
|
||||
# torch.cuda.amp.custom_fwd is deprecated >= 2.4
|
||||
import torch
|
||||
if Version(torch.__version__) < Version("2.4.0"):
|
||||
torch_version = torch.__version__
|
||||
if Version(torch_version) < Version("2.4.0"):
|
||||
torch_amp_custom_fwd = torch.cuda.amp.custom_fwd
|
||||
torch_amp_custom_bwd = torch.cuda.amp.custom_bwd
|
||||
else:
|
||||
|
|
@ -184,7 +185,7 @@ pass
|
|||
|
||||
# Check TRL version
|
||||
from trl import __version__ as trl_version
|
||||
if Version(xformers_version) >= Version("0.9.0"):
|
||||
if Version(trl_version) >= Version("0.9.0"):
|
||||
raise ImportError(
|
||||
"Unsloth: If you are in Colab, we updated the top cell install instructions - please change it to below "\
|
||||
"then press Disconnect Runtime and then Restart it.\n"\
|
||||
|
|
@ -199,7 +200,24 @@ if Version(xformers_version) >= Version("0.9.0"):
|
|||
)
|
||||
pass
|
||||
|
||||
# Confirm versions
|
||||
# =============================================
|
||||
if Version(torch_version) < Version("2.2.0") and Version(xformers_version) >= Version("0.0.24"):
|
||||
raise ImportError(
|
||||
f"Unsloth: You have torch = {torch_version} but xformers = {xformers_version}.\n"\
|
||||
f"Please install xformers < 0.0.24 for torch = {torch_version}."
|
||||
)
|
||||
elif Version(torch_version) < Version("2.3.0") and Version(xformers_version) >= Version("0.0.26"):
|
||||
raise ImportError(
|
||||
f"Unsloth: You have torch = {torch_version} but xformers = {xformers_version}.\n"\
|
||||
f"Please install xformers < 0.0.26 for torch = {torch_version}."
|
||||
)
|
||||
elif Version(torch_version) < Version("2.4.0") and Version(xformers_version) >= Version("0.0.27"):
|
||||
raise ImportError(
|
||||
f"Unsloth: You have torch = {torch_version} but xformers = {xformers_version}.\n"\
|
||||
f"Please install xformers < 0.0.27 for torch = {torch_version}."
|
||||
)
|
||||
pass
|
||||
|
||||
# =============================================
|
||||
# Torch compile settings
|
||||
|
|
|
|||
|
|
@ -1282,7 +1282,7 @@ class FastLlamaModel:
|
|||
max_memory = round(gpu_stats.total_memory / 1024 / 1024 / 1024, 3)
|
||||
|
||||
statistics = \
|
||||
f"==((====))== Unsloth {__version__}: Fast {model_patcher.__name__[4:-5]} patching. Transformers = {transformers_version}\n"\
|
||||
f"==((====))== Unsloth {__version__}: Fast {model_patcher.__name__[4:-5]} patching. Transformers = {transformers_version}.\n"\
|
||||
f" \\\ /| GPU: {gpu_stats.name}. Max memory: {max_memory} GB. Platform = {platform_system}.\n"\
|
||||
f"O^O/ \_/ \\ Pytorch: {torch.__version__}. CUDA = {gpu_stats.major}.{gpu_stats.minor}. CUDA Toolkit = {torch.version.cuda}.\n"\
|
||||
f"\ / Bfloat16 = {str(SUPPORTS_BFLOAT16).upper()}. FA [Xformers = {xformers_version}. FA2 = {HAS_FLASH_ATTENTION}]\n"\
|
||||
|
|
|
|||
|
|
@ -35,9 +35,14 @@ if SUPPORTS_GEMMA2:
|
|||
pass
|
||||
|
||||
|
||||
def _get_model_name(model_name, load_in_4bit = True):
|
||||
def __get_model_name(
|
||||
model_name,
|
||||
load_in_4bit = True,
|
||||
INT_TO_FLOAT_MAPPER = None,
|
||||
FLOAT_TO_INT_MAPPER = None,
|
||||
):
|
||||
|
||||
if not SUPPORTS_FOURBIT and model_name in INT_TO_FLOAT_MAPPER:
|
||||
if not SUPPORTS_FOURBIT and model_name.lower() in INT_TO_FLOAT_MAPPER:
|
||||
model_name = INT_TO_FLOAT_MAPPER[model_name.lower()]
|
||||
logger.warning_once(
|
||||
f"Unsloth: Your transformers version of {transformers_version} does not support native "\
|
||||
|
|
@ -46,25 +51,71 @@ def _get_model_name(model_name, load_in_4bit = True):
|
|||
f"to obtain the latest transformers build, then restart this session.\n"\
|
||||
f"For now, we shall load `{model_name}` instead (still 4bit, just slower downloading)."
|
||||
)
|
||||
return model_name
|
||||
|
||||
elif not load_in_4bit and model_name in INT_TO_FLOAT_MAPPER:
|
||||
elif not load_in_4bit and model_name.lower() in INT_TO_FLOAT_MAPPER:
|
||||
new_model_name = INT_TO_FLOAT_MAPPER[model_name.lower()]
|
||||
# logger.warning_once(
|
||||
# f"Unsloth: You passed in `{model_name}` which is a 4bit model, yet you set\n"\
|
||||
# f"`load_in_4bit = False`. We shall load `{new_model_name}` instead."
|
||||
# )
|
||||
model_name = new_model_name
|
||||
return new_model_name
|
||||
|
||||
elif load_in_4bit and SUPPORTS_FOURBIT and model_name in FLOAT_TO_INT_MAPPER:
|
||||
elif load_in_4bit and SUPPORTS_FOURBIT and model_name.lower() in FLOAT_TO_INT_MAPPER:
|
||||
new_model_name = FLOAT_TO_INT_MAPPER[model_name.lower()]
|
||||
# logger.warning_once(
|
||||
# f"Unsloth: You passed in `{model_name}` and `load_in_4bit = True`.\n"\
|
||||
# f"We shall load `{new_model_name}` for 4x faster loading."
|
||||
# )
|
||||
model_name = new_model_name
|
||||
return new_model_name
|
||||
pass
|
||||
|
||||
return model_name
|
||||
return None
|
||||
pass
|
||||
|
||||
|
||||
def _get_new_mapper():
|
||||
try:
|
||||
import requests
|
||||
new_mapper = "https://raw.githubusercontent.com/unslothai/unsloth/main/unsloth/models/mapper.py"
|
||||
with requests.get(new_mapper, timeout = 3) as new_mapper: new_mapper = new_mapper.text
|
||||
new_mapper = new_mapper\
|
||||
.replace("INT_TO_FLOAT_MAPPER", "NEW_INT_TO_FLOAT_MAPPER")\
|
||||
.replace("FLOAT_TO_INT_MAPPER", "NEW_FLOAT_TO_INT_MAPPER")
|
||||
exec(new_mapper, locals())
|
||||
return NEW_INT_TO_FLOAT_MAPPER, NEW_FLOAT_TO_INT_MAPPER
|
||||
except:
|
||||
return {}, {}
|
||||
pass
|
||||
pass
|
||||
|
||||
|
||||
def _get_model_name(model_name, load_in_4bit = True):
|
||||
new_model_name = __get_model_name(
|
||||
model_name = model_name,
|
||||
load_in_4bit = load_in_4bit,
|
||||
INT_TO_FLOAT_MAPPER = INT_TO_FLOAT_MAPPER,
|
||||
FLOAT_TO_INT_MAPPER = FLOAT_TO_INT_MAPPER,
|
||||
)
|
||||
if new_model_name is None and \
|
||||
model_name.count("/") == 1 and \
|
||||
model_name[0].isalnum():
|
||||
# Try checking if a new Unsloth version allows it!
|
||||
NEW_INT_TO_FLOAT_MAPPER, NEW_FLOAT_TO_INT_MAPPER = _get_new_mapper()
|
||||
upgraded_model_name = __get_model_name(
|
||||
model_name = model_name,
|
||||
load_in_4bit = load_in_4bit,
|
||||
INT_TO_FLOAT_MAPPER = NEW_INT_TO_FLOAT_MAPPER,
|
||||
FLOAT_TO_INT_MAPPER = NEW_FLOAT_TO_INT_MAPPER,
|
||||
)
|
||||
if upgraded_model_name is not None:
|
||||
raise NotImplementedError(
|
||||
f"Unsloth: {model_name} is not supported in your current Unsloth version! Please update Unsloth via:\n\n"\
|
||||
'pip install --upgrade --force-reinstall --no-cache-dir "unsloth[colab-new] @ git+https://github.com/unslothai/unsloth.git"'
|
||||
)
|
||||
pass
|
||||
pass
|
||||
return new_model_name if new_model_name is not None else model_name
|
||||
pass
|
||||
|
||||
|
||||
|
|
@ -98,16 +149,22 @@ class FastLanguageModel(FastLlamaModel):
|
|||
from huggingface_hub.utils import disable_progress_bars, enable_progress_bars, are_progress_bars_disabled
|
||||
was_disabled = are_progress_bars_disabled()
|
||||
disable_progress_bars()
|
||||
|
||||
autoconfig_error = None
|
||||
peft_error = None
|
||||
try:
|
||||
model_config = AutoConfig.from_pretrained(model_name, token = token, revision = revision)
|
||||
is_model = True
|
||||
except:
|
||||
except Exception as autoconfig_error:
|
||||
autoconfig_error = str(autoconfig_error)
|
||||
is_model = False
|
||||
try:
|
||||
peft_config = PeftConfig .from_pretrained(model_name, token = token, revision = revision)
|
||||
is_peft = True
|
||||
except:
|
||||
except Exception as peft_error:
|
||||
peft_error = str(peft_error)
|
||||
is_peft = False
|
||||
pass
|
||||
|
||||
# Cannot be both!
|
||||
if is_model and is_peft:
|
||||
|
|
@ -118,11 +175,7 @@ class FastLanguageModel(FastLlamaModel):
|
|||
"Please separate the LoRA and base models to 2 repos."
|
||||
)
|
||||
elif not is_model and not is_peft:
|
||||
raise RuntimeError(
|
||||
f"Unsloth: `{model_name}` is not a base model or a PEFT model.\n"\
|
||||
"We could not locate a `config.json` or `adapter_config.json` file.\n"\
|
||||
"Are you certain the model name is correct? Does it actually exist?"
|
||||
)
|
||||
raise RuntimeError(autoconfig_error or peft_error)
|
||||
pass
|
||||
|
||||
# Get base model for PEFT:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue