☰
week5
2026/10/3 16:11:12 网站建设 项目流程

目标:
训练基于transformer的单向语言模型,并完成文本生成。

示意图:

内容:

importmathimportosimporttorchimporttorch.nnasnnimporttorch.nn.functionalasF# ====================================# 1. 模型超参数# ====================================torch.manual_seed(42)batch_size=32block_size=64# 最大上下文长度n_embd=128# 词向量维度n_head=4# 注意力头数n_layer=3# Transformer 层数dropout=0.1learning_rate=3e-4max_iters=1500eval_interval=300device=("cuda"iftorch.cuda.is_available()else"mps"iftorch.backends.mps.is_available()else"cpu")print("当前设备:",device)# ====================================# 2. 准备训练数据# ====================================# 建议准备自己的中文语料 data.txt。# 没有文件时使用演示语料,仅用于验证训练流程。ifos.path.exists("data.txt"):withopen("data.txt","r",encoding="utf-8")asf:text=f.read()else:sentences=["今天天气很好,我们一起出去散步。","人工智能正在改变我们的生活。","深度学习是机器学习的重要分支。","自然语言处理可以帮助计算机理解文本。","神经网络可以通过训练学习数据中的规律。","Transformer使用注意力机制处理序列信息。","语言模型的任务是预测下一个字符。","机器学习需要大量数据进行训练。","学习编程需要不断地练习和思考。","使用PyTorch可以方便地构建深度学习模型。","我们正在学习如何训练一个语言模型。","文本生成是自然语言处理的重要应用。",]text="\n".join(sentences*200)# 字符级 Tokenizerchars=sorted(list(set(text)))vocab_size=len(chars)stoi={ch:ifori,chinenumerate(chars)}itos={i:chforch,iinstoi.items()}defencode(s):return[stoi[c]forcins]defdecode(ids):return"".join(itos[int(i)]foriinids)data=torch.tensor(encode(text),dtype=torch.long)# 训练集与验证集n=int(0.9*len(data))train_data=data[:n]val_data=data[n:]assertmin(len(train_data),len(val_data))>block_sizeprint("词表大小:",vocab_size)print("训练字符数:",len(train_data))# ====================================# 3. 构建训练批次# ====================================defget_batch(split):source=train_dataifsplit=="train"elseval_data ix=torch.randint(0,len(source)-block_size,(batch_size,))x=torch.stack([source[i:i+block_size]foriinix])# y 相对于 x 向右移动一个位置y=torch.stack([source[i+1:i+block_size+1]foriinix])returnx.to(device),y.to(device)# ====================================# 4. 多头因果自注意力# ====================================classCausalSelfAttention(nn.Module):def__init__(self):super().__init__()assertn_embd%n_head==0self.num_heads=n_head self.head_dim=n_embd//n_head self.qkv=nn.Linear(n_embd,3*n_embd)self.proj=nn.Linear(n_embd,n_embd)self.attn_dropout=nn.Dropout(dropout)self.resid_dropout=nn.Dropout(dropout)# 下三角因果 Maskself.register_buffer("mask",torch.tril(torch.ones(block_size,block_size)).view(1,1,block_size,block_size))defforward(self,x):B,T,C=x.shape# 一次线性映射生成 Q、K、Vq,k,v=self.qkv(x).split(C,dim=2)# [B, T, C] -> [B, Head, T, HeadDim]q=q.view(B,T,self.num_heads,self.head_dim).transpose(1,2)k=k.view(B,T,self.num_heads,self.head_dim).transpose(1,2)v=v.view(B,T,self.num_heads,self.head_dim).transpose(1,2)# Scaled Dot-Product Attentionscores=q @ k.transpose(-2,-1)scores=scores/math.sqrt(self.head_dim)# 遮挡未来 Tokenscores=scores.masked_fill(self.mask[:,:,:T,:T]==0,float("-inf"))weights=F.softmax(scores,dim=-1)weights=self.attn_dropout(weights)out=weights @ v# 拼接所有注意力头out=out.transpose(1,2).contiguous()out=out.view(B,T,C)returnself.resid_dropout(self.proj(out))# ====================================# 5. 前馈神经网络# ====================================classFeedForward(nn.Module):def__init__(self):super().__init__()self.net=nn.Sequential(nn.Linear(n_embd,4*n_embd),nn.GELU(),nn.Linear(4*n_embd,n_embd),nn.Dropout(dropout))defforward(self,x):returnself.net(x)# ====================================# 6. Transformer Block# ====================================classTransformerBlock(nn.Module):def__init__(self):super().__init__()self.ln1=nn.LayerNorm(n_embd)self.attn=CausalSelfAttention()self.ln2=nn.LayerNorm(n_embd)self.ffn=FeedForward()defforward(self,x):# Pre-Norm + 残差连接x=x+self.attn(self.ln1(x))x=x+self.ffn(self.ln2(x))returnx# ====================================# 7. GPT 单向语言模型# ====================================classMiniGPT(nn.Module):def__init__(self):super().__init__()self.token_embedding=nn.Embedding(vocab_size,n_embd)self.position_embedding=nn.Embedding(block_size,n_embd)self.blocks=nn.Sequential(*[TransformerBlock()for_inrange(n_layer)])self.ln_f=nn.LayerNorm(n_embd)self.lm_head=nn.Linear(n_embd,vocab_size)defforward(self,idx,targets=None):B,T=idx.shapeassertT<=block_size token_emb=self.token_embedding(idx)positions=torch.arange(T,device=idx.device)pos_emb=self.position_embedding(positions)x=token_emb+pos_emb x=self.blocks(x)x=self.ln_f(x)logits=self.lm_head(x)loss=NoneiftargetsisnotNone:B,T,V=logits.shape loss=F.cross_entropy(logits.reshape(B*T,V),targets.reshape(B*T))returnlogits,loss@torch.no_grad()defgenerate(self,idx,max_new_tokens=100,temperature=1.0,top_k=None):asserttemperature>0self.eval()for_inrange(max_new_tokens):# 只使用最后 block_size 个 Tokenidx_cond=idx[:,-block_size:]logits,_=self(idx_cond)# 获取最后一个位置的预测logits=logits[:,-1,:]# 温度调节logits=logits/temperature# Top-K 采样iftop_kisnotNone:k=min(top_k,logits.size(-1))values,_=torch.topk(logits,k)logits[logits<values[:,[-1]]]=-float("inf")probs=F.softmax(logits,dim=-1)next_token=torch.multinomial(probs,num_samples=1)idx=torch.cat([idx,next_token],dim=1)returnidx# ====================================# 8. 训练与验证# ====================================model=MiniGPT().to(device)print("模型参数量:",sum(p.numel()forpinmodel.parameters()))optimizer=torch.optim.AdamW(model.parameters(),lr=learning_rate)@torch.no_grad()defestimate_loss():model.eval()results={}forsplitin["train","val"]:losses=torch.zeros(20)foriinrange(20):xb,yb=get_batch(split)_,loss=model(xb,yb)losses[i]=loss.item()results[split]=losses.mean().item()model.train()returnresultsforstepinrange(max_iters):ifstep%eval_interval==0orstep==max_iters-1:losses=estimate_loss()print(f"step{step:4d}| "f"train loss:{losses['train']:.4f}| "f"val loss:{losses['val']:.4f}")xb,yb=get_batch("train")logits,loss=model(xb,yb)optimizer.zero_grad(set_to_none=True)loss.backward()# 梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(),1.0)optimizer.step()# ====================================# 9. 保存模型# ====================================torch.save({"model_state_dict":model.state_dict(),"stoi":stoi,"itos":itos,"block_size":block_size,"n_embd":n_embd,"n_head":n_head,"n_layer":n_layer,},"mini_gpt.pth")print("模型已保存:mini_gpt.pth")# ====================================# 10. 文本生成# ====================================prompt="人工智能"# 字符级词表无法识别训练集外的字符unknown=set(prompt)-set(stoi)ifunknown:raiseValueError(f"提示词包含词表外字符:{unknown}")context=torch.tensor([encode(prompt)],dtype=torch.long,device=device)generated=model.generate(context,max_new_tokens=100,temperature=0.8,top_k=10)print("\n========== 生成结果 ==========")print(decode(generated[0].tolist()))

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询