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