课程设计报告实战指南:从系统设计到数据库建模
2026/10/3 9:28:33
目标:
对比文本分类不同训练方法效果。
示意图:
data.csv text,label 这家餐厅服务很好,1今天的体验非常糟糕,0产品质量非常不错,1再也不会购买这个产品,0importcopyimportrandomimporttimefromcollectionsimportCounterimportnumpyasnpimportpandasaspdimporttorchimporttorch.nnasnnfromsklearn.model_selectionimporttrain_test_splitfromsklearn.metricsimport(accuracy_score,precision_recall_fscore_support)fromtorch.utils.dataimportDataset,DataLoader# ========================================# 1. 全局配置# ========================================SEED=42MAX_LEN=64BATCH_SIZE=32EMBED_DIM=128HIDDEN_DIM=128EPOCHS=10LR=0.001DEVICE=("cuda"iftorch.cuda.is_available()else"mps"iftorch.backends.mps.is_available()else"cpu")defset_seed(seed):random.seed(seed)np.random.seed(seed)torch.manual_seed(seed)iftorch.cuda.is_available():torch.cuda.manual_seed_all(seed)set_seed(SEED)print("Device:",DEVICE)# ========================================# 2. 数据加载和划分# ========================================df=pd.read_csv("data.csv")df=df[["text","label"]].dropna().copy()df["text"]=df["text"].astype(str)# 保证标签从 0 开始且连续labels=sorted(df["label"].unique())label_map={label:ifori,labelinenumerate(labels)}df["label"]=df["label"].map(label_map)num_classes=len(label_map)# 先划分训练集和临时数据集train_df,temp_df=train_test_split(df,test_size=0.2,random_state=SEED,stratify=df["label"])# 临时数据集平均分为验证集、测试集val_df,test_df=train_test_split(temp_df,test_size=0.5,random_state=SEED,stratify=temp_df["label"])print("训练集:",len(train_df))print("验证集:",len(val_df))print("测试集:",len(test_df))# ========================================# 3. 构建字符级词表# ========================================counter=Counter(chfortextintrain_df["text"]forchintext)PAD=0UNK=1vocab={"<PAD>":PAD,"<UNK>":UNK}forcharincounter:vocab[char]=len(vocab)vocab_size=len(vocab)print("词表大小:",vocab_size)defencode(text):tokens=[vocab.get(ch,UNK)forchintext[:MAX_LEN]]length=len(tokens)tokens+=[PAD]*(MAX_LEN-length)returntokens,length# ========================================# 4. Dataset# ========================================classTextDataset(Dataset):def__init__(self,dataframe):self.texts=dataframe["text"].tolist()self.labels=dataframe["label"].tolist()def__len__(self):returnlen(self.texts)def__getitem__(self,index):tokens,length=encode(self.texts[index])return(torch.tensor(tokens,dtype=torch.long),torch.tensor(length,dtype=torch.long),torch.tensor(self.labels[index],dtype=torch.long))train_loader=DataLoader(TextDataset(train_df),batch_size=BATCH_SIZE,shuffle=True)val_loader=DataLoader(TextDataset(val_df),batch_size=BATCH_SIZE)test_loader=DataLoader(TextDataset(test_df),batch_size=BATCH_SIZE)# ========================================# 5. 模型一:MLP# ========================================classMLPClassifier(nn.Module):def__init__(self):super().__init__()self.embedding=nn.Embedding(vocab_size,EMBED_DIM,padding_idx=PAD)self.fc=nn.Sequential(nn.Linear(EMBED_DIM,HIDDEN_DIM),nn.ReLU(),nn.Dropout(0.2),nn.Linear(HIDDEN_DIM,num_classes))defforward(self,x,lengths):emb=self.embedding(x)# 不将 Padding 位置计入平均值mask=(x!=PAD).unsqueeze(-1)emb=emb*mask pooled=emb.sum(dim=1)/lengths.clamp(min=1).unsqueeze(1)returnself.fc(pooled)# ========================================# 6. 模型二:TextCNN# ========================================classTextCNN(nn.Module):def__init__(self):super().__init__()self.embedding=nn.Embedding(vocab_size,EMBED_DIM,padding_idx=PAD)self.convs=nn.ModuleList([nn.Conv1d(EMBED_DIM,64,kernel_size=k)forkin[2,3,4]])self.dropout=nn.Dropout(0.2)self.fc=nn.Linear(64*len(self.convs),num_classes)defforward(self,x,lengths):emb=self.embedding(x)# [B, T, E] -> [B, E, T]emb=emb.transpose(1,2)features=[]forconvinself.convs:feature=torch.relu(conv(emb))# 去掉完全位于 Padding 中的窗口k=conv.kernel_size[0]positions=torch.arange(feature.size(-1),device=x.device)valid=positions.unsqueeze(0)<(lengths-k+1).clamp(min=1).unsqueeze(1)feature=feature.masked_fill(~valid.unsqueeze(1),float("-inf"))pooled=feature.max(dim=-1).values features.append(pooled)out=torch.cat(features,dim=1)returnself.fc(self.dropout(out))# ========================================# 7. 模型三:BiLSTM# ========================================classBiLSTMClassifier(nn.Module):def__init__(self):super().__init__()self.embedding=nn.Embedding(vocab_size,EMBED_DIM,padding_idx=PAD)self.lstm=nn.LSTM(input_size=EMBED_DIM,hidden_size=HIDDEN_DIM,num_layers=1,batch_first=True,bidirectional=True)self.dropout=nn.Dropout(0.2)self.fc=nn.Linear(HIDDEN_DIM*2,num_classes)defforward(self,x,lengths):emb=self.embedding(x)packed=nn.utils.rnn.pack_padded_sequence(emb,lengths.cpu().clamp(min=1),batch_first=True,enforce_sorted=False)_,(hidden,_)=self.lstm(packed)# 拼接最后一层正向和反向隐藏状态out=torch.cat([hidden[-2],hidden[-1]],dim=1)returnself.fc(self.dropout(out))# ========================================# 8. 模型四:Transformer Encoder# ========================================classTransformerClassifier(nn.Module):def__init__(self):super().__init__()self.embedding=nn.Embedding(vocab_size,EMBED_DIM,padding_idx=PAD)self.position=nn.Embedding(MAX_LEN,EMBED_DIM)encoder_layer=nn.TransformerEncoderLayer(d_model=EMBED_DIM,nhead=4,dim_feedforward=256,dropout=0.2,batch_first=True)self.encoder=nn.TransformerEncoder(encoder_layer,num_layers=2,enable_nested_tensor=False)self.fc=nn.Linear(EMBED_DIM,num_classes)defforward(self,x,lengths):batch,seq_len=x.shape pos=torch.arange(seq_len,device=x.device).unsqueeze(0)emb=self.embedding(x)+self.position(pos)padding_mask=x==PAD out=self.encoder(emb,src_key_padding_mask=padding_mask)# Masked Mean Poolingvalid_mask=(~padding_mask).unsqueeze(-1)out=out*valid_mask pooled=out.sum(dim=1)/lengths.clamp(min=1).unsqueeze(1)returnself.fc(pooled)# ========================================# 9. 训练和评估# ========================================criterion=nn.CrossEntropyLoss()defsync_device():ifDEVICE=="cuda":torch.cuda.synchronize()elifDEVICE=="mps":torch.mps.synchronize()defevaluate(model,loader):model.eval()total_loss=0all_preds=[]all_labels=[]withtorch.no_grad():forx,lengths,yinloader:x=x.to(DEVICE)lengths=lengths.to(DEVICE)y=y.to(DEVICE)logits=model(x,lengths)loss=criterion(logits,y)total_loss+=loss.item()*len(y)preds=logits.argmax(dim=1)all_preds.extend(preds.cpu().tolist())all_labels.extend(y.cpu().tolist())precision,recall,f1,_=(precision_recall_fscore_support(all_labels,all_preds,average="macro",zero_division=0))return{"loss":total_loss/len(loader.dataset),"accuracy":accuracy_score(all_labels,all_preds),"precision":precision,"recall":recall,"f1":f1}deftrain_model(name,model):model=model.to(DEVICE)optimizer=torch.optim.AdamW(model.parameters(),lr=LR)best_f1=-1best_state=Nonesync_device()start_time=time.perf_counter()print(f"\n=========={name}==========")forepochinrange(EPOCHS):model.train()total_loss=0forx,lengths,yintrain_loader:x=x.to(DEVICE)lengths=lengths.to(DEVICE)y=y.to(DEVICE)optimizer.zero_grad()logits=model(x,lengths)loss=criterion(logits,y)loss.backward()nn.utils.clip_grad_norm_(model.parameters(),max_norm=1.0)optimizer.step()total_loss+=loss.item()*len(y)train_loss=(total_loss/len(train_loader.dataset))val_metrics=evaluate(model,val_loader)print(f"Epoch{epoch+1:02d}| "f"Train Loss:{train_loss:.4f}| "f"Val Acc:{val_metrics['accuracy']:.4f}| "f"Val F1:{val_metrics['f1']:.4f}")# 用验证集挑选最佳模型ifval_metrics["f1"]>best_f1:best_f1=val_metrics["f1"]best_state=copy.deepcopy(model.state_dict())sync_device()elapsed=time.perf_counter()-start_time# 加载验证集表现最好的模型model.load_state_dict(best_state)# 测试集只用于最终评估test_metrics=evaluate(model,test_loader)params=sum(p.numel()forpinmodel.parameters())torch.save(best_state,f"{name.lower()}.pth")result={"Model":name,"Accuracy":test_metrics["accuracy"],"Precision":test_metrics["precision"],"Recall":test_metrics["recall"],"Macro-F1":test_metrics["f1"],"Time(s)":elapsed,"Parameters":params}returnresult# ========================================# 10. 四种模型对比# ========================================models={"MLP":MLPClassifier(),"TextCNN":TextCNN(),"BiLSTM":BiLSTMClassifier(),"Transformer":TransformerClassifier()}results=[]forname,modelinmodels.items():set_seed(SEED)result=train_model(name,model)results.append(result)# ========================================# 11. 输出实验结果# ========================================result_df=pd.DataFrame(results)result_df.to_csv("comparison_results.csv",index=False,encoding="utf-8-sig")print("\n========== 最终结果 ==========")print(result_df.round(4).to_string(index=False))