diff --git a/unsloth/__init__.py b/unsloth/__init__.py index ea2fe76858..6a2d999b41 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -61,10 +61,10 @@ except: pass # Hugging Face Hub faster downloads (only enable during Colab and Kaggle sessions) -# keynames = "\n" + "\n".join(os.environ.keys()) -# if "\nCOLAB_" in keynames or "\nKAGGLE_" in keynames: -# os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1" -# pass +keynames = "\n" + "\n".join(os.environ.keys()) +if "\nCOLAB_" in keynames or "\nKAGGLE_" in keynames: + os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1" +pass # We support Pytorch 2 # Fixes https://github.com/unslothai/unsloth/issues/38 diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 9bea364ca4..0fcfe2a270 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -335,6 +335,9 @@ def LlamaAttention_fast_forward( if past_key_value is not None: 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) + if position_ids is None: cos = self.rotary_emb.cos_cached sin = self.rotary_emb.sin_cached @@ -971,19 +974,21 @@ class LlamaRotaryEmbedding(torch.nn.Module): self.dim = dim self.max_position_embeddings = 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(4 * 8192, self.max_position_embeddings) # Build here to make `torch.jit.trace` work. - self._set_cos_sin_cache(seq_len=max_position_embeddings, device=device, dtype=torch.get_default_dtype()) + self._set_cos_sin_cache(seq_len=self.current_rope_size, device=device, dtype=torch.get_default_dtype()) 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.max_seq_len_cached = seq_len + self.current_rope_size = seq_len inv_freq = 1.0 / ( self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64, device="cpu").float() / self.dim) ) - t = torch.arange(self.max_seq_len_cached, device="cpu", dtype=torch.int64).float() + t = torch.arange(self.current_rope_size, device="cpu", dtype=torch.int64).float() freqs = torch.outer(t, inv_freq) # Different from paper, but it uses a different permutation in order to obtain the same calculation @@ -994,14 +999,21 @@ class LlamaRotaryEmbedding(torch.nn.Module): def forward(self, x, position_ids=None, seq_len=None): # x: [bs, num_attention_heads, seq_len, head_size] - if seq_len > self.max_seq_len_cached: + if seq_len > self.current_rope_size: self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=x.dtype) return ( - self.cos_cached[:seq_len].to(dtype=x.dtype), - self.sin_cached[:seq_len].to(dtype=x.dtype), + self.cos_cached[:seq_len].to(dtype = x.dtype), + self.sin_cached[:seq_len].to(dtype = x.dtype), ) 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 = int(round(seq_len / 8192)) * 8192 + self._set_cos_sin_cache(self.current_rope_size, device = "cuda:0", dtype = x.dtype) + pass pass @@ -1016,11 +1028,11 @@ class LlamaLinearScalingRotaryEmbedding(LlamaRotaryEmbedding): pass def _set_cos_sin_cache(self, seq_len, device, dtype): - self.max_seq_len_cached = seq_len + self.current_rope_size = seq_len inv_freq = 1.0 / ( self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64, device="cpu").float() / self.dim) ) - t = torch.arange(self.max_seq_len_cached, device="cpu", dtype=torch.int64).float() + t = torch.arange(self.current_rope_size, device="cpu", dtype=torch.int64).float() t = t / self.scaling_factor freqs = torch.outer(t, inv_freq)