恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
week6
首页
资讯中心
/
week6
week6
发布时间:2026/10/3 6:46:50
目标对比文本分类不同训练方法效果。示意图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. 全局配置# SEED42MAX_LEN64BATCH_SIZE32EMBED_DIM128HIDDEN_DIM128EPOCHS10LR0.001DEVICE(cudaiftorch.cuda.is_available()elsempsiftorch.backends.mps.is_available()elsecpu)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. 数据加载和划分# dfpd.read_csv(data.csv)dfdf[[text,label]].dropna().copy()df[text]df[text].astype(str)# 保证标签从 0 开始且连续labelssorted(df[label].unique())label_map{label:ifori,labelinenumerate(labels)}df[label]df[label].map(label_map)num_classeslen(label_map)# 先划分训练集和临时数据集train_df,temp_dftrain_test_split(df,test_size0.2,random_stateSEED,stratifydf[label])# 临时数据集平均分为验证集、测试集val_df,test_dftrain_test_split(temp_df,test_size0.5,random_stateSEED,stratifytemp_df[label])print(训练集:,len(train_df))print(验证集:,len(val_df))print(测试集:,len(test_df))# # 3. 构建字符级词表# counterCounter(chfortextintrain_df[text]forchintext)PAD0UNK1vocab{PAD:PAD,UNK:UNK}forcharincounter:vocab[char]len(vocab)vocab_sizelen(vocab)print(词表大小:,vocab_size)defencode(text):tokens[vocab.get(ch,UNK)forchintext[:MAX_LEN]]lengthlen(tokens)tokens[PAD]*(MAX_LEN-length)returntokens,length# # 4. Dataset# classTextDataset(Dataset):def__init__(self,dataframe):self.textsdataframe[text].tolist()self.labelsdataframe[label].tolist()def__len__(self):returnlen(self.texts)def__getitem__(self,index):tokens,lengthencode(self.texts[index])return(torch.tensor(tokens,dtypetorch.long),torch.tensor(length,dtypetorch.long),torch.tensor(self.labels[index],dtypetorch.long))train_loaderDataLoader(TextDataset(train_df),batch_sizeBATCH_SIZE,shuffleTrue)val_loaderDataLoader(TextDataset(val_df),batch_sizeBATCH_SIZE)test_loaderDataLoader(TextDataset(test_df),batch_sizeBATCH_SIZE)# # 5. 模型一MLP# classMLPClassifier(nn.Module):def__init__(self):super().__init__()self.embeddingnn.Embedding(vocab_size,EMBED_DIM,padding_idxPAD)self.fcnn.Sequential(nn.Linear(EMBED_DIM,HIDDEN_DIM),nn.ReLU(),nn.Dropout(0.2),nn.Linear(HIDDEN_DIM,num_classes))defforward(self,x,lengths):embself.embedding(x)# 不将 Padding 位置计入平均值mask(x!PAD).unsqueeze(-1)embemb*mask pooledemb.sum(dim1)/lengths.clamp(min1).unsqueeze(1)returnself.fc(pooled)# # 6. 模型二TextCNN# classTextCNN(nn.Module):def__init__(self):super().__init__()self.embeddingnn.Embedding(vocab_size,EMBED_DIM,padding_idxPAD)self.convsnn.ModuleList([nn.Conv1d(EMBED_DIM,64,kernel_sizek)forkin[2,3,4]])self.dropoutnn.Dropout(0.2)self.fcnn.Linear(64*len(self.convs),num_classes)defforward(self,x,lengths):embself.embedding(x)# [B, T, E] - [B, E, T]embemb.transpose(1,2)features[]forconvinself.convs:featuretorch.relu(conv(emb))# 去掉完全位于 Padding 中的窗口kconv.kernel_size[0]positionstorch.arange(feature.size(-1),devicex.device)validpositions.unsqueeze(0)(lengths-k1).clamp(min1).unsqueeze(1)featurefeature.masked_fill(~valid.unsqueeze(1),float(-inf))pooledfeature.max(dim-1).values features.append(pooled)outtorch.cat(features,dim1)returnself.fc(self.dropout(out))# # 7. 模型三BiLSTM# classBiLSTMClassifier(nn.Module):def__init__(self):super().__init__()self.embeddingnn.Embedding(vocab_size,EMBED_DIM,padding_idxPAD)self.lstmnn.LSTM(input_sizeEMBED_DIM,hidden_sizeHIDDEN_DIM,num_layers1,batch_firstTrue,bidirectionalTrue)self.dropoutnn.Dropout(0.2)self.fcnn.Linear(HIDDEN_DIM*2,num_classes)defforward(self,x,lengths):embself.embedding(x)packednn.utils.rnn.pack_padded_sequence(emb,lengths.cpu().clamp(min1),batch_firstTrue,enforce_sortedFalse)_,(hidden,_)self.lstm(packed)# 拼接最后一层正向和反向隐藏状态outtorch.cat([hidden[-2],hidden[-1]],dim1)returnself.fc(self.dropout(out))# # 8. 模型四Transformer Encoder# classTransformerClassifier(nn.Module):def__init__(self):super().__init__()self.embeddingnn.Embedding(vocab_size,EMBED_DIM,padding_idxPAD)self.positionnn.Embedding(MAX_LEN,EMBED_DIM)encoder_layernn.TransformerEncoderLayer(d_modelEMBED_DIM,nhead4,dim_feedforward256,dropout0.2,batch_firstTrue)self.encodernn.TransformerEncoder(encoder_layer,num_layers2,enable_nested_tensorFalse)self.fcnn.Linear(EMBED_DIM,num_classes)defforward(self,x,lengths):batch,seq_lenx.shape postorch.arange(seq_len,devicex.device).unsqueeze(0)embself.embedding(x)self.position(pos)padding_maskxPAD outself.encoder(emb,src_key_padding_maskpadding_mask)# Masked Mean Poolingvalid_mask(~padding_mask).unsqueeze(-1)outout*valid_mask pooledout.sum(dim1)/lengths.clamp(min1).unsqueeze(1)returnself.fc(pooled)# # 9. 训练和评估# criterionnn.CrossEntropyLoss()defsync_device():ifDEVICEcuda:torch.cuda.synchronize()elifDEVICEmps:torch.mps.synchronize()defevaluate(model,loader):model.eval()total_loss0all_preds[]all_labels[]withtorch.no_grad():forx,lengths,yinloader:xx.to(DEVICE)lengthslengths.to(DEVICE)yy.to(DEVICE)logitsmodel(x,lengths)losscriterion(logits,y)total_lossloss.item()*len(y)predslogits.argmax(dim1)all_preds.extend(preds.cpu().tolist())all_labels.extend(y.cpu().tolist())precision,recall,f1,_(precision_recall_fscore_support(all_labels,all_preds,averagemacro,zero_division0))return{loss:total_loss/len(loader.dataset),accuracy:accuracy_score(all_labels,all_preds),precision:precision,recall:recall,f1:f1}deftrain_model(name,model):modelmodel.to(DEVICE)optimizertorch.optim.AdamW(model.parameters(),lrLR)best_f1-1best_stateNonesync_device()start_timetime.perf_counter()print(f\n{name})forepochinrange(EPOCHS):model.train()total_loss0forx,lengths,yintrain_loader:xx.to(DEVICE)lengthslengths.to(DEVICE)yy.to(DEVICE)optimizer.zero_grad()logitsmodel(x,lengths)losscriterion(logits,y)loss.backward()nn.utils.clip_grad_norm_(model.parameters(),max_norm1.0)optimizer.step()total_lossloss.item()*len(y)train_loss(total_loss/len(train_loader.dataset))val_metricsevaluate(model,val_loader)print(fEpoch{epoch1:02d}| fTrain Loss:{train_loss:.4f}| fVal Acc:{val_metrics[accuracy]:.4f}| fVal F1:{val_metrics[f1]:.4f})# 用验证集挑选最佳模型ifval_metrics[f1]best_f1:best_f1val_metrics[f1]best_statecopy.deepcopy(model.state_dict())sync_device()elapsedtime.perf_counter()-start_time# 加载验证集表现最好的模型model.load_state_dict(best_state)# 测试集只用于最终评估test_metricsevaluate(model,test_loader)paramssum(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)resulttrain_model(name,model)results.append(result)# # 11. 输出实验结果# result_dfpd.DataFrame(results)result_df.to_csv(comparison_results.csv,indexFalse,encodingutf-8-sig)print(\n 最终结果 )print(result_df.round(4).to_string(indexFalse))