A minimal PyTorch training loop
Teach a small model y = 2x + 1
A CPU teaching example for observing gradient updates, not a benchmark of real-world generalisation.
1. Prepare the environment and tensors
Install the appropriate PyTorch version in an isolated Python environment using the official installation guide. This exercise needs no GPU, dataset download or torchvision package.
There are four input samples with one feature each: shape [4, 1]. Inspecting shape, dtype and device is a useful first step when debugging tensors; see the official tensor tutorial.
2. Complete example
1 | import torch |
3. Five actions in one update
Clear gradients, predict, calculate loss, backpropagate and update parameters. PyTorch accumulates gradients by default; omitting the clearing step changes this example’s optimisation. The API and steps are documented in the official optimisation tutorial.
nn.Linear(1, 1) represents a line with one weight and one bias. MSE measures squared prediction error. A reasonable check here is that the weight approaches 2, the bias approaches 1, and the prediction at x=3 approaches 7. Do not describe these expected teaching-example values as measured research results.
4. Separate inference from evaluation
model.eval() changes the behaviour of training-sensitive layers. torch.no_grad() disables gradient recording in a block. They do different things. This model has no Dropout, but both steps remain good habits; see the Module API.
A falling training loss does not establish performance on new data. For real tasks, split training, validation and test data before model selection. Consider grouped splits for related samples, such as measurements from the same patient or source, to avoid leakage.
Exercise
Reduce the learning rate and compare the parameter errors after the same number of steps. Then change the labels to 3*x-2: predict the appropriate result before running the code.
The deep-learning handbook has additional Chinese evaluation chapters. Return to the English collection for other English guides.