check stride

This commit is contained in:
Daniel Han 2025-07-23 02:52:29 -07:00
commit 64e4e26066
3 changed files with 17 additions and 14 deletions

View file

@ -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),

View file

@ -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,

View file

@ -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(