关于d2l中train_ch3以及softmax中图在pycharm中显示不出来的解决方案

发布时间:2026/7/23 9:41:08
关于d2l中train_ch3以及softmax中图在pycharm中显示不出来的解决方案 使用最新版d2l【1.0.3】的同学由于新版d2l包不包含第三章训练代码需要手动在【环境目录\d2l\Lib\site-packages\d2l\torch.py】文件中添加以下代码defevaluate_accuracy(net,data_iter:torch.utils.data.DataLoader):ifisinstance(net,torch.nn.Module):net.eval()metricAccumulator(2)withtorch.no_grad():forX,yindata_iter:metric.add(accuracy(net(X),y),y.numel())returnmetric[0]/metric[1]deftrain_epoch_ch3(net,train_iter,loss,updater):metricsAccumulator(3)ifisinstance(net,torch.nn.Module):net.train()forX,yintrain_iter:y_hatnet(X)lloss(y_hat,y)ifisinstance(updater,torch.optim.Optimizer):updater.zero_grad()l.mean().backward()updater.step()else:l.sum().backward()updater(X.shape[0])# number of Xs samplesmetrics.add(float(l.detach().sum()),accuracy(y_hat,y),y.numel())returnmetrics[0]/metrics[2],metrics[1]/metrics[2]deftrain_ch3(net,train_iter,test_iter,loss,num_epochs,updater):训练模型定义见第3章animatorAnimator(xlabelepoch,xlim[1,num_epochs],ylim[0.3,0.9],legend[train loss,train acc,test acc])foriinrange(num_epochs):train_metricstrain_epoch_ch3(net,train_iter,loss,updater)test_accevaluate_accuracy(net,test_iter)animator.add(i1,train_metrics(test_acc,))train_loss,train_acctrain_metricsasserttrain_loss0.5,train_lossasserttrain_acc1andtrain_acc0.7,train_accasserttest_acc1andtest_acc0.7,test_acc改动点将本书中的metrics.add(float(l.sum()), accuracy(y_hat, y), y.numel())改为metrics.add(float(l.detach().sum()), accuracy(y_hat, y), y.numel())同时针对使用在Pycharm专业版中使用Jupyter时没有图像输出或者图像一闪而过的情况需修改以下代码在【环境目录\d2l\Lib\site-packages\d2l\torch.py】文件中将【Animator】类中的display.display(self.fig)display.clear_output(waitTrue)注释掉即可。新的【Animator】类如下【注3.6中存在这样的问题也直接去掉就好了3.7中需要在上面提到的torch文件中手动注释】classAnimator:For plotting data in animation.def__init__(self,xlabelNone,ylabelNone,legendNone,xlimNone,ylimNone,xscalelinear,yscalelinear,fmts(-,m--,g-.,r:),nrows1,ncols1,figsize(3.5,2.5)):Defined in :numref:sec_utils# Incrementally plot multiple linesiflegendisNone:legend[]d2l.use_svg_display()self.fig,self.axesd2l.plt.subplots(nrows,ncols,figsizefigsize)ifnrows*ncols1:self.axes[self.axes,]# Use a lambda function to capture argumentsself.config_axeslambda:d2l.set_axes(self.axes[0],xlabel,ylabel,xlim,ylim,xscale,yscale,legend)self.X,self.Y,self.fmtsNone,None,fmtsdefadd(self,x,y):# Add multiple data points into the figureifnothasattr(y,__len__):y[y]nlen(y)ifnothasattr(x,__len__):x[x]*nifnotself.X:self.X[[]for_inrange(n)]ifnotself.Y:self.Y[[]for_inrange(n)]fori,(a,b)inenumerate(zip(x,y)):ifaisnotNoneandbisnotNone:self.X[i].append(a)self.Y[i].append(b)self.axes[0].cla()forx,y,fmtinzip(self.X,self.Y,self.fmts):self.axes[0].plot(x,y,fmt)self.config_axes()# display.display(self.fig)# display.clear_output(waitTrue)如果还是没有显示的话一般来说之前显示图没问题的话这一步是不用设置的并且如果全部都设置了还是无法显示就要考虑matplotlib和d2l的冲突了就在以下设置需要打开【不同的版本可能没有pillow那个无需理会全部勾选就ok】成功结果如下所示注本文参考d2l参考书中Zhang_Xuhui同学的评论并且贴主进行了一定的修改和优化