Files
ModelHub XC 4244787e58 初始化项目,由ModelHub XC社区提供模型
Model: sabari2005/cyberslm-instruct
Source: Original Platform
2026-08-29 19:29:19 +08:00

128 lines
4.0 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
SwiGLU Feed-Forward Network (FFN)
===================================
Reference: "GLU Variants Improve Transformer" (Noam Shazeer, 2020)
https://arxiv.org/abs/2002.05202
Mathematical definition
-----------------------
Standard FFN (for contrast):
FFN(x) = activation(x W₁) W₂
SwiGLU FFN:
SwiGLU(x) = (x W_gate ⊙ swish(x W_gate)) W₂ ← WRONG shorthand
Correct form with separate gate and value projections:
gate(x) = x W_gate ∈ ^{B×T×ffn_dim}
val(x) = x W_val ∈ ^{B×T×ffn_dim}
hidden = swish(gate(x)) ⊙ val(x)
out = hidden W_out ∈ ^{B×T×hidden_dim}
where swish(z) = z · sigmoid(z) = z · σ(z).
Why SwiGLU
----------
- The gating mechanism (⊙) gives the network a multiplicative path to
control information flow — intuitively "how much" of each feature passes
through at each position.
- swish is smooth and non-monotonic, empirically outperforming ReLU and
GELU in large-scale experiments (PaLM, LLaMA, etc.).
- The factored W_gate / W_val structure adds one extra matrix but improves
quality relative to a standard 2-layer FFN of the same parameter budget.
Parameter count
---------------
Three matrices: W_gate, W_val, W_out
W_gate : hidden_dim × ffn_hidden_dim (384 × 1024 = 393 216)
W_val : hidden_dim × ffn_hidden_dim (384 × 1024 = 393 216)
W_out : ffn_hidden_dim × hidden_dim (1024 × 384 = 393 216)
Total per layer: 1 179 648 ≈ 1.18 M
No bias on any projection (consistent with modern practice).
Numerical stability
-------------------
- swish(z) = z · σ(z) is numerically safe for all real z.
- σ(z) = 1/(1+exp(z)) in PyTorch uses a numerically stable implementation.
- The element-wise product of two bounded (after sigmoid) quantities does
not amplify values explosively.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from cyberslm.model.config import CyberSLMConfig
class SwiGLUFeedForward(nn.Module):
"""
SwiGLU Feed-Forward Network.
Consists of three bias-free linear projections:
``gate_proj``, ``val_proj``, and ``out_proj``.
Parameters
----------
config : CyberSLMConfig
Validated model configuration.
Shape
-----
Input : ``(batch, seq_len, hidden_dim)``
Output : ``(batch, seq_len, hidden_dim)``
"""
def __init__(self, config: CyberSLMConfig) -> None:
super().__init__()
self.hidden_dim = config.hidden_dim
self.ffn_hidden_dim = config.ffn_hidden_dim
# Gate projection: produces the gating signal fed through swish.
self.gate_proj = nn.Linear(
config.hidden_dim, config.ffn_hidden_dim, bias=False
)
# Value projection: produces the value signal gated element-wise.
self.val_proj = nn.Linear(
config.hidden_dim, config.ffn_hidden_dim, bias=False
)
# Output projection: maps back to residual stream dimension.
self.out_proj = nn.Linear(
config.ffn_hidden_dim, config.hidden_dim, bias=False
)
def forward(self, x: Tensor) -> Tensor:
"""
Apply SwiGLU feed-forward transformation.
Parameters
----------
x : Tensor
Input of shape ``(batch, seq_len, hidden_dim)``.
Returns
-------
Tensor
Output of shape ``(batch, seq_len, hidden_dim)``.
"""
# gate: (B, T, ffn_hidden_dim)
# val: (B, T, ffn_hidden_dim)
gate = self.gate_proj(x)
val = self.val_proj(x)
# SwiGLU: swish(gate) ⊙ val
# F.silu is swish: silu(z) = z * sigmoid(z)
hidden = F.silu(gate) * val # (B, T, ffn_hidden_dim)
# Project back to hidden_dim
return self.out_proj(hidden) # (B, T, hidden_dim)
def extra_repr(self) -> str:
return (
f"hidden_dim={self.hidden_dim}, "
f"ffn_hidden_dim={self.ffn_hidden_dim}"
)