AI手撕代码笔记

AI手撕代码笔记 AI手撕代码笔记Attn相关基础attncausal mask多头注意力cross attn辅助函数SoftmaxRLadvantage计算GAEAttn相关基础attn细节注意K的转秩注意mask写法是置为-1e9(masked_fill需要转bool取反)dropout在softmax之后softmax只做最后一个维度importtorchimporttorch.nnasnnfromtypingimportOptional,Tupleimporttorch.nn.functionalasFimportmathclassScaledDotProductAttention(nn.Module):def__init__(self,dropout_p:float0.0):super().__init__()self.dropoutnn.Dropout(dropout_p)defforward(self,q:torch.Tensor,k:torch.Tensor,v:torch.Tensor,mask:Optional[torch.Tensor]None,)-Tuple[torch.Tensor,torch.Tensor]:# q/k/v: [B, H, S, D]返回 (output, attn_weights)# print(q, k, v, mask)dq.shape[-1]scoresq k.transpose(-1,-2)/math.sqrt(d)ifmaskisnotNone:scoresscores.masked_fill_(~mask.bool(),-1e9)attn_weightsF.softmax(scores,dim-1)attn_weightsself.dropout(attn_weights)outputattn_weights vreturn(output,attn_weights)causal mask常规版本(多头注意力)defcreate_batch_causal_mask(batch_size:int,n_head:int,seq_len:int,deviceNone):输出 [B, H, L, L]causal_2dtorch.tril(torch.ones(seq_len,seq_len,dtypetorch.bool,devicedevice))maskcausal_2d[None,None,:,:].expand(batch_size,n_head,seq_len,seq_len)returnmask对于decoder部分如果以及有了一部分生成带滑窗便宜的casual maskq_idxtorch.arange(seq_len,devicex.device).unsqueeze(1)# [seq_len, 1]k_idxtorch.arange(total_len,devicex.device).unsqueeze(0)# [1, total_len]causal_maskk_idx(start_posq_idx)# [seq_len, total_len]多头注意力细节分头先reshape然后permute记得要转回来classMultiHeadAttention(nn.Module):def__init__(self,embed_dim:int,num_heads:int):super().__init__()assertembed_dim%num_heads0,embed_dim必须可以被num_heads整除self.embed_dimembed_dim self.num_headsnum_heads self.head_dimembed_dim//num_heads# Q K V 投影self.w_qnn.Linear(embed_dim,embed_dim)self.w_knn.Linear(embed_dim,embed_dim)self.w_vnn.Linear(embed_dim,embed_dim)# 输出投影self.w_onn.Linear(embed_dim,embed_dim)defforward(self,x,maskNone): Args: x: [B, L, C] 输入 mask: [B, 1, Lq, Lk] bool, True允许访问, Falsemask掉 Returns: out: [B, L, C] B,Lq,_x.shape# 1. 投影qself.w_q(x)# [B, Lq, C]kself.w_k(x)# [B, Lk, C]vself.w_v(x)Lkk.size(1)# 2. split heads: [B, H, L, head_dim]defsplit_head(t):B,L,Ct.shapereturnt.reshape(B,L,self.num_heads,self.head_dim).permute(0,2,1,3)qsplit_head(q)# [B, H, Lq, hd]ksplit_head(k)# [B, H, Lk, hd]vsplit_head(v)# [B, H, Lk, hd]# 3. scaled dot‑productscaleself.head_dim**(-0.5)attn_scoretorch.matmul(q,k.transpose(-1,-2))*scale# [B, H, Lq, Lk]# apply mask: False的位置填‑infsoftmax后权重为0ifmaskisnotNone:attn_scoreattn_score.masked_fill(~mask,float(-inf))attn_weightF.softmax(attn_score,dim-1)# [B, H, Lq, Lk]outtorch.matmul(attn_weight,v)# [B, H, Lq, hd]# concat headsoutout.permute(0,2,1,3).reshape(B,Lq,self.embed_dim)# [B, Lq, C]outself.w_o(out)returnout,attn_weightcross attn注意Q和K/V来源不一致其他与self-attn一样辅助函数Softmaxdefsoftmax(x,dim-1):e_xnp.exp(x-np.max(x,axis-1,keepdimsTrue))returne_x/np.sum(e_x,axis-1,keepdimsTrue)RLadvantage计算importtorchdefcompute_mc_advantage(rewards,dones,values,gamma0.99):Tlen(rewards)returnstorch.zeros_like(rewards)running_ret0.0fortinreversed(range(T)):ifdones[t]:running_ret0.0running_retrewards[t]gamma*running_ret returns[t]running_ret advreturns-values adv(adv-adv.mean())/(adv.std()1e-8)returnadv,returnsGAEimporttorchdefcompute_gae(rewards:torch.Tensor,dones:torch.Tensor,values:torch.Tensor,gamma:float0.99,lam:float0.95): Args: rewards: [T] 单段轨迹奖励 dones: [T] bool/float1代表episode结束 values: [T] critic预测的state value gamma: 折扣因子 lam: GAE λ参数 Returns: advantages: [T] returns: [T] advantages values Tlen(rewards)advantagestorch.zeros_like(rewards)last_advantage0.0fortinreversed(range(T)):# 下一个时刻value最后一步tT‑1没有next state → 0next_valvalues[t1].item()if(t1T)else0.0# TD‑error δ_tdeltarewards[t]gamma*next_val*(1.0-dones[t])-values[t].item()# GAE递推advantages[t]deltagamma*lam*(1.0-dones[t])*last_advantage last_advantageadvantages[t].item()returnsadvantagesvalues# 可选advantage标准化ppo训练必用advantages(advantages-advantages.mean())/(advantages.std()1e-8)returnadvantages,returnsimporttorchdefcompute_group_advantages(rewards):# GRPO组内优势无critic用组均值做基线group_meanrewards.mean()advantagesrewards-group_mean advantages(advantages-advantages.mean())/(advantages.std()1e-8)returnadvantagesdefgrpo_calc_loss(old_log_probs,new_log_probs,advantages,ref_log_probs,clip_epsilon0.2,kl_beta0.04):ratiotorch.exp(new_log_probs-old_log_probs)surr1ratio*advantages surr2torch.clamp(ratio,1-clip_epsilon,1clip_epsilon)*advantages policy_loss-torch.min(surr1,surr2).mean()# KL惩罚 ref || currentkltorch.exp(new_log_probs)*(new_log_probs-ref_log_probs)kl_divkl.mean()total_losspolicy_losskl_beta*kl_divreturntotal_lossdefmain():devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)# 超参group_size4clip_epsilon0.2kl_beta0.04update_epochs2total_iters50foriter_idxinrange(total_iters):# -------- rollout 采样得到这批数据不展开模型、采样细节 --------old_log_probstorch.randn(group_size,devicedevice)ref_log_probstorch.randn(group_size,devicedevice)rewardstorch.randn(group_size,devicedevice)# -------- 计算advantage --------advantagescompute_group_advantages(rewards)# -------- PPO‑clip 更新循环 --------for_inrange(update_epochs):new_log_probstorch.randn(group_size,devicedevice)# 模型前向不展开# -------- 计算loss --------lossgrpo_calc_loss(old_log_probs,new_log_probs,advantages,ref_log_probs,clip_epsilon,kl_beta)# -------- 优化器步骤不展开模型与optimizer定义 --------loss.backward()# optimizer.step()# optimizer.zero_grad()ifiter_idx%100:print(fiter{iter_idx}, loss:{loss.item():.3f})if__name____main__:main()