Update llama.py
This commit is contained in:
parent
fdce7bff90
commit
a654779617
1 changed files with 9 additions and 5 deletions
|
|
@ -1065,9 +1065,15 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module):
|
|||
config = None, # [TODO] Hack to pass in config - need to remove later
|
||||
):
|
||||
super().__init__()
|
||||
print(1068)
|
||||
# if config is not None: return # [TODO] Hack to pass in config - need to remove later
|
||||
|
||||
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.base = base
|
||||
|
|
@ -1080,7 +1086,6 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module):
|
|||
)
|
||||
inv_freq = self.apply_scaling(inv_freq)
|
||||
self.register_buffer("inv_freq", inv_freq, persistent = False)
|
||||
print(1083)
|
||||
|
||||
# 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())
|
||||
|
|
@ -1090,7 +1095,6 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module):
|
|||
# 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
|
||||
print(1093)
|
||||
|
||||
t = torch.arange(self.current_rope_size, device="cpu", dtype=torch.int64).float()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue