diff --git a/unsloth/kernels/rope_embedding.py b/unsloth/kernels/rope_embedding.py index a9527520ab..c1167393fb 100644 --- a/unsloth/kernels/rope_embedding.py +++ b/unsloth/kernels/rope_embedding.py @@ -39,24 +39,28 @@ def _rope_embedding( half_head_dim = head_dim // 2 mask = col_offsets < half_head_dim - Q1 = tl.load(Q + row_position*Q_row_stride + head_position*head_dim + \ - half_head_dim*0 + col_offsets, mask = mask, other = 0) - Q2 = tl.load(Q + row_position*Q_row_stride + head_position*head_dim + \ - half_head_dim*1 + col_offsets, mask = mask, other = 0) sin1 = tl.load(sin + (row_position % seqlen)*sin_row_stride + \ half_head_dim*0 + col_offsets, mask = mask, other = 0) cos1 = tl.load(cos + (row_position % seqlen)*cos_row_stride + \ half_head_dim*0 + col_offsets, mask = mask, other = 0) + # For Gemma - sometimes RoPE must be done in float32 and not bfloat16 + Q1 = tl.load(Q + row_position*Q_row_stride + head_position*head_dim + \ + half_head_dim*0 + col_offsets, mask = mask, other = 0).to(sin1.dtype) + Q2 = tl.load(Q + row_position*Q_row_stride + head_position*head_dim + \ + half_head_dim*1 + col_offsets, mask = mask, other = 0).to(sin1.dtype) + if BACKWARD_PASS: # See our blog post for more info. sin1 = -sin1 pass tl.store(Q + row_position*Q_row_stride + head_position*head_dim + \ - half_head_dim*0 + col_offsets, Q1*cos1 - Q2*sin1, mask = mask) + half_head_dim*0 + col_offsets, + Q1*cos1 - Q2*sin1, mask = mask) tl.store(Q + row_position*Q_row_stride + head_position*head_dim + \ - half_head_dim*1 + col_offsets, Q2*cos1 + Q1*sin1, mask = mask) + half_head_dim*1 + col_offsets, + Q2*cos1 + Q1*sin1, mask = mask) pass diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index f4bb6d877e..bcd0e1abd9 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -156,9 +156,8 @@ def GemmaModel_fast_forward_inference( hidden_states = self.embed_tokens(input_ids) # 3072**0.5 = 55.5000 in bfloat16, whilst 55.4256 in float32 # 2048**0.5 = 45.2500 in bfloat16, whilst 45.2548 in float32 - # hidden_states *= torch.tensor(math_sqrt(self.config.hidden_size), dtype = hidden_states.dtype) - hidden_states *= math_sqrt(self.config.hidden_size) - + hidden_states *= torch.tensor(math_sqrt(self.config.hidden_size), dtype = hidden_states.dtype) + next_decoder_cache = [] for idx, decoder_layer in enumerate(self.layers): # Self Attention @@ -221,9 +220,12 @@ class GemmaFixedRotaryEmbedding(torch.nn.Module): radians_new = positions[..., None] / timescale[None, None, :] radians_new = radians_new.squeeze(0) - emb = torch.cat((radians_new, radians_new), dim=-1) - self.register_buffer("cos_cached", emb.cos().to(dtype=dtype, device=device, non_blocking=True), persistent=False) - self.register_buffer("sin_cached", emb.sin().to(dtype=dtype, device=device, non_blocking=True), persistent=False) + emb = torch.cat((radians_new, radians_new), dim = -1) + # We must do RoPE in float32! + cos = emb.cos().to(device = device, non_blocking = True)#, dtype = dtype) + sin = emb.sin().to(device = device, non_blocking = True)#, dtype = dtype) + self.register_buffer("cos_cached", cos, persistent = False) + self.register_buffer("sin_cached", sin, persistent = False) pass def forward(self, x, position_ids=None, seq_len=None): @@ -307,13 +309,14 @@ class FastGemmaModel(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): + # RoPE must be done in float32 for Gemma + # 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) - pass - pass + # module.cos_cached = module.cos_cached.to(correct_dtype) + # module.sin_cached = module.sin_cached.to(correct_dtype) + # pass + # pass pass # Add 1 to weight diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 4b9161a1a5..3f281a09c8 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -523,8 +523,8 @@ def LlamaModel_fast_forward( # inputs_embeds *= math_sqrt(self.config.hidden_size) # Ie 3072**0.5 = 55.5000 in bfloat16, whilst 55.4256 in float32 # & 2048**0.5 = 45.2500 in bfloat16, whilst 45.2548 in float32 - # inputs_embeds *= torch.tensor(math_sqrt(self.config.hidden_size), dtype = inputs_embeds.dtype) - inputs_embeds *= math_sqrt(self.config.hidden_size) + inputs_embeds *= torch.tensor(math_sqrt(self.config.hidden_size), dtype = inputs_embeds.dtype) + # inputs_embeds *= math_sqrt(self.config.hidden_size) if inputs_requires_grad: inputs_embeds.requires_grad_(True) pass