check stride
This commit is contained in:
parent
07c6cb8f08
commit
64e4e26066
3 changed files with 17 additions and 14 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue