Compare commits
2 Commits
c840c9159f
...
a465dd1d75
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a465dd1d75 | ||
|
|
a12d070d82 |
@@ -66,7 +66,7 @@ for k in range(top_k):
|
||||
eid = expert_ids[k].item()
|
||||
w = expert_weights[k].item()
|
||||
gate_up = F.linear(hidden, w13[eid])
|
||||
gate = torch.silu(gate_up[:, :I])
|
||||
gate = F.silu(gate_up[:, :I])
|
||||
up = gate_up[:, I:]
|
||||
act = gate * up
|
||||
expert_out = F.linear(act, w2[eid])
|
||||
@@ -131,7 +131,7 @@ for _ in range(3):
|
||||
eid = expert_ids[k].item()
|
||||
w = expert_weights[k].item()
|
||||
gate_up = F.linear(hidden, w13[eid])
|
||||
gate = torch.silu(gate_up[:, :I])
|
||||
gate = F.silu(gate_up[:, :I])
|
||||
up = gate_up[:, I:]
|
||||
act = gate * up
|
||||
out_py += w * F.linear(act, w2[eid])
|
||||
@@ -144,7 +144,7 @@ for _ in range(100):
|
||||
eid = expert_ids[k].item()
|
||||
w = expert_weights[k].item()
|
||||
gate_up = F.linear(hidden, w13[eid])
|
||||
gate = torch.silu(gate_up[:, :I])
|
||||
gate = F.silu(gate_up[:, :I])
|
||||
up = gate_up[:, I:]
|
||||
act = gate * up
|
||||
out_py += w * F.linear(act, w2[eid])
|
||||
|
||||
@@ -49,7 +49,7 @@ torch::Tensor moe_decode(
|
||||
auto gate_up = torch::mm(hidden, gate_up_weights[eid].t());
|
||||
|
||||
// SiLU and mul
|
||||
auto gate = torch::silu(gate_up.slice(1, 0, inter));
|
||||
auto gate_slice = gate_up.slice(1, 0, inter); auto gate = gate_slice * torch::sigmoid(gate_slice);
|
||||
auto up = gate_up.slice(1, inter, inter2);
|
||||
auto act = gate * up; // (1, I)
|
||||
|
||||
@@ -117,7 +117,7 @@ torch::Tensor moe_prefill(
|
||||
auto gate_up = torch::mm(tokens, gate_up_weights[eid].t());
|
||||
|
||||
// SiLU and mul
|
||||
auto gate = torch::silu(gate_up.slice(1, 0, inter));
|
||||
auto gate_slice = gate_up.slice(1, 0, inter); auto gate = gate_slice * torch::sigmoid(gate_slice);
|
||||
auto up = gate_up.slice(1, inter, inter2);
|
||||
auto act = gate * up; // (count, I)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user