File size: 4,999 Bytes
d91766b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | import re
import torch
import torch.nn as nn
from diffulex.layer.linear import ReplicatedLinear
from diffulex.moe.layer.base import FusedMoE
from diffulex.utils.checkpoint import LoadContext, ResolvedWeight
class NaiveFusedMoE(FusedMoE):
"""
single device moe layer
"""
def __init__(
self,
hidden_size: int,
intermediate_size: int,
num_experts: int,
top_k: int,
*,
hidden_act: str = "silu",
norm_topk_prob: bool = True,
moe_gemm_impl: str = "triton",
num_shared_experts: int = 0,
shared_expert_intermediate_size: int | None = None,
) -> None:
super().__init__(
hidden_size,
intermediate_size,
num_experts,
top_k,
hidden_act=hidden_act,
norm_topk_prob=norm_topk_prob,
moe_gemm_impl=moe_gemm_impl,
num_shared_experts=num_shared_experts,
shared_expert_intermediate_size=shared_expert_intermediate_size,
)
self.gate = ReplicatedLinear(hidden_size, self.num_experts, bias=False)
self.w13 = nn.Parameter(
torch.empty(self.num_experts, self.intermediate_size * 2, hidden_size)
)
self.w2 = nn.Parameter(
torch.empty(self.num_experts, hidden_size, self.intermediate_size)
)
def forward(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
original_shape = hidden_states.shape
flat_hidden_states = hidden_states.reshape(-1, original_shape[-1])
router_logits = self.gate(flat_hidden_states)
topk_output = self.router(router_logits)
topk_weights = topk_output.weights
topk_ids = topk_output.ids
final_hidden_states = self.expert_gemm(
impl=self.moe_gemm_impl,
hidden_states=flat_hidden_states,
w13=self.w13,
w2=self.w2,
topk_ids=topk_ids,
topk_weights=topk_weights,
local_expert_start=0,
hidden_act=self.hidden_act,
)
final_hidden_states = final_hidden_states.reshape(original_shape)
final_hidden_states = self.add_shared_experts(final_hidden_states, hidden_states)
return final_hidden_states, router_logits
def load_w1(self, loaded_weight: torch.Tensor, expert_idx: int) -> None:
self.w13.data[expert_idx, 0 : self.intermediate_size].copy_(loaded_weight)
def load_w3(self, loaded_weight: torch.Tensor, expert_idx: int) -> None:
self.w13.data[expert_idx, self.intermediate_size : 2 * self.intermediate_size].copy_(loaded_weight)
def load_w2(self, loaded_weight: torch.Tensor, expert_idx: int) -> None:
self.w2.data[expert_idx].copy_(loaded_weight)
def resolve_checkpoint_weight(self, suffix: str, ctx: LoadContext) -> ResolvedWeight | None:
# Stacked format: experts.gate_proj.weight ([num_experts, ...])
stacked_match = re.fullmatch(r"experts\.(gate_proj|up_proj|down_proj)\.weight", suffix)
if stacked_match is not None:
proj_name = stacked_match.group(1)
return ResolvedWeight(
loader=lambda loaded_weight, proj_name=proj_name: self._load_stacked_expert(
loaded_weight, proj_name
)
)
# Individual expert format: experts.0.gate_proj.weight
match = re.fullmatch(r"experts\.(\d+)\.(gate_proj|up_proj|down_proj)\.weight", suffix)
if match is None:
return None
expert_idx = int(match.group(1))
proj_name = match.group(2)
if proj_name == "gate_proj":
return ResolvedWeight(
loader=lambda loaded_weight, expert_idx=expert_idx: self.load_w1(
loaded_weight,
expert_idx,
)
)
if proj_name == "up_proj":
return ResolvedWeight(
loader=lambda loaded_weight, expert_idx=expert_idx: self.load_w3(
loaded_weight,
expert_idx,
)
)
if proj_name == "down_proj":
return ResolvedWeight(
loader=lambda loaded_weight, expert_idx=expert_idx: self.load_w2(
loaded_weight,
expert_idx,
)
)
return None
def _load_stacked_expert(self, loaded_weight: torch.Tensor, proj_name: str) -> None:
if proj_name == "gate_proj":
self.w13.data[:, :self.intermediate_size].copy_(loaded_weight)
elif proj_name == "up_proj":
self.w13.data[:, self.intermediate_size:].copy_(loaded_weight)
elif proj_name == "down_proj":
if loaded_weight.shape == self.w2.data.shape:
self.w2.data.copy_(loaded_weight)
else:
self.w2.data.copy_(loaded_weight.transpose(1, 2).contiguous())
__all__ = ["NaiveFusedMoE"]
|