(
X, # pointer to the input
# Y, # pointer to the output to be recomputed
DY, # pointer to the output gradient
DX, # pointer to the input gradient
stride_x_row, # how much to increase the pointer when moving by 1 row
N, # number of columns in X
eps, # epsilon to avoid division by zero
BLOCK_N: tl.constexpr,
)
| 62 | # @triton.heuristics({"RECOMPUTE_OUTPUT": lambda args: args["Y"] is not None}) |
| 63 | @triton.jit |
| 64 | def _l2_norm_bwd_kernel( |
| 65 | X, # pointer to the input |
| 66 | # Y, # pointer to the output to be recomputed |
| 67 | DY, # pointer to the output gradient |
| 68 | DX, # pointer to the input gradient |
| 69 | stride_x_row, # how much to increase the pointer when moving by 1 row |
| 70 | N, # number of columns in X |
| 71 | eps, # epsilon to avoid division by zero |
| 72 | BLOCK_N: tl.constexpr, |
| 73 | ): |
| 74 | # Map the program id to the elements of X, DX, and DY it should compute. |
| 75 | # Map the program id to the row of X and Y it should compute. |
| 76 | row = tl.program_id(0) |
| 77 | X += row * stride_x_row |
| 78 | DX += row * stride_x_row |
| 79 | DY += row * stride_x_row |
| 80 | |
| 81 | # Y += row * stride_y_row |
| 82 | cols = tl.arange(0, BLOCK_N) |
| 83 | x = tl.load(X + cols, mask=cols < N, other=0.0).to(tl.float32) |
| 84 | x = tl.where(cols < N, x, 0.0) |
| 85 | var = tl.sum(x * x) |
| 86 | rstd = 1 / tl.sqrt(var + eps) |
| 87 | # tl.store(Rstd + row, rstd) |
| 88 | # Normalize and apply linear transformation |
| 89 | mask = cols < N |
| 90 | # y = x * rstd |
| 91 | dy = tl.load(DY + cols, mask=cols < N, other=0.0).to(tl.float32) |
| 92 | dy = tl.where(cols < N, dy, 0.0) |
| 93 | # dx = dy * rstd - tl.sum(dy * x) * (1 / (var+eps)) * rstd * x |
| 94 | dx = dy * rstd - tl.sum(dy * x) * (1 / (var+eps)) * rstd * x |
| 95 | tl.store(DX + cols, dx, mask=mask) |
| 96 | |
| 97 | |
| 98 | def _l2_norm_fwd( |
nothing calls this directly
no outgoing calls
no test coverage detected