diff --git a/unsloth/__init__.py b/unsloth/__init__.py index 7080c92894..0cc76cd6f2 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -17,22 +17,16 @@ import importlib # Currently only supports 1 GPU, or else seg faults will occur. if "CUDA_VISIBLE_DEVICES" in os.environ: - device = os.environ["CUDA_VISIBLE_DEVICES"] - if not device.isdigit(): + devices = os.environ["CUDA_VISIBLE_DEVICES"] + # check if there are multiple cuda devices set in env + if not devices.isdigit(): + first_id = devices.split(',')[0] warnings.warn( - f"Unsloth: 'CUDA_VISIBLE_DEVICES' is currently {device} "\ - "but we require 'CUDA_VISIBLE_DEVICES=0'\n"\ - "We shall set it ourselves." + f"Unsloth: 'CUDA_VISIBLE_DEVICES' is currently {devices} \n"\ + "Multiple CUDA devices detected but we require a single device.\n"\ + f"We will override CUDA_VISIBLE_DEVICES to first device: {first_id}." ) - os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" - os.environ["CUDA_VISIBLE_DEVICES"] = "0" - elif "CUDA_DEVICE_ORDER" not in os.environ: - warnings.warn( - f"Unsloth: 'CUDA_DEVICE_ORDER' is not set "\ - "but we require 'CUDA_DEVICE_ORDER=PCI_BUS_ID'\n"\ - "We shall set it ourselves." - ) - os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" + os.environ["CUDA_VISIBLE_DEVICES"] = str(first_id) else: # warnings.warn("Unsloth: 'CUDA_VISIBLE_DEVICES' is not set. We shall set it ourselves.") os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" diff --git a/unsloth/kernels/rope_embedding.py b/unsloth/kernels/rope_embedding.py index c1167393fb..99a99dc320 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,