d2l-ai/d2l-zh · error · AssertionError
test_acc <= 1 and test_acc > 0.7
Error message
test_acc <= 1 and test_acc > 0.7
What it means
Python AssertionError from d2l.torch.train_ch3's final sanity check: test accuracy must be in (0.7, 1] after training on Fashion-MNIST. It verifies the trained model generalizes as the book claims (~0.83 test acc); failure means underfitting, divergence, or an evaluation-path bug rather than a library defect.
Source
Thrown at d2l/torch.py:341
self.axes[0].plot(x, y, fmt)
self.config_axes()
display.display(self.fig)
display.clear_output(wait=True)
def train_ch3(net, train_iter, test_iter, loss, num_epochs, updater):
"""训练模型(定义见第3章)
Defined in :numref:`sec_softmax_scratch`"""
animator = Animator(xlabel='epoch', xlim=[1, num_epochs], ylim=[0.3, 0.9],
legend=['train loss', 'train acc', 'test acc'])
for epoch in range(num_epochs):
train_metrics = train_epoch_ch3(net, train_iter, loss, updater)
test_acc = evaluate_accuracy(net, test_iter)
animator.add(epoch + 1, train_metrics + (test_acc,))
train_loss, train_acc = train_metrics
assert train_loss < 0.5, train_loss
assert train_acc <= 1 and train_acc > 0.7, train_acc
assert test_acc <= 1 and test_acc > 0.7, test_acc
def predict_ch3(net, test_iter, n=6):
"""预测标签(定义见第3章)
Defined in :numref:`sec_softmax_scratch`"""
for X, y in test_iter:
break
trues = d2l.get_fashion_mnist_labels(y)
preds = d2l.get_fashion_mnist_labels(d2l.argmax(net(X), axis=1))
titles = [true +'\n' + pred for true, pred in zip(trues, preds)]
d2l.show_images(
d2l.reshape(X[0:n], (n, 28, 28)), 1, n, titles=titles[0:n])
def evaluate_loss(net, data_iter, loss):
"""评估给定数据集上模型的损失
Defined in :numref:`sec_model_selection`"""
metric = d2l.Accumulator(2) # 损失的总和,样本数量View on GitHub (pinned to e6b18ccea7)
Solutions
- Rerun with the book's configuration: fresh net, num_epochs=10, lr=0.1, batch_size=256
- Build train_iter and test_iter with identical transforms (only shuffle=True vs False differs)
- Confirm net.eval() semantics if your custom model has dropout/batchnorm (d2l's evaluate_accuracy handles the standard softmax net)
- If train acc is high but test acc ~0.1, check label order/argmax axis and that the same net object is passed to evaluation
Example fix
# before train_iter = load_data_fashion_mnist(batch_size, resize=None)[0] # normalized _, test_iter = load_data_fashion_mnist(batch_size) # different path train_ch3(net, train_iter, test_iter, loss, 2, updater) # test_acc 0.5 -> AssertionError # after train_iter, test_iter = load_data_fashion_mnist(256) net = nn.Sequential(nn.Flatten(), nn.Linear(784, 10)) trainer = torch.optim.SGD(net.parameters(), lr=0.1) train_ch3(net, train_iter, test_iter, loss, 10, trainer.step) # test_acc ~0.83
Defensive patterns
Strategy: validation
Validate before calling
train_iter, test_iter = d2l.load_data_fashion_mnist(256) # same transforms net = build_fresh_net() # rebuild, never reuse stale weights assert next(iter(test_iter)) is not None
Type guard
def eval_pipeline_ready(net, test_iter) -> bool:
X, y = next(iter(test_iter))
return net(X).shape[1] == 10 and X.shape[0] == y.shape[0] Try / catch
try:
train_ch3(net, train_iter, test_iter, loss, num_epochs, updater)
except AssertionError as e:
print(f'test_acc={e.args[0]} outside (0.7, 1]; check epochs, transforms, net freshness')
raise Prevention
- Create both DataLoaders from load_data_fashion_mnist so transforms match
- Instantiate a fresh net per experiment
- Keep num_epochs=10; do not trim for speed in assert-bearing code
- Verify transforms (ToTensor/normalize) applied to both splits
When it happens
Trigger: evaluate_accuracy(net, test_iter) returning <= 0.7: model undertrained (num_epochs=1-3), diverged lr, net reused from a previous overfit state, or test_iter normalized differently from train_iter (e.g. ToTensor only on the test transform path by mistake); also evaluating with net still in train() mode so dropout/batchnorm skew results if a custom net uses them.
Common situations: CPU-only environments where users trim num_epochs; reusing a net variable across notebook cells without re-instantiation; transforms applied inconsistently between the two DataLoaders; running on torch>=2.6 where default DataLoader/factory settings changed and iterators need explicit handling.
Related errors
- test_acc <= 1 and test_acc > 0.7
- train_loss < 0.5
- train_loss < 0.5
- train_acc <= 1 and train_acc > 0.7
- f"{name} 不存在于 {DATA_HUB}"
AI-assisted analysis of d2l-ai/d2l-zh@e6b18ccea7 (2026-08-14).
Data as JSON: /api/errors/fe78a917d06d2588.
Report an issue: GitHub.