NanoGPT
所有之前:这篇是关于NanoGPT项目的核心代码的总结和理解,我会持续把主流技术加入到这个项目中,并对核心代码进行解释。完整项目代码需要查询NanoGPT. 模型架构 Base 最基础的模型架构是一个Decoder-only的Transformer模型,代码实现参考CodeByMyself,在CodeByMyself的基础上,使用了更现代化的Rope位置编码和SwiGLU激活函数, 并且加入了Moe用于和FFN对比,可在配置文件中选择是否使用Moe。部分代码如下: # swiglu class ffn_swiglu(nn.Module): def __init__(self, d_model, d_hidden, dropout=0.1): super().__init__() #第一个 Linear 输出 2 * d_hidden self.w_gate_up = nn.Linear(d_model, 2 * d_hidden) #同时生成 gate 和 up self.w_down = nn.Linear(d_hidden, d_model) self.dropout = nn.Dropout(dropout) def forward(self, x): # x: [B, L, d_model] gate_up = self.w_gate_up(x) # [B, L, 2 * d_hidden] gate, up = gate_up.chunk(2, dim=-1) # each: [B, L, d_hidden] swiglu_out = F.silu(gate) * up # [B, L, d_hidden] out = self.w_down(self.dropout(swiglu_out)) # [B, L, d_model] return out # moe class moe(nn.Module): def __init__(self,n_expert,d_model,top_k=2,dropout=0.1): super().__init__() self.d_model = d_model self.n_expert = n_expert self.top_k = top_k self.dropout = dropout self.gate = nn.Linear(d_model,n_expert,bias=False) self.softmax = nn.Softmax(dim=-1) # self.experts = nn.ModuleList([ffn(self.d_model,self.d_model*4,dropout) for _ in range(n_expert)]) relu激活的expert self.experts = nn.ModuleList([ffn_swiglu(self.d_model,self.d_model*2,dropout) for _ in range(n_expert)]) def forward(self,x): b,t,d = x.shape assert d == self.d_model,f"输入维度和moe设置维度不匹配" x_flat = x.view(-1,d) N = x_flat.shape[0] gate_logits = self.gate(x_flat) topk_weights,topk_indices = torch.topk(gate_logits,self.top_k,dim=-1) topk_weights= self.softmax(topk_weights) out = torch.zeros_like(x_flat) for i,expert in enumerate(self.experts): mask = (topk_indices == i) if not mask.any(): continue token_indices,expert_pos = torch.where(mask) select_x = x_flat[token_indices] expert_out = expert(select_x) weights = topk_weights[token_indices, expert_pos] out.index_add_(0, token_indices, expert_out * weights.unsqueeze(1)) return out.view(b,t,d) def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0): """ 预计算旋转角度的复数表示(cos + i*sin) """ freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) t = torch.arange(end, device=freqs.device) freqs = torch.outer(t, freqs).float() # [end, dim//2] freqs_cos = freqs.cos() freqs_sin = freqs.sin() return freqs_cos, freqs_sin # 分开存储更便于后续操作,不使用torch.complex类型,因为部分系统不支持 SwiGLU激活函数的实现是通过将线性层的输出分成两部分,一部分作为gate,另一部分作为up,然后使用SILU函数对gate进行激活,并与up相乘得到最终的输出。 ...