From 3e93a78794feb10914e9074a204e41fb901a3e42 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Fri, 15 Mar 2024 20:14:56 +1100 Subject: [PATCH] Update rope_embedding.py --- unsloth/kernels/rope_embedding.py | 50 ++++++++++++++++++------------- 1 file changed, 30 insertions(+), 20 deletions(-) diff --git a/unsloth/kernels/rope_embedding.py b/unsloth/kernels/rope_embedding.py index c1167393fb..7e5c2e9d8c 100644 --- a/unsloth/kernels/rope_embedding.py +++ b/unsloth/kernels/rope_embedding.py @@ -24,7 +24,7 @@ def _rope_embedding( Q, Q_row_stride, cos, cos_row_stride, sin, sin_row_stride, - seqlen, head_dim, + seqlen, head_dim, group_size, n_heads, BACKWARD_PASS: tl.constexpr, BLOCK_SIZE : tl.constexpr, ): @@ -34,7 +34,7 @@ def _rope_embedding( See our blog post for more info """ row_position = tl.program_id(0) - head_position = tl.program_id(1) + group_head_position = tl.program_id(1) col_offsets = tl.arange(0, BLOCK_SIZE) half_head_dim = head_dim // 2 mask = col_offsets < half_head_dim @@ -44,23 +44,25 @@ def _rope_embedding( 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) - tl.store(Q + row_position*Q_row_stride + head_position*head_dim + \ - half_head_dim*1 + col_offsets, - Q2*cos1 + Q1*sin1, mask = mask) + head_start = group_head_position * group_size + head_end = tl.math.min((head_start + group_size), n_heads) + + for i in range(head_start, head_end): + offs_q1 = row_position * Q_row_stride + i * head_dim + col_offsets + offs_q2 = row_position * Q_row_stride + i * head_dim + col_offsets + half_head_dim + + # For Gemma - sometimes RoPE must be done in float32 and not bfloat16 + Q1 = tl.load(Q + offs_q1, mask = mask, other = 0).to(sin1.dtype) + Q2 = tl.load(Q + offs_q2, mask = mask, other = 0).to(sin1.dtype) + + tl.store(Q + offs_q1, Q1*cos1 - Q2*sin1, mask = mask) + tl.store(Q + offs_q2, Q2*cos1 + Q1*sin1, mask = mask) + pass pass @@ -75,12 +77,16 @@ class Fast_RoPE_Embedding(torch.autograd.Function): # [TODO] Changing blocksize to head_dim//2 seems to have # some concurrency / un-deterministic issues. - BLOCK_SIZE, num_warps = calculate_settings(head_dim) # (head_dim//2) - _rope_embedding[(n_rows, n_heads,)]( + BLOCK_SIZE, num_warps = calculate_settings(head_dim//2) # (head_dim//2) + group_size = 4 # 4 or 8, too large group_size can hurt performance. + n_groups = triton.cdiv(n_heads, group_size) + + grid = (n_rows, n_groups, ) + _rope_embedding[grid]( Q, Q.stride(0), cos, cos.stride(0), sin, sin.stride(0), - seq_len, head_dim, + seq_len, head_dim, group_size, n_heads, BACKWARD_PASS = False, BLOCK_SIZE = BLOCK_SIZE, num_warps = num_warps, @@ -102,11 +108,15 @@ class Fast_RoPE_Embedding(torch.autograd.Function): cos = ctx.cos sin = ctx.sin - _rope_embedding[(n_rows, n_heads,)]( + group_size = 4 # 4 or 8, too large group_size can hurt performance. + n_groups = triton.cdiv(n_heads, group_size) + + grid = (n_rows, n_groups, ) + _rope_embedding[grid]( dY, dY .stride(0), cos, cos.stride(0), sin, sin.stride(0), - seq_len, head_dim, + seq_len, head_dim, group_size, n_heads, BACKWARD_PASS = True, BLOCK_SIZE = ctx.BLOCK_SIZE, num_warps = ctx.num_warps, @@ -164,4 +174,4 @@ def inplace_rope_embedding(Q, K, cos, sin, position_ids): Q = Slow_RoPE_Embedding.apply(Q, cos, sin, position_ids) K = Slow_RoPE_Embedding.apply(K, cos, sin, position_ids) return Q, K -pass +pass \ No newline at end of file