ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

week2

week2 目标尝试完成一个多分类任务的训练:一个随机向量哪一维数字最大就属于第几类。内容importtorchimporttorch.nnasnnimporttorch.optimasoptim# 1. 固定随机种子torch.manual_seed(42)# 2. 定义参数input_dim5num_classes5num_samples10000epochs100batch_size128# 3. 生成训练数据X_traintorch.rand(num_samples,input_dim)# 找出每个向量中最大值的位置y_traintorch.argmax(X_train,dim1)# 生成独立测试数据X_testtorch.rand(2000,input_dim)y_testtorch.argmax(X_test,dim1)# 4. 定义神经网络classModel(nn.Module):def__init__(self):super().__init__()self.fcnn.Linear(5,5)defforward(self,x):returnself.fc(x)modelModel()# 5. 定义损失函数和优化器criterionnn.CrossEntropyLoss()optimizeroptim.Adam(model.parameters(),lr0.01)# 6. 训练模型forepochinrange(epochs):model.train()indicestorch.randperm(num_samples)total_loss0total_correct0forstartinrange(0,num_samples,batch_size):idxindices[start:startbatch_size]X_batchX_train[idx]y_batchy_train[idx]# 前向传播outputsmodel(X_batch)# 计算损失losscriterion(outputs,y_batch)# 梯度清零optimizer.zero_grad()# 反向传播loss.backward()# 更新参数optimizer.step()total_lossloss.item()*len(idx)predstorch.argmax(outputs,dim1)total_correct(predsy_batch).sum().item()if(epoch1)%100:print(fEpoch{epoch1}, fLoss:{total_loss/num_samples:.4f}, fAccuracy:{total_correct/num_samples:.2%})# 7. 测试模型model.eval()withtorch.no_grad():outputsmodel(X_test)predictionstorch.argmax(outputs,dim1)accuracy(predictionsy_test).float().mean()print(f\n测试准确率:{accuracy.item():.2%})# 8. 预测新的随机向量xtorch.tensor([[0.12,0.35,0.91,0.43,0.28]])withtorch.no_grad():outputmodel(x)predictiontorch.argmax(output,dim1)print(f预测类别第{prediction.item()1}类)print(f真实类别第{torch.argmax(x).item()1}类)
返回列表