Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
21 lines
683 B
Python
21 lines
683 B
Python
"""SwiGLU FFN: x [B,T,D] -> y [B,T,D]."""
|
|
from __future__ import annotations
|
|
|
|
import torch.nn.functional as F
|
|
from torch import nn
|
|
|
|
|
|
class SwiGLUMLP(nn.Module):
|
|
def __init__(self, hidden_size: int, intermediate_size: int):
|
|
super().__init__()
|
|
self.w1 = nn.Linear(hidden_size, intermediate_size, bias=False)
|
|
self.w3 = nn.Linear(hidden_size, intermediate_size, bias=False)
|
|
self.w2 = nn.Linear(intermediate_size, hidden_size, bias=False)
|
|
|
|
@classmethod
|
|
def from_config(cls, config) -> SwiGLUMLP:
|
|
return cls(config.hidden_size, config.intermediate_size)
|
|
|
|
def forward(self, x):
|
|
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|