Bug fixes
This commit is contained in:
parent
3b908c36e8
commit
1576396cd0
2 changed files with 7 additions and 3 deletions
|
|
@ -285,7 +285,11 @@ if major_version >= 8:
|
|||
if _is_package_available("flash_attn"):
|
||||
# Check for CUDA linking errors "undefined symbol: _ZNK3c106SymIntltEl"
|
||||
try:
|
||||
from flash_attn.flash_attn_interface import flash_attn_cuda
|
||||
try:
|
||||
# See https://github.com/unslothai/unsloth/issues/1437
|
||||
from flash_attn.flash_attn_interface import flash_attn_gpu
|
||||
except:
|
||||
from flash_attn.flash_attn_interface import flash_attn_cuda
|
||||
HAS_FLASH_ATTENTION = True
|
||||
|
||||
# Also check for softcapping
|
||||
|
|
@ -799,7 +803,7 @@ def patch_linear_scaling(
|
|||
f"from typing import Union, Optional, List, Any, Callable, Tuple\n"\
|
||||
f"from {model_filepath} import logger, "\
|
||||
f"{model_name.title()}Attention, {model_name.title()}Config"
|
||||
|
||||
|
||||
try:
|
||||
function = inspect.getsource(attention_module.__init__)
|
||||
except:
|
||||
|
|
|
|||
|
|
@ -304,7 +304,7 @@ class FastMistralModel(FastLlamaModel):
|
|||
attention_module = MistralAttention,
|
||||
)
|
||||
# Just for Mistral Nemo models!
|
||||
if function is not None:
|
||||
if function is not None and init_name is not None:
|
||||
function = patch_mistral_nemo_attention(function)
|
||||
# if True:#init_name is not None:
|
||||
exec(function, globals())
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue