Mistral Nemo (#782)

* Update __init__.py

* dynamic RoPE

* Update mistral.py

* Update llama.py

* Update tokenizer_utils.py

* Update mistral.py

* Update llama.py

* Update __init__.py

* Update flex_attention.py

* Update llama.py

* Update llama.py

* Mistral Nemo
This commit is contained in:
Daniel Han 2024-07-19 00:14:24 -07:00 committed by GitHub
commit bd2959d300
3 changed files with 54 additions and 6 deletions

View file

@ -65,8 +65,26 @@ logging.getLogger("transformers.tokenization_utils_base").setLevel(logging.CRITI
# =============================================
# Edits all Config files to enable RoPE Scaling for all models
from transformers import PretrainedConfig
# Transformers had to update for Mistral Nemo 12b since Attention is (5120, 4096) now.
def patch_mistral_nemo_config(config):
if "head_dim (" not in config:
add_head_dim = "If it is not specified, will default to `8`.\n"\
" head_dim (`int`, *optional*, defaults to `hidden_size // num_attention_heads`):\n"\
" The attention head dimension."
config = config.replace("If it is not specified, will default to `8`.", add_head_dim)
add_head_dim = "num_key_value_heads=8,\n head_dim=None,"
config = config.replace("num_key_value_heads=8,", add_head_dim)
add_head_dim = "self.sliding_window = sliding_window\n self.head_dim = head_dim or hidden_size // num_attention_heads\n"
config = config.replace("self.sliding_window = sliding_window", add_head_dim)
pass
return config
pass
from transformers import __version__ as transformers_version
from transformers import PretrainedConfig
model_architectures = ["llama", "mistral", "gemma", "gemma2", "qwen2",]
for model_name in model_architectures:
@ -87,8 +105,14 @@ for model_name in model_architectures:
r"\n self.rope_scaling = rope_scaling\n",
config,
)
exec(config, globals())
# Just for Mistral Nemo
if model_name == "mistral":
if Version(transformers_version) <= Version("4.42.4"):
config = patch_mistral_nemo_config(config)
pass
exec(config, globals())
exec(f"import {config_filepath}", globals())
exec(f"{config_filepath}.{config_filename} = {config_filename}", globals())
pass
@ -97,7 +121,6 @@ pass
# =============================================
# torch.cuda.amp.custom_fwd is deprecated >= 2.4
import torch
from packaging.version import 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
@ -748,7 +771,7 @@ def patch_linear_scaling(
"self.rotary_emb = .+?\)", function,
flags = re.DOTALL | re.MULTILINE,
)
if len(rotary_emb) == 0: return
if len(rotary_emb) == 0: return None, function
rotary_emb = rotary_emb[0]
function = function.replace(rotary_emb, fix_rope_function, 1)
function = exec_code + "\n\n" + function

View file

@ -1162,9 +1162,12 @@ class FastLlamaModel:
print(statistics)
# Warn about fast transfers
old_hf_transfer = os.environ.get("HF_HUB_ENABLE_HF_TRANSFER", "0")
if os.environ.get("HF_HUB_ENABLE_HF_TRANSFER", "0") == "1":
logger.warning_once("Unsloth: Fast downloading is enabled - ignore downloading bars which are red colored!")
print("Unsloth: Fast downloading is enabled - ignore downloading bars which are red colored!")
pass
# Return old flag
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = old_hf_transfer
model_patcher.pre_patch()
get_statistics() # For debugging - we use a download counter to see if environments are not breaking
@ -1247,6 +1250,8 @@ class FastLlamaModel:
attn_implementation = "eager",
**kwargs,
)
# Return old flag
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = old_hf_transfer
# We currently only support NVIDIA GPUs - AMD / Intel is a work in progress!
post_check = check_nvidia()

View file

@ -270,6 +270,24 @@ def MistralForCausalLM_fast_forward(
pass
# Transformers had to update for Mistral Nemo 12b since Attention is (5120, 4096) now.
def patch_mistral_nemo_attention(function):
function = function.replace(
"(self.head_dim * self.num_heads) != self.hidden_size",
"False",
)
function = function.replace(
"self.head_dim = self.hidden_size // self.num_heads",
"self.head_dim = config.head_dim",
)
function = function.replace(
"self.o_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)",
"self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False)",
)
return function
pass
class FastMistralModel(FastLlamaModel):
@staticmethod
@ -280,7 +298,9 @@ class FastMistralModel(FastLlamaModel):
scaled_rope_module = LlamaLinearScalingRotaryEmbedding,
attention_module = MistralAttention,
)
if init_name is not None:
# Just for Mistral Nemo models!
function = patch_mistral_nemo_attention(function)
if True:#init_name is not None:
exec(function, globals())
MistralAttention.__init__ = eval(init_name)
pass