From 64e4e26066117ff1315a92441874fcb1208b4092 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 23 Jul 2025 02:52:29 -0700 Subject: [PATCH] check stride --- unsloth/kernels/cross_entropy_loss.py | 9 ++++++--- unsloth/kernels/rms_layernorm.py | 20 ++++++++++---------- unsloth/models/_utils.py | 2 +- 3 files changed, 17 insertions(+), 14 deletions(-) diff --git a/unsloth/kernels/cross_entropy_loss.py b/unsloth/kernels/cross_entropy_loss.py index df331fcd91..25b26aaa16 100644 --- a/unsloth/kernels/cross_entropy_loss.py +++ b/unsloth/kernels/cross_entropy_loss.py @@ -107,7 +107,7 @@ _cross_entropy_forward = triton.heuristics( def _chunked_cross_entropy_forward( logits_ptr , - logits_row_stride , + logits_row_stride : tl.constexpr, loss_ptr , logsumexp_ptr , labels_ptr , @@ -191,9 +191,9 @@ _chunked_cross_entropy_forward = triton.heuristics( def _cross_entropy_backward( logits_ptr , - logits_row_stride , + logits_row_stride : tl.constexpr, dloss_ptr , - dloss_row_stride , + dloss_row_stride : tl.constexpr, logsumexp_ptr , labels_ptr , VOCAB_SIZE : tl.constexpr, @@ -301,6 +301,7 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): BLOCK_SIZE, num_warps = calculate_settings(vocab_size) logsumexp = torch.empty(n_rows, dtype = torch.float32, device = device) + print("logits.stride(0)", logits.stride(0)) with torch_gpu_device(device): _cross_entropy_forward[(n_rows,)]( logits, logits.stride(0), @@ -363,6 +364,8 @@ class Fast_CrossEntropyLoss(torch.autograd.Function): div, mod = divmod(vocab_size, BLOCK_SIZE) n_blocks : int = div + (mod != 0) + print("logits.stride(0) dY", logits.stride(0)) + print("dlosses.stride(0) dY", dlosses.stride(0)) with torch_gpu_device(dlosses.device): _cross_entropy_backward[(n_rows, n_blocks,)]( logits, logits.stride(0), diff --git a/unsloth/kernels/rms_layernorm.py b/unsloth/kernels/rms_layernorm.py index fba7e56a84..ec45c6033b 100644 --- a/unsloth/kernels/rms_layernorm.py +++ b/unsloth/kernels/rms_layernorm.py @@ -19,9 +19,9 @@ from .utils import calculate_settings, torch_gpu_device @triton.jit def _rms_layernorm_forward( - Y, Y_row_stride, - X, X_row_stride, - W, W_row_stride, + Y, Y_row_stride : tl.constexpr, + X, X_row_stride : tl.constexpr, + W, W_row_stride : tl.constexpr, r, r_row_stride : tl.constexpr, n_cols : tl.constexpr, eps : tl.constexpr, @@ -54,10 +54,10 @@ pass def _rms_layernorm_backward( - dY, dY_row_stride, - dX, dX_row_stride, - X, X_row_stride, - W, W_row_stride, + dY, dY_row_stride : tl.constexpr, + dX, dX_row_stride : tl.constexpr, + X, X_row_stride : tl.constexpr, + W, W_row_stride : tl.constexpr, r, r_row_stride : tl.constexpr, # dW, dW_row_stride, n_cols : tl.constexpr, @@ -106,9 +106,9 @@ _rms_layernorm_backward = triton.heuristics( @triton.jit def _gemma_rms_layernorm_forward( - Y, Y_row_stride, - X, X_row_stride, - W, W_row_stride, + Y, Y_row_stride : tl.constexpr, + X, X_row_stride : tl.constexpr, + W, W_row_stride : tl.constexpr, r, r_row_stride : tl.constexpr, n_cols : tl.constexpr, eps : tl.constexpr, diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index c445e98b07..883a05f82a 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -431,7 +431,7 @@ if DEVICE_TYPE == "cuda": "Unsloth: If you want to finetune Gemma 2, upgrade flash-attn to version 2.6.3 or higher!\n"\ "Newer versions support faster and less memory usage kernels for Gemma 2's attention softcapping!\n"\ "To update flash-attn, do the below:\n"\ - '\npip install --no-deps --upgrade "flash-attn>=2.6.3"' + '\npip install --no-deps --no-build-isolation --upgrade "flash-attn>=2.6.3"' ) except: print(