class SwiGLU(nn.Module): def __init__(self, d_model, d_ff): super().__init__() self.w1 = nn.Linear(d_model, d_ff) self.w2 = nn.Linear(d_model, d_ff) self.w3 = nn.Linear(d_ff, d_model) self.act = nn.SiLU() def forward(self, x): return self.w3(self.act(self.w1(x)) * self.w2(x))