From a5e0867f89e04c78dd99be20e2963c09279b777d Mon Sep 17 00:00:00 2001 From: Kaitao Yang <21039614+ykaitao@users.noreply.github.com> Date: Wed, 18 Feb 2026 17:13:33 -0800 Subject: [PATCH] reduce code duplicaton by inheritting from LlamaRotaryEmbedding (#3878) * simplify_code_using_apply_time_scaling * modify LlamaRotaryEmbedding for better inheritance * reduce_code_duplication_LlamaExtendedRotaryEmbedding --- unsloth/models/llama.py | 157 +++++++++------------------------------- 1 file changed, 33 insertions(+), 124 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index dcaa770baa..a3fc9ab49a 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -1640,6 +1640,17 @@ class LlamaRotaryEmbedding(torch.nn.Module): self.multi_gpu_cos_cached = [None] * DEVICE_COUNT self.multi_gpu_sin_cached = [None] * DEVICE_COUNT + # Normal Llama-3 RoPE + inv_freq = 1.0 / ( + self.base + ** ( + torch.arange(0, self.dim, 2, dtype = torch.int64, device = "cpu").float() + / self.dim + ) + ) + inv_freq = self._apply_inv_freq_scaling(inv_freq) + self.register_buffer("inv_freq", inv_freq, persistent = False) + # Build here to make `torch.jit.trace` work. for device_idx in range(DEVICE_COUNT): self._set_cos_sin_cache( @@ -1656,22 +1667,24 @@ class LlamaRotaryEmbedding(torch.nn.Module): 1, device = get_current_device(), dtype = torch.get_default_dtype() ) + def _apply_inv_freq_scaling(self, inv_freq): + """Override to apply custom inv_freq scaling (e.g., extended RoPE).""" + return inv_freq + + def _apply_time_scaling(self, t): + """Override to apply custom time scaling (e.g., linear scaling).""" + return t + 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 - inv_freq = 1.0 / ( - self.base - ** ( - torch.arange(0, self.dim, 2, dtype = torch.int64, device = "cpu").float() - / self.dim - ) - ) t = torch.arange( - self.current_rope_size, device = "cpu", dtype = torch.int64 + self.current_rope_size, device = self.inv_freq.device, dtype = torch.int64 ).float() + t = self._apply_time_scaling(t) - freqs = torch.outer(t, inv_freq) + freqs = torch.outer(t, self.inv_freq) # Different from paper, but it uses a different permutation in order to obtain the same calculation emb = torch.cat((freqs, freqs), dim = -1) cos = emb.cos().to(dtype = dtype, device = device, non_blocking = True) @@ -1733,33 +1746,14 @@ class LlamaLinearScalingRotaryEmbedding(LlamaRotaryEmbedding): config = config, ) - def _set_cos_sin_cache(self, seq_len, device, dtype): - 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.current_rope_size, device = "cpu", dtype = torch.int64 - ).float() - t = t / self.scaling_factor - - freqs = torch.outer(t, inv_freq) - # Different from paper, but it uses a different permutation in order to obtain the same calculation - emb = torch.cat((freqs, freqs), dim = -1) - cos = emb.cos().to(dtype = dtype, device = device, non_blocking = True) - sin = emb.sin().to(dtype = dtype, device = device, non_blocking = True) - self.multi_gpu_cos_cached[device.index] = cos - self.multi_gpu_sin_cached[device.index] = sin - return cos, sin + def _apply_time_scaling(self, t): + """Apply linear scaling to time indices.""" + return t / self.scaling_factor # See https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/rotary_embedding.py#L736 # For Llama 3.1 -class LlamaExtendedRotaryEmbedding(torch.nn.Module): +class LlamaExtendedRotaryEmbedding(LlamaRotaryEmbedding): def __init__( self, dim = None, @@ -1768,101 +1762,16 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module): device = None, config = None, # [TODO] Hack to pass in config - need to remove later ): - super().__init__() - if config is not None: - # [TODO] Hack to pass in config - need to remove later - base = _get_rope_theta(config, default = base) - 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 = DEVICE_TYPE_TORCH - max_position_embeddings = config.max_position_embeddings - - 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) - self.multi_gpu_cos_cached = [None] * DEVICE_COUNT - self.multi_gpu_sin_cached = [None] * DEVICE_COUNT - - # Normal Llama-3 RoPE - inv_freq = 1.0 / ( - self.base - ** ( - torch.arange(0, self.dim, 2, dtype = torch.int64, device = "cpu").float() - / self.dim - ) + super().__init__( + dim = dim, + max_position_embeddings = max_position_embeddings, + base = base, + device = device, + config = config, ) - inv_freq = self.apply_scaling(inv_freq) - self.register_buffer("inv_freq", inv_freq, persistent = False) - - # Build here to make `torch.jit.trace` work. - for device_idx in range(DEVICE_COUNT): - self._set_cos_sin_cache( - seq_len = self.current_rope_size, - device = torch.device(device_idx), - dtype = torch.get_default_dtype(), - ) - - # dummy so that patch_utils doesn't fail for now - self.cos_cached = torch.empty( - 1, device = get_current_device(), dtype = torch.get_default_dtype() - ) - self.sin_cached = torch.empty( - 1, device = get_current_device(), dtype = torch.get_default_dtype() - ) - - 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() - - freqs = torch.outer(t, self.inv_freq) - # Different from paper, but it uses a different permutation in order to obtain the same calculation - emb = torch.cat((freqs, freqs), dim = -1) - cos = emb.cos().to(dtype = dtype, device = device, non_blocking = True) - sin = emb.sin().to(dtype = dtype, device = device, non_blocking = True) - self.multi_gpu_cos_cached[device.index] = cos - self.multi_gpu_sin_cached[device.index] = sin - return cos, sin - - def forward(self, x, position_ids = None, seq_len = None): - # x: [bs, num_attention_heads, seq_len, head_size] - if seq_len is not None and seq_len > self.current_rope_size: - self._set_cos_sin_cache(seq_len = seq_len, device = x.device, dtype = x.dtype) - device_index = x.device.index - return ( - self.multi_gpu_cos_cached[device_index][:seq_len], - self.multi_gpu_sin_cached[device_index][:seq_len], - ) - - def get_cached(self, seq_len = None, device_index = None): - if device_index is None: - device_index = get_current_device() - return self.multi_gpu_cos_cached[device_index], self.multi_gpu_sin_cached[ - device_index - ] - - 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 = ((seq_len // 8192) + ((seq_len % 8192) != 0)) * 8192 - for device_idx in range(DEVICE_COUNT): - self._set_cos_sin_cache( - self.current_rope_size, device = torch.device(device_idx), dtype = x.dtype - ) # From https://github.com/meta-llama/llama-models/blob/main/models/llama3_1/api/model.py#L41 - def apply_scaling(self, freqs: torch.Tensor): + def _apply_inv_freq_scaling(self, freqs: torch.Tensor): # Values obtained from grid search scale_factor = 8 low_freq_factor = 1