logger
This commit is contained in:
parent
c05af38c22
commit
76788e0f8a
2 changed files with 20 additions and 18 deletions
|
|
@ -16,7 +16,6 @@ import warnings, importlib, sys
|
|||
from packaging.version import Version
|
||||
import os, re, subprocess, inspect, functools
|
||||
import numpy as np
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
# Log Unsloth is being used
|
||||
os.environ["UNSLOTH_IS_PRESENT"] = "1"
|
||||
|
|
@ -27,7 +26,6 @@ from .import_fixes import (
|
|||
check_fbgemm_gpu_version,
|
||||
torchvision_compatibility_check,
|
||||
)
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
fix_message_factory_issue()
|
||||
check_fbgemm_gpu_version()
|
||||
|
|
|
|||
|
|
@ -13,16 +13,15 @@
|
|||
# limitations under the License.
|
||||
|
||||
import os
|
||||
import sys
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
from importlib.metadata import version as importlib_version
|
||||
from packaging.version import Version as TrueVersion
|
||||
import re
|
||||
import logging
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
from unsloth_zoo.log import logger
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
# Cannot import logger here since it'll import transformers
|
||||
# from unsloth_zoo.log import logger
|
||||
|
||||
|
||||
def Version(version):
|
||||
try:
|
||||
|
|
@ -42,7 +41,7 @@ def Version(version):
|
|||
f"Unsloth: Could not get version for `{version}`\n"
|
||||
f"File name = [{caller.filename}] Line number = [{caller.lineno}]"
|
||||
)
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
# Ignore logging messages
|
||||
class HideLoggingMessage(logging.Filter):
|
||||
|
|
@ -53,7 +52,7 @@ class HideLoggingMessage(logging.Filter):
|
|||
|
||||
def filter(self, x):
|
||||
return not (self.text in x.getMessage())
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
# Fix up AttributeError: 'MessageFactory' object has no attribute 'GetPrototype'
|
||||
# MUST do this at the start primarily due to tensorflow causing issues
|
||||
|
|
@ -71,6 +70,7 @@ def fix_message_factory_issue():
|
|||
def GetPrototype(self, *args, **kwargs):
|
||||
return
|
||||
|
||||
from unsloth_zoo.log import logger
|
||||
if not hasattr(google.protobuf.message_factory, "MessageFactory"):
|
||||
logger.info("Unsloth: Patching protobuf.MessageFactory as it doesn't exist")
|
||||
google.protobuf.message_factory.MessageFactory = MessageFactory
|
||||
|
|
@ -100,7 +100,7 @@ def fix_message_factory_issue():
|
|||
pass
|
||||
except:
|
||||
pass
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
# Fix Xformers performance issues since 0.0.25
|
||||
def fix_xformers_performance_issue():
|
||||
|
|
@ -108,6 +108,7 @@ def fix_xformers_performance_issue():
|
|||
return
|
||||
xformers_version = importlib_version("xformers")
|
||||
if Version(xformers_version) < Version("0.0.29"):
|
||||
from unsloth_zoo.log import logger
|
||||
xformers_location = importlib.util.find_spec("xformers").origin
|
||||
xformers_location = os.path.split(xformers_location)[0]
|
||||
cutlass = Path(xformers_location) / "ops" / "fmha" / "cutlass.py"
|
||||
|
|
@ -129,7 +130,7 @@ def fix_xformers_performance_issue():
|
|||
)
|
||||
except Exception as e:
|
||||
logger.info(f"Unsloth: Failed patching Xformers with error = {str(e)}")
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
# ValueError: 'aimv2' is already used by a Transformers config, pick another name.
|
||||
def fix_vllm_aimv2_issue():
|
||||
|
|
@ -137,6 +138,7 @@ def fix_vllm_aimv2_issue():
|
|||
return
|
||||
vllm_version = importlib_version("vllm")
|
||||
if Version(vllm_version) < Version("0.10.1"):
|
||||
from unsloth_zoo.log import logger
|
||||
vllm_version = importlib.util.find_spec("vllm").origin
|
||||
vllm_version = os.path.split(vllm_version)[0]
|
||||
ovis_config = Path(vllm_version) / "transformers_utils" / "configs" / "ovis.py"
|
||||
|
|
@ -168,7 +170,7 @@ def fix_vllm_aimv2_issue():
|
|||
)
|
||||
except Exception as e:
|
||||
logger.info(f"Unsloth: Failed patching vLLM with error = {str(e)}")
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
def fix_vllm_guided_decoding_params():
|
||||
if importlib.util.find_spec("vllm") is None:
|
||||
|
|
@ -184,7 +186,7 @@ def fix_vllm_guided_decoding_params():
|
|||
vllm.sampling_params.GuidedDecodingParams = (
|
||||
vllm.sampling_params.StructuredOutputsParams
|
||||
)
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
def ignore_logger_messages():
|
||||
# Ignore Environment variable `HF_TOKEN` is set
|
||||
|
|
@ -195,7 +197,7 @@ def ignore_logger_messages():
|
|||
del huggingface_hub_logger
|
||||
except:
|
||||
pass
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
def patch_ipykernel_hf_xet():
|
||||
# HF-XET == 1.1.10 and ipykernel == 7.0.0 / 7.0.1 causes issues
|
||||
|
|
@ -227,7 +229,7 @@ def patch_ipykernel_hf_xet():
|
|||
from huggingface_hub.utils import disable_progress_bars
|
||||
|
||||
disable_progress_bars()
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
def patch_trackio():
|
||||
# Set some environment variables to customize the Trackio dashboard for experiment tracking
|
||||
|
|
@ -239,7 +241,7 @@ def patch_trackio():
|
|||
"https://raw.githubusercontent.com/unslothai/unsloth/main/images/unsloth%20logo%20white%20text.png"
|
||||
)
|
||||
os.environ["TRACKIO_PLOT_ORDER"] = "train/reward"
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
def patch_datasets():
|
||||
# Datasets 4.4.0 and 4.4.1 weirdly have some weird `_thread.RLock_recursion_count` issues
|
||||
|
|
@ -254,7 +256,7 @@ def patch_datasets():
|
|||
f"#### Unsloth: Using `datasets = {str(datasets_version)}` will cause recursion errors.\n"
|
||||
"Please downgrade datasets to `datasets==4.3.0"
|
||||
)
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
|
||||
def check_fbgemm_gpu_version():
|
||||
if importlib.util.find_spec("fbgemm_gpu") is None:
|
||||
|
|
@ -268,9 +270,9 @@ def check_fbgemm_gpu_version():
|
|||
raise ImportError(
|
||||
f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu_version} detected. It might cause unexpected issues like segmentation faults. Please uninstall the current one by doing `pip uninstall fbgemm-gpu` && `pip install fbgemm-gpu` to install fbgemm-gpu 1.4.0 or newer!"
|
||||
)
|
||||
from unsloth_zoo.log import logger
|
||||
logger.info(f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu_version} detected.")
|
||||
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
def patch_enable_input_require_grads():
|
||||
"""
|
||||
|
|
@ -330,11 +332,11 @@ def patch_enable_input_require_grads():
|
|||
self._require_grads_hook = hooks[0]
|
||||
|
||||
PreTrainedModel.enable_input_require_grads = _patched_enable_input_require_grads
|
||||
from unsloth_zoo.log import logger
|
||||
logger.info(
|
||||
"Unsloth: Patched enable_input_require_grads for vision model compatibility"
|
||||
)
|
||||
|
||||
print([mod for mod in ["trl", "transformers", "peft"] if mod in sys.modules])
|
||||
|
||||
def torchvision_compatibility_check():
|
||||
if importlib.util.find_spec("torch") is None:
|
||||
|
|
@ -371,6 +373,7 @@ def torchvision_compatibility_check():
|
|||
f"but found torchvision=={torchvision_version}. "
|
||||
f"Please refer to https://pytorch.org/get-started/previous-versions/ for more information."
|
||||
)
|
||||
from unsloth_zoo.log import logger
|
||||
logger.info(
|
||||
f"Unsloth: torch=={torch_version} and torchvision=={torchvision_version} are compatible."
|
||||
)
|
||||
|
|
@ -385,6 +388,7 @@ def fix_openenv_no_vllm():
|
|||
openenv = Path(trl_location) / "experimental" / "openenv" / "utils.py"
|
||||
if not openenv.exists():
|
||||
return
|
||||
from unsloth_zoo.log import logger
|
||||
try:
|
||||
with open(openenv, "r+", encoding = "utf-8") as f:
|
||||
text = f.read()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue