diff --git a/README.md b/README.md index b23acffcb8..0590415f43 100644 --- a/README.md +++ b/README.md @@ -10,7 +10,7 @@ -### Finetune Llama 3.1, Mistral, Phi-3 & Gemma 2-5x faster with 80% less memory! +### Finetune Llama 3.1, Mistral, Phi-3.5 & Gemma 2-5x faster with 80% less memory! ![](https://i.ibb.co/sJ7RhGG/image-41.png) diff --git a/unsloth/chat_templates.py b/unsloth/chat_templates.py index 82f6aba148..f83df579b1 100644 --- a/unsloth/chat_templates.py +++ b/unsloth/chat_templates.py @@ -473,7 +473,7 @@ pass # =========================================== Phi-3 phi3_template = \ - "{{ bos_token }}"\ + # "{{ bos_token }}"\ # Phi-3.5 removes BOS? "{% for message in messages %}"\ "{% if message['role'] == 'user' %}"\ "{{'<|user|>\n' + message['content'] + '<|end|>\n'}}"\ @@ -505,7 +505,9 @@ PARAMETER stop "<|assistant|>" ''' phi3_template_eos_token = "<|end|>" -CHAT_TEMPLATES["phi-3"] = (phi3_template, phi3_template_eos_token, False, phi3_ollama,) +CHAT_TEMPLATES["phi-3"] = (phi3_template, phi3_template_eos_token, False, phi3_ollama,) +CHAT_TEMPLATES["phi-35"] = CHAT_TEMPLATES["phi-3"] +CHAT_TEMPLATES["phi-3.5"] = CHAT_TEMPLATES["phi-3"] pass # =========================================== Llama-3.1 diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index d8904aa12b..1c48e8e588 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -969,6 +969,7 @@ def patch_llama_rope_scaling( scaled_rope_module = None, extended_rope_module = None, attention_module = None, + longrope_module = None, ): assert(\ rope_module is not None and \ @@ -1026,14 +1027,26 @@ def patch_llama_rope_scaling( max_position_embeddings=self.max_position_embeddings, base=self.rope_theta, ) + elif scaling_type == "longrope": + self.rotary_emb = {longrope_rope_function}( + dim = self.head_dim, + max_position_embeddings = self.max_position_embeddings, + original_max_position_embeddings = self.config.original_max_position_embeddings, + base = self.rope_theta, + short_factor = self.config.rope_scaling['short_factor'], + long_factor = self.config.rope_scaling['long_factor' ], + ) else: raise ValueError(f"Unknown RoPE scaling type {{scaling_type}}") pass """ + fix_rope_function = fix_rope_function.format( rope_function = rope_module.__name__, scaled_rope_function = scaled_rope_module.__name__, extended_rope_function = extended_rope_module.__name__, + longrope_rope_function = \ + (longrope_module if longrope_module is not None else rope_module).__name__ ) rotary_emb = re.findall( "self.rotary_emb = .+?\)", function, diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index a0894ec7a0..45f14c1131 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -254,6 +254,10 @@ class GemmaFixedRotaryEmbedding(torch.nn.Module): ) pass + def get_cached(self, seq_len = None): + return self.cos_cached, self.sin_cached + pass + def extend_rope_embedding(self, x, seq_len): if seq_len <= self.current_rope_size: return # Iteratively grow by increments of 8192 diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 048ba69193..376b4b4ebc 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -187,8 +187,9 @@ def LlamaAttention_fast_forward_inference( # cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len) # Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids) - cos = self.rotary_emb.cos_cached[position_ids].unsqueeze(1) - sin = self.rotary_emb.sin_cached[position_ids].unsqueeze(1) + cos, sin = self.rotary_emb.get_cached(kv_seq_len) + cos = cos[position_ids].unsqueeze(1) + sin = sin[position_ids].unsqueeze(1) h = self.half_head_dim RH_Q = self.RH_Q @@ -346,14 +347,17 @@ def LlamaAttention_fast_forward( kv_seq_len += past_key_value[0].shape[-2] # Extend RoPE dynamically to fit in VRAM - self.rotary_emb.extend_rope_embedding(V, seq_len = kv_seq_len) + rotary_emb = self.rotary_emb + rotary_emb.extend_rope_embedding(V, seq_len = kv_seq_len) if position_ids is None: - cos = self.rotary_emb.cos_cached - sin = self.rotary_emb.sin_cached + # Useful for LongRoPE + cos, sin = rotary_emb.get_cached(kv_seq_len) + # cos = self.rotary_emb.cos_cached + # sin = self.rotary_emb.sin_cached Q, K = fast_rope_embedding(Q, K, cos, sin) else: - cos, sin = self.rotary_emb(V, seq_len = kv_seq_len) + cos, sin = rotary_emb(V, seq_len = kv_seq_len) Q, K = inplace_rope_embedding(Q, K, cos, sin, position_ids) pass @@ -1048,6 +1052,10 @@ class LlamaRotaryEmbedding(torch.nn.Module): ) pass + def get_cached(self, seq_len = None): + return self.cos_cached, self.sin_cached + pass + def extend_rope_embedding(self, x, seq_len): if seq_len <= self.current_rope_size: return # Iteratively grow by increments of 8192 @@ -1170,6 +1178,125 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module): ) pass + def get_cached(self, seq_len = None): + return self.cos_cached, self.sin_cached + pass + + def extend_rope_embedding(self, x, seq_len): + if seq_len <= self.current_rope_size: return + # Iteratively grow by increments of 8192 + self.current_rope_size = math.ceil(seq_len / 8192) * 8192 + self._set_cos_sin_cache(self.current_rope_size, device = "cuda:0", dtype = x.dtype) + pass +pass + + +class LongRopeRotaryEmbedding(torch.nn.Module): + # For Phi 3.5 128K https://huggingface.co/microsoft/Phi-3.5-mini-instruct/blob/main/modeling_phi3.py + def __init__(self, + dim = None, + max_position_embeddings = 131072, + original_max_position_embeddings = 4096, + base = 10000, + short_factor = None, + long_factor = None, + device = None, + config = None, # [TODO] Hack to pass in config - need to remove later + ): + super().__init__() + assert(short_factor is not None) + assert(long_factor is not None) + assert(type(original_max_position_embeddings) is int) + + if config is not None: + # [TODO] Hack to pass in config - need to remove later + base = config.rope_theta + partial_rotary_factor = config.partial_rotary_factor if hasattr(config, "partial_rotary_factor") else 1.0 + dim = int((config.hidden_size // config.num_attention_heads)) + device = "cuda" + max_position_embeddings = config.max_position_embeddings + pass + + self.dim = dim + self.max_position_embeddings = max_position_embeddings + self.original_max_position_embeddings = original_max_position_embeddings + self.base = base + # Dynamic RoPE we first set it to a max of 4 * 8192 tokens then we iteratively grow this + self.current_rope_size = min(original_max_position_embeddings, self.max_position_embeddings) + + # Long RoPE similar to RoPE except short sequences have 1 cos / sin + # and long sequences have another cos / sin + inv_freq_shape = torch.arange(0, self.dim, 2, dtype=torch.int64, device="cpu").float() / self.dim + short_factor = torch.tensor(short_factor, device = "cpu", dtype = torch.float32) + long_factor = torch.tensor(long_factor, device = "cpu", dtype = torch.float32) + short_inv_freq = 1.0 / (short_factor * self.base**inv_freq_shape) + long_inv_freq = 1.0 / (long_factor * self.base**inv_freq_shape) + + # Phi-3 Scale factor + scale = self.max_position_embeddings / self.original_max_position_embeddings + if scale <= 1.0: + scaling_factor = 1.0 + else: + scaling_factor = math.sqrt(1 + math.log(scale) / math.log(self.original_max_position_embeddings)) + pass + self.scaling_factor = scaling_factor + + # Short and long inv_freq + self.register_buffer("short_inv_freq", short_inv_freq, persistent = False) + self.register_buffer("long_inv_freq", long_inv_freq, persistent = False) + # Build here to make `torch.jit.trace` work. + # self._set_cos_sin_cache(seq_len=self.current_rope_size, device=device, dtype=torch.get_default_dtype()) + + # Short sequences + dtype = torch.bfloat16 if is_bfloat16_supported() else torch.float16 + t = torch.arange(original_max_position_embeddings, device=self.short_inv_freq.device, dtype=torch.int64).float() + freqs = torch.outer(t, self.short_inv_freq) + emb = torch.cat((freqs, freqs), dim=-1) + cos_cached = (emb.cos() * self.scaling_factor).to(dtype=dtype, device=device, non_blocking=True) + sin_cached = (emb.sin() * self.scaling_factor).to(dtype=dtype, device=device, non_blocking=True) + self.register_buffer("short_cos_cached", cos_cached, persistent=False) + self.register_buffer("short_sin_cached", sin_cached, persistent=False) + pass + + def _set_cos_sin_cache(self, seq_len, device, dtype): + # Note: on the original Llama codebase, these tensors are created on the target device (and not on CPU) and + # in FP32. They are applied (multiplied) in FP32 as well. + self.current_rope_size = seq_len + + t = torch.arange(self.current_rope_size, device=self.inv_freq.device, dtype=torch.int64).float() + # Long sequences + freqs = torch.outer(t, self.long_inv_freq) + emb = torch.cat((freqs, freqs), dim=-1) + cos_cached = (emb.cos() * self.scaling_factor).to(dtype=dtype, device=device, non_blocking=True) + sin_cached = (emb.sin() * self.scaling_factor).to(dtype=dtype, device=device, non_blocking=True) + self.register_buffer("long_cos_cached", cos_cached, persistent=False) + self.register_buffer("long_sin_cached", sin_cached, persistent=False) + pass + + def forward(self, x, position_ids=None, seq_len=None): + # x: [bs, num_attention_heads, seq_len, head_size] + if seq_len > self.current_rope_size: + self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=x.dtype) + + if seq_len < self.original_max_position_embeddings: + return ( + self.short_cos_cached[:seq_len].to(dtype = x.dtype), + self.short_sin_cached[:seq_len].to(dtype = x.dtype), + ) + else: + return ( + self.long_cos_cached[:seq_len].to(dtype = x.dtype), + self.long_sin_cached[:seq_len].to(dtype = x.dtype), + ) + pass + pass + + def get_cached(self, seq_len = None): + if seq_len < self.original_max_position_embeddings: + return self.short_cos_cached, self.short_sin_cached + return self.long_cos_cached, self.long_sin_cached + pass + def extend_rope_embedding(self, x, seq_len): if seq_len <= self.current_rope_size: return # Iteratively grow by increments of 8192 @@ -1242,6 +1369,7 @@ class FastLlamaModel: scaled_rope_module = LlamaLinearScalingRotaryEmbedding, extended_rope_module = LlamaExtendedRotaryEmbedding, attention_module = LlamaAttention, + longrope_module = LongRopeRotaryEmbedding, ) if init_name is not None: exec(function, globals()) @@ -1657,11 +1785,19 @@ class FastLlamaModel: pass pass # Downcast RoPE embedding to correct data type - if (name.endswith("rotary_emb") or hasattr(module, "cos_cached")) \ - and (module.cos_cached.dtype != correct_dtype): - - module.cos_cached = module.cos_cached.to(correct_dtype) - module.sin_cached = module.sin_cached.to(correct_dtype) + if (name.endswith("rotary_emb") or hasattr(module, "cos_cached")): + + if hasattr(module, "cos_cached") and \ + (module.cos_cached.dtype != correct_dtype): + + module.cos_cached = module.cos_cached.to(correct_dtype) + module.sin_cached = module.sin_cached.to(correct_dtype) + + elif hasattr(module, "short_cos_cached") and \ + (module.short_cos_cached.dtype != correct_dtype): + + module.short_cos_cached = module.short_cos_cached.to(correct_dtype) + module.short_sin_cached = module.short_sin_cached.to(correct_dtype) pass pass pass diff --git a/unsloth/models/mapper.py b/unsloth/models/mapper.py index b8259a073c..3f49c96551 100644 --- a/unsloth/models/mapper.py +++ b/unsloth/models/mapper.py @@ -249,6 +249,10 @@ __INT_TO_FLOAT_MAPPER = \ "unsloth/gemma-2-2b-it", "google/gemma-2-2b-it", ), + "unsloth/Phi-3.5-mini-instruct-bnb-4bit" : ( + "unsloth/Phi-3.5-mini-instruct", + "microsoft/Phi-3.5-mini-instruct", + ), } INT_TO_FLOAT_MAPPER = {}