embedding
decoder_block
lm_head
this part
can be stacked
for N times
group_0
group_1
group_2
expert_0
expert_1
expert_2
activation
weight
element-wise add
element-wise multiply
matmul
L
lookup
N
rmsnorm
S
silu
P
rope
H
head
R
router
C
concat
L
N
P
P
H
P
P
H
P
P
H
C
N
R
S
S
S
C
N
input_ids
seq_len
1867
318
352
embedding_lookup (L)
def
embedding_lookup(input_ids, embedding_table):
# input_ids: [seq_len]
# embedding_table: [vocab_size, hidden_size]
# embed_out: [seq_len, hidden_size]
embed_out = embedding_table[input_ids]
return
embed_out
embedding_table
hidden_size
vocab_size
embed_out
hidden_size
seq_len
rms_norm (N)
def
rms_norm(embed_out, gamma, eps=1e-6):
# embed_out: [seq_len, hidden_size]
# gamma: [hidden_size]
# x_sq: [seq_len, hidden_size]
x_sq = embed_out ** 2
# sum_sq: [seq_len, 1]
sum_sq = np.sum(x_sq, axis=-1, keepdims=True)
# mean_sq: [seq_len, 1]
mean_sq = sum_sq / hidden_size
# rms: [seq_len, 1]
rms = np.sqrt(mean_sq)
# denom: [seq_len, 1]
denom = rms + eps
# normed: [seq_len, hidden_size]
normed = embed_out / denom
# rms_out: [seq_len, hidden_size]
rms_out = normed * gamma
return
rms_out
gamma
hidden_size
rms_out
hidden_size
seq_len
q_proj_0
def
q_proj(rms_out, w_q_0):
# rms_out: [seq_len, hidden_size]
# w_q_0: [hidden_size, d_head, q_heads]
# q_0: [q_heads, seq_len, d_head]
q_0 = rms_out @ w_q_0
return
q_0
k_proj_0
def
k_proj(rms_out, w_k_0):
# rms_out: [seq_len, hidden_size]
# w_k_0: [hidden_size, d_head]
# k_0: [seq_len, d_head]
k_0 = rms_out @ w_k_0
return
k_0
v_proj_0
def
v_proj(rms_out, w_v_0):
# rms_out: [seq_len, hidden_size]
# w_v_0: [hidden_size, d_head]
# v_0: [seq_len, d_head]
v_0 = rms_out @ w_v_0
return
v_0
w_q_0
d_head
hidden_size
q_heads
w_k_0
d_head
hidden_size
w_v_0
d_head
hidden_size
q_0
d_head
seq_len
q_heads
k_0
d_head
seq_len
v_0
d_head
seq_len
rope_q_0 (P)
def
rope(q_0):
# q_0: [q_heads, seq_len, d_head]
# cos: [seq_len, d_head // 2]
# sin: [seq_len, d_head // 2]
cos, sin = rope_freq(seq_len, d_head)
# q_even: [q_heads, seq_len, d_head // 2]
q_even = q_0[..., 0::2]
# q_odd: [q_heads, seq_len, d_head // 2]
q_odd = q_0[..., 1::2]
# term1: [q_heads, seq_len, d_head // 2]
term1 = q_even * cos
# term2: [q_heads, seq_len, d_head // 2]
term2 = q_odd * sin
# q_rot1: [q_heads, seq_len, d_head // 2]
q_rot1 = term1 - term2
# term3: [q_heads, seq_len, d_head // 2]
term3 = q_even * sin
# term4: [q_heads, seq_len, d_head // 2]
term4 = q_odd * cos
# q_rot2: [q_heads, seq_len, d_head // 2]
q_rot2 = term3 + term4
# stacked: [q_heads, seq_len, d_head // 2, 2]
stacked = np.stack([q_rot1, q_rot2], axis=-1)
# q_rot_0: [q_heads, seq_len, d_head]
q_rot_0 = stacked.reshape(q_0.shape)
return
q_rot_0
rope_k_0 (P)
def
rope(k_0):
# k_0: [seq_len, d_head]
# cos: [seq_len, d_head // 2]
# sin: [seq_len, d_head // 2]
cos, sin = rope_freq(seq_len, d_head)
# k_even: [seq_len, d_head // 2]
k_even = k_0[..., 0::2]
# k_odd: [seq_len, d_head // 2]
k_odd = k_0[..., 1::2]
# term1: [seq_len, d_head // 2]
term1 = k_even * cos
# term2: [seq_len, d_head // 2]
term2 = k_odd * sin
# k_rot1: [seq_len, d_head // 2]
k_rot1 = term1 - term2
# term3: [seq_len, d_head // 2]
term3 = k_even * sin
# term4: [seq_len, d_head // 2]
term4 = k_odd * cos
# k_rot2: [seq_len, d_head // 2]
k_rot2 = term3 + term4
# stacked: [seq_len, d_head // 2, 2]
stacked = np.stack([k_rot1, k_rot2], axis=-1)
# k_rot_0: [seq_len, d_head]
k_rot_0 = stacked.reshape(k_0.shape)
return
k_rot_0
q_rot_0
d_head
seq_len
q_heads
k_rot_0
d_head
seq_len
attention_head_0 (H)
def
softmax(scores):
# scores: [q_heads, seq_len, seq_len]
# max_score: [q_heads, seq_len, 1]
max_score = np.max(scores, axis=-1, keepdims=True)
# diff: [q_heads, seq_len, seq_len]
diff = scores - max_score
# exp_scores: [q_heads, seq_len, seq_len]
exp_scores = np.exp(diff)
# sum_exp: [q_heads, seq_len, 1]
sum_exp = np.sum(exp_scores, axis=-1, keepdims=True)
# weights: [q_heads, seq_len, seq_len]
weights = exp_scores / sum_exp
return
weights
def
attention_head(q_rot_0, k_rot_0, v_0):
# q_rot_0: [q_heads, seq_len, d_head]
# k_rot_0: [seq_len, d_head]
# v_0: [seq_len, d_head]
# k_trans: [d_head, seq_len]
k_trans = k_rot_0.T
# dot_prod: [q_heads, seq_len, seq_len]
dot_prod = q_rot_0 @ k_trans
# scale: []
scale = np.sqrt(d_head)
# scores: [q_heads, seq_len, seq_len]
scores = dot_prod / scale
# weights: [q_heads, seq_len, seq_len]
weights = softmax(scores)
# head_outs: [q_heads, seq_len, d_head]
head_outs = weights @ v_0
# group_0_out: [seq_len, head_dim]
group_0_out = np.concatenate(head_outs, axis=-1)
return
group_0_out
group_0_out
head_dim
seq_len
q_proj_1
def
q_proj(rms_out, w_q_1):
# rms_out: [seq_len, hidden_size]
# w_q_1: [hidden_size, d_head, q_heads]
# q_1: [q_heads, seq_len, d_head]
q_1 = rms_out @ w_q_1
return
q_1
k_proj_1
def
k_proj(rms_out, w_k_1):
# rms_out: [seq_len, hidden_size]
# w_k_1: [hidden_size, d_head]
# k_1: [seq_len, d_head]
k_1 = rms_out @ w_k_1
return
k_1
v_proj_1
def
v_proj(rms_out, w_v_1):
# rms_out: [seq_len, hidden_size]
# w_v_1: [hidden_size, d_head]
# v_1: [seq_len, d_head]
v_1 = rms_out @ w_v_1
return
v_1
w_q_1
d_head
hidden_size
q_heads
w_k_1
d_head
hidden_size
w_v_1
d_head
hidden_size
q_1
d_head
seq_len
q_heads
k_1
d_head
seq_len
v_1
d_head
seq_len
rope_q_1 (P)
def
rope(q_1):
# q_1: [q_heads, seq_len, d_head]
# cos: [seq_len, d_head // 2]
# sin: [seq_len, d_head // 2]
cos, sin = rope_freq(seq_len, d_head)
# q_even: [q_heads, seq_len, d_head // 2]
q_even = q_1[..., 0::2]
# q_odd: [q_heads, seq_len, d_head // 2]
q_odd = q_1[..., 1::2]
# term1: [q_heads, seq_len, d_head // 2]
term1 = q_even * cos
# term2: [q_heads, seq_len, d_head // 2]
term2 = q_odd * sin
# q_rot1: [q_heads, seq_len, d_head // 2]
q_rot1 = term1 - term2
# term3: [q_heads, seq_len, d_head // 2]
term3 = q_even * sin
# term4: [q_heads, seq_len, d_head // 2]
term4 = q_odd * cos
# q_rot2: [q_heads, seq_len, d_head // 2]
q_rot2 = term3 + term4
# stacked: [q_heads, seq_len, d_head // 2, 2]
stacked = np.stack([q_rot1, q_rot2], axis=-1)
# q_rot_1: [q_heads, seq_len, d_head]
q_rot_1 = stacked.reshape(q_1.shape)
return
q_rot_1
rope_k_1 (P)
def
rope(k_1):
# k_1: [seq_len, d_head]
# cos: [seq_len, d_head // 2]
# sin: [seq_len, d_head // 2]
cos, sin = rope_freq(seq_len, d_head)
# k_even: [seq_len, d_head // 2]
k_even = k_1[..., 0::2]
# k_odd: [seq_len, d_head // 2]
k_odd = k_1[..., 1::2]
# term1: [seq_len, d_head // 2]
term1 = k_even * cos
# term2: [seq_len, d_head // 2]
term2 = k_odd * sin
# k_rot1: [seq_len, d_head // 2]
k_rot1 = term1 - term2
# term3: [seq_len, d_head // 2]
term3 = k_even * sin
# term4: [seq_len, d_head // 2]
term4 = k_odd * cos
# k_rot2: [seq_len, d_head // 2]
k_rot2 = term3 + term4
# stacked: [seq_len, d_head // 2, 2]
stacked = np.stack([k_rot1, k_rot2], axis=-1)
# k_rot_1: [seq_len, d_head]
k_rot_1 = stacked.reshape(k_1.shape)
return
k_rot_1
q_rot_1
d_head
seq_len
q_heads
k_rot_1
d_head
seq_len
attention_head_1 (H)
def
softmax(scores):
# scores: [q_heads, seq_len, seq_len]
# max_score: [q_heads, seq_len, 1]
max_score = np.max(scores, axis=-1, keepdims=True)
# diff: [q_heads, seq_len, seq_len]
diff = scores - max_score
# exp_scores: [q_heads, seq_len, seq_len]
exp_scores = np.exp(diff)
# sum_exp: [q_heads, seq_len, 1]
sum_exp = np.sum(exp_scores, axis=-1, keepdims=True)
# weights: [q_heads, seq_len, seq_len]
weights = exp_scores / sum_exp
return
weights
def
attention_head(q_rot_1, k_rot_1, v_1):
# q_rot_1: [q_heads, seq_len, d_head]
# k_rot_1: [seq_len, d_head]
# v_1: [seq_len, d_head]
# k_trans: [d_head, seq_len]
k_trans = k_rot_1.T
# dot_prod: [q_heads, seq_len, seq_len]
dot_prod = q_rot_1 @ k_trans
# scale: []
scale = np.sqrt(d_head)
# scores: [q_heads, seq_len, seq_len]
scores = dot_prod / scale
# weights: [q_heads, seq_len, seq_len]
weights = softmax(scores)
# head_outs: [q_heads, seq_len, d_head]
head_outs = weights @ v_1
# group_1_out: [seq_len, head_dim]
group_1_out = np.concatenate(head_outs, axis=-1)
return
group_1_out
group_1_out
head_dim
seq_len
q_proj_2
def
q_proj(rms_out, w_q_2):
# rms_out: [seq_len, hidden_size]
# w_q_2: [hidden_size, d_head, q_heads]
# q_2: [q_heads, seq_len, d_head]
q_2 = rms_out @ w_q_2
return
q_2
k_proj_2
def
k_proj(rms_out, w_k_2):
# rms_out: [seq_len, hidden_size]
# w_k_2: [hidden_size, d_head]
# k_2: [seq_len, d_head]
k_2 = rms_out @ w_k_2
return
k_2
v_proj_2
def
v_proj(rms_out, w_v_2):
# rms_out: [seq_len, hidden_size]
# w_v_2: [hidden_size, d_head]
# v_2: [seq_len, d_head]
v_2 = rms_out @ w_v_2
return
v_2
w_q_2
d_head
hidden_size
q_heads
w_k_2
d_head
hidden_size
w_v_2
d_head
hidden_size
q_2
d_head
seq_len
q_heads
k_2
d_head
seq_len
v_2
d_head
seq_len
rope_q_2 (P)
def
rope(q_2):
# q_2: [q_heads, seq_len, d_head]
# cos: [seq_len, d_head // 2]
# sin: [seq_len, d_head // 2]
cos, sin = rope_freq(seq_len, d_head)
# q_even: [q_heads, seq_len, d_head // 2]
q_even = q_2[..., 0::2]
# q_odd: [q_heads, seq_len, d_head // 2]
q_odd = q_2[..., 1::2]
# term1: [q_heads, seq_len, d_head // 2]
term1 = q_even * cos
# term2: [q_heads, seq_len, d_head // 2]
term2 = q_odd * sin
# q_rot1: [q_heads, seq_len, d_head // 2]
q_rot1 = term1 - term2
# term3: [q_heads, seq_len, d_head // 2]
term3 = q_even * sin
# term4: [q_heads, seq_len, d_head // 2]
term4 = q_odd * cos
# q_rot2: [q_heads, seq_len, d_head // 2]
q_rot2 = term3 + term4
# stacked: [q_heads, seq_len, d_head // 2, 2]
stacked = np.stack([q_rot1, q_rot2], axis=-1)
# q_rot_2: [q_heads, seq_len, d_head]
q_rot_2 = stacked.reshape(q_2.shape)
return
q_rot_2
rope_k_2 (P)
def
rope(k_2):
# k_2: [seq_len, d_head]
# cos: [seq_len, d_head // 2]
# sin: [seq_len, d_head // 2]
cos, sin = rope_freq(seq_len, d_head)
# k_even: [seq_len, d_head // 2]
k_even = k_2[..., 0::2]
# k_odd: [seq_len, d_head // 2]
k_odd = k_2[..., 1::2]
# term1: [seq_len, d_head // 2]
term1 = k_even * cos
# term2: [seq_len, d_head // 2]
term2 = k_odd * sin
# k_rot1: [seq_len, d_head // 2]
k_rot1 = term1 - term2
# term3: [seq_len, d_head // 2]
term3 = k_even * sin
# term4: [seq_len, d_head // 2]
term4 = k_odd * cos
# k_rot2: [seq_len, d_head // 2]
k_rot2 = term3 + term4
# stacked: [seq_len, d_head // 2, 2]
stacked = np.stack([k_rot1, k_rot2], axis=-1)
# k_rot_2: [seq_len, d_head]
k_rot_2 = stacked.reshape(k_2.shape)
return
k_rot_2
q_rot_2
d_head
seq_len
q_heads
k_rot_2
d_head
seq_len
attention_head_2 (H)
def
softmax(scores):
# scores: [q_heads, seq_len, seq_len]
# max_score: [q_heads, seq_len, 1]
max_score = np.max(scores, axis=-1, keepdims=True)
# diff: [q_heads, seq_len, seq_len]
diff = scores - max_score
# exp_scores: [q_heads, seq_len, seq_len]
exp_scores = np.exp(diff)
# sum_exp: [q_heads, seq_len, 1]
sum_exp = np.sum(exp_scores, axis=-1, keepdims=True)
# weights: [q_heads, seq_len, seq_len]
weights = exp_scores / sum_exp
return
weights
def
attention_head(q_rot_2, k_rot_2, v_2):
# q_rot_2: [q_heads, seq_len, d_head]
# k_rot_2: [seq_len, d_head]
# v_2: [seq_len, d_head]
# k_trans: [d_head, seq_len]
k_trans = k_rot_2.T
# dot_prod: [q_heads, seq_len, seq_len]
dot_prod = q_rot_2 @ k_trans
# scale: []
scale = np.sqrt(d_head)
# scores: [q_heads, seq_len, seq_len]
scores = dot_prod / scale
# weights: [q_heads, seq_len, seq_len]
weights = softmax(scores)
# head_outs: [q_heads, seq_len, d_head]
head_outs = weights @ v_2
# group_2_out: [seq_len, head_dim]
group_2_out = np.concatenate(head_outs, axis=-1)
return
group_2_out
group_2_out
head_dim
seq_len
concat (C)
def
concat_groups(group_0_out, group_1_out, group_2_out):
# group_0_out: [seq_len, head_dim]
# group_1_out: [seq_len, head_dim]
# group_2_out: [seq_len, head_dim]
# gqa_out: [seq_len, hidden_size]
gqa_out = np.concatenate([group_0_out, group_1_out, group_2_out], axis=-1)
return
gqa_out
gqa_out
hidden_size
seq_len
out_matmul
def
out_matmul(gqa_out, w_matmul):
# gqa_out: [seq_len, hidden_size]
# w_matmul: [hidden_size, hidden_size]
# gqa_block_out: [seq_len, hidden_size]
gqa_block_out = gqa_out @ w_matmul
return
gqa_block_out
w_matmul
hidden_size
hidden_size
residual_add (+)
def
residual_add(embed_out, gqa_block_out):
# embed_out: [seq_len, hidden_size]
# gqa_block_out: [seq_len, hidden_size]
# res_mid: [seq_len, hidden_size]
res_mid = embed_out + gqa_block_out
return
res_mid
gqa_block_out
hidden_size
seq_len
rms_norm (N)
def
rms_norm(res_mid, gamma, eps=1e-6):
# res_mid: [seq_len, hidden_size]
# gamma: [hidden_size]
# x_sq: [seq_len, hidden_size]
x_sq = res_mid ** 2
# sum_sq: [seq_len, 1]
sum_sq = np.sum(x_sq, axis=-1, keepdims=True)
# mean_sq: [seq_len, 1]
mean_sq = sum_sq / hidden_size
# rms: [seq_len, 1]
rms = np.sqrt(mean_sq)
# denom: [seq_len, 1]
denom = rms + eps
# normed: [seq_len, hidden_size]
normed = res_mid / denom
# rms_out: [seq_len, hidden_size]
rms_out = normed * gamma
return
rms_out
gamma
hidden_size
rms_out
hidden_size
seq_len
router (R)
def
softmax(top_logits):
# top_logits: [seq_len, top_k]
# max_logit: [seq_len, 1]
max_logit = np.max(top_logits, axis=-1, keepdims=True)
# diff: [seq_len, top_k]
diff = top_logits - max_logit
# exp_logits: [seq_len, top_k]
exp_logits = np.exp(diff)
# sum_exp: [seq_len, 1]
sum_exp = np.sum(exp_logits, axis=-1, keepdims=True)
# top_weights: [seq_len, top_k]
top_weights = exp_logits / sum_exp
return
top_weights
def
router(rms_out, w_router, top_k=2):
# rms_out: [seq_len, hidden_size]
# w_router: [hidden_size, num_experts]
# logits: [seq_len, num_experts]
logits = rms_out @ w_router
# neg_logits: [seq_len, num_experts]
neg_logits = -logits
# sorted_indices: [seq_len, num_experts]
sorted_indices = np.argsort(neg_logits, axis=-1)
# top_indices: [seq_len, top_k]
top_indices = sorted_indices[..., :top_k]
# top_logits: [seq_len, top_k]
top_logits = np.take_along_axis(logits, top_indices, axis=-1)
# top_weights: [seq_len, top_k]
top_weights = softmax(top_logits)
# router_weights: [seq_len, num_experts]
router_weights = np.zeros_like(logits)
# router_weights: [seq_len, num_experts]
np.put_along_axis(router_weights, top_indices, top_weights, axis=-1)
return
router_weights
w_router
num_experts
hidden_size
router_weights
num_experts
seq_len
0.6
0.4
0
0
0.7
0.3
0.5
0
0.5
The router activates different experts for different tokens.
Entry [i, j] = 0 means expert_j is not activated for token_i.
In moe_matmul, the 0 is multiplied by that expert's result to make it 0.
gate_proj_0
def
gate_proj(rms_out, w_gate_0):
# rms_out: [seq_len, hidden_size]
# w_gate_0: [hidden_size, inter_size]
# x_gate_0: [seq_len, inter_size]
x_gate_0 = rms_out @ w_gate_0
return
x_gate_0
up_proj_0
def
up_proj(rms_out, w_up_0):
# rms_out: [seq_len, hidden_size]
# w_up_0: [hidden_size, inter_size]
# x_up_0: [seq_len, inter_size]
x_up_0 = rms_out @ w_up_0
return
x_up_0
w_gate_0
inter_size
hidden_size
w_up_0
inter_size
hidden_size
x_gate_0
inter_size
seq_len
x_up_0
inter_size
seq_len
silu_0 (S)
def
silu(x_gate_0):
# x_gate_0: [seq_len, inter_size]
# neg_x: [seq_len, inter_size]
neg_x = -x_gate_0
# exp_neg: [seq_len, inter_size]
exp_neg = np.exp(neg_x)
# denom: [seq_len, inter_size]
denom = 1.0 + exp_neg
# x_act_0: [seq_len, inter_size]
x_act_0 = x_gate_0 / denom
return
x_act_0
elementwise_mul_0
def
elementwise_mul(x_act_0, x_up_0):
# x_act_0: [seq_len, inter_size]
# x_up_0: [seq_len, inter_size]
# x_inter_0: [seq_len, inter_size]
x_inter_0 = x_act_0 * x_up_0
return
x_inter_0
down_proj_0
def
down_proj(x_inter_0, w_down_0):
# x_inter_0: [seq_len, inter_size]
# w_down_0: [inter_size, hidden_size]
# x_down_0: [seq_len, hidden_size]
x_down_0 = x_inter_0 @ w_down_0
return
x_down_0
w_down_0
hidden_size
inter_size
x_down_0
hidden_size
seq_len
gate_proj_1
def
gate_proj(rms_out, w_gate_1):
# rms_out: [seq_len, hidden_size]
# w_gate_1: [hidden_size, inter_size]
# x_gate_1: [seq_len, inter_size]
x_gate_1 = rms_out @ w_gate_1
return
x_gate_1
up_proj_1
def
up_proj(rms_out, w_up_1):
# rms_out: [seq_len, hidden_size]
# w_up_1: [hidden_size, inter_size]
# x_up_1: [seq_len, inter_size]
x_up_1 = rms_out @ w_up_1
return
x_up_1
w_gate_1
inter_size
hidden_size
w_up_1
inter_size
hidden_size
x_gate_1
inter_size
seq_len
x_up_1
inter_size
seq_len
silu_1 (S)
def
silu(x_gate_1):
# x_gate_1: [seq_len, inter_size]
# neg_x: [seq_len, inter_size]
neg_x = -x_gate_1
# exp_neg: [seq_len, inter_size]
exp_neg = np.exp(neg_x)
# denom: [seq_len, inter_size]
denom = 1.0 + exp_neg
# x_act_1: [seq_len, inter_size]
x_act_1 = x_gate_1 / denom
return
x_act_1
elementwise_mul_1
def
elementwise_mul(x_act_1, x_up_1):
# x_act_1: [seq_len, inter_size]
# x_up_1: [seq_len, inter_size]
# x_inter_1: [seq_len, inter_size]
x_inter_1 = x_act_1 * x_up_1
return
x_inter_1
down_proj_1
def
down_proj(x_inter_1, w_down_1):
# x_inter_1: [seq_len, inter_size]
# w_down_1: [inter_size, hidden_size]
# x_down_1: [seq_len, hidden_size]
x_down_1 = x_inter_1 @ w_down_1
return
x_down_1
w_down_1
hidden_size
inter_size
x_down_1
hidden_size
seq_len
gate_proj_2
def
gate_proj(rms_out, w_gate_2):
# rms_out: [seq_len, hidden_size]
# w_gate_2: [hidden_size, inter_size]
# x_gate_2: [seq_len, inter_size]
x_gate_2 = rms_out @ w_gate_2
return
x_gate_2
up_proj_2
def
up_proj(rms_out, w_up_2):
# rms_out: [seq_len, hidden_size]
# w_up_2: [hidden_size, inter_size]
# x_up_2: [seq_len, inter_size]
x_up_2 = rms_out @ w_up_2
return
x_up_2
w_gate_2
inter_size
hidden_size
w_up_2
inter_size
hidden_size
x_gate_2
inter_size
seq_len
x_up_2
inter_size
seq_len
silu_2 (S)
def
silu(x_gate_2):
# x_gate_2: [seq_len, inter_size]
# neg_x: [seq_len, inter_size]
neg_x = -x_gate_2
# exp_neg: [seq_len, inter_size]
exp_neg = np.exp(neg_x)
# denom: [seq_len, inter_size]
denom = 1.0 + exp_neg
# x_act_2: [seq_len, inter_size]
x_act_2 = x_gate_2 / denom
return
x_act_2
elementwise_mul_2
def
elementwise_mul(x_act_2, x_up_2):
# x_act_2: [seq_len, inter_size]
# x_up_2: [seq_len, inter_size]
# x_inter_2: [seq_len, inter_size]
x_inter_2 = x_act_2 * x_up_2
return
x_inter_2
down_proj_2
def
down_proj(x_inter_2, w_down_2):
# x_inter_2: [seq_len, inter_size]
# w_down_2: [inter_size, hidden_size]
# x_down_2: [seq_len, hidden_size]
x_down_2 = x_inter_2 @ w_down_2
return
x_down_2
w_down_2
hidden_size
inter_size
x_down_2
hidden_size
seq_len
concat (C)
def
concat_experts(x_down_0, x_down_1, x_down_2):
# x_down_0: [seq_len, hidden_size]
# x_down_1: [seq_len, hidden_size]
# x_down_2: [seq_len, hidden_size]
# expert_outs: [num_experts, seq_len, hidden_size]
expert_outs = np.stack([x_down_0, x_down_1, x_down_2], axis=0)
return
expert_outs
expert_outs
hidden_size
seq_len
num_experts
moe_matmul (X)
def
moe_matmul(router_weights, expert_outs):
# router_weights: [seq_len, num_experts]
# expert_outs: [num_experts, seq_len, hidden_size]
# exp_trans: [seq_len, num_experts, hidden_size]
exp_trans = np.swapaxes(expert_outs, 0, 1)
# w_exp: [seq_len, 1, num_experts]
w_exp = np.expand_dims(router_weights, axis=1)
# moe_mat: [seq_len, 1, hidden_size]
moe_mat = w_exp @ exp_trans
# moe_out: [seq_len, hidden_size]
moe_out = np.squeeze(moe_mat, axis=1)
return
moe_out
moe_out
hidden_size
seq_len
residual_add (+)
def
residual_add(res_mid, moe_out):
# res_mid: [seq_len, hidden_size]
# moe_out: [seq_len, hidden_size]
# res_out: [seq_len, hidden_size]
res_out = res_mid + moe_out
return
res_out
moe_block_out
hidden_size
seq_len
rms_norm (N)
def
rms_norm(moe_block_out, gamma, eps=1e-6):
# moe_block_out: [seq_len, hidden_size]
# gamma: [hidden_size]
# x_sq: [seq_len, hidden_size]
x_sq = moe_block_out ** 2
# sum_sq: [seq_len, 1]
sum_sq = np.sum(x_sq, axis=-1, keepdims=True)
# mean_sq: [seq_len, 1]
mean_sq = sum_sq / hidden_size
# rms: [seq_len, 1]
rms = np.sqrt(mean_sq)
# denom: [seq_len, 1]
denom = rms + eps
# normed: [seq_len, hidden_size]
normed = moe_block_out / denom
# rms_out: [seq_len, hidden_size]
rms_out = normed * gamma
return
rms_out
gamma
hidden_size
rms_out
hidden_size
seq_len
logits_matmul
def
logits_matmul(rms_out, embedding_table_T):
# rms_out: [seq_len, hidden_size]
# embedding_table_T: [hidden_size, vocab_size]
# all_logits: [seq_len, vocab_size]
all_logits = rms_out @ embedding_table_T
return
all_logits
embedding_table.T
vocab_size
hidden_size
all_logits
vocab_size
seq_len