Compare commits

...

2 Commits

Author SHA1 Message Date
Claude
a465dd1d75 fix: use F.silu in test script for old corex torch 2026-08-15 05:22:31 +00:00
Claude
a12d070d82 fix: replace torch::silu with x*sigmoid(x) for old corex torch 2026-08-15 05:22:25 +00:00
2 changed files with 5 additions and 5 deletions

View File

@@ -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])

View File

@@ -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)