MCPcopy Create free account
hub / github.com/OpenSparseLLMs/Linear-MoE / _l2_norm_bwd_kernel

Function _l2_norm_bwd_kernel

linear_moe/model/common_modules/l2norm.py:64–95  ·  view source on GitHub ↗
(
    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,
)

Source from the content-addressed store, hash-verified

62# @triton.heuristics({"RECOMPUTE_OUTPUT": lambda args: args["Y"] is not None})
63@triton.jit
64def _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
98def _l2_norm_fwd(

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected