ML Tips and Tricks
ML Tips and Tricks
Many students have low productivity or struggle to get neural networks to work properly. This is often not because of a lack of understanding of neural networks from a conceptual standpoint, the mathematics behind them or lack of coding skills but rather due to the peculiar challenges neural network training pose. The problem is exacerbated by:
- Neural networks are hard to debug: It is not clear where something has gone wrong, they do not break in an intelligible way. In traditional code one can write unit tests for all important functions and check human-understandable variable values. In neural networks, this is not possible.
- Entanglement: you cannot separate components well to study them in isolation.
- Many bugs are subtle and affect only the final performance but are not seen in isolation.
- Debugging must often be done on a statistical, aggregate level: statistics on gradients, parameters, activations and training dynamics.
- What makes neural networks work well is to some extent a range of tricks and techniques. When not present, an otherwise identical and reasonable neural network trains poorly.
Getting Up and Running
- Get the smallest and simplest reasonable baseline running. The start of most projects will be boring!
- Remove any fanciness: No data augmentation, learning rate decay, no superfluous loss terms that are not strictly necessary, remove architectural complexity etc.
- Beware however: There are cases where a baseline only works due to all the employed fancy tricks! You must judge whether simplification is possible.
- Overfit on a single training data point. You must be able to drive loss to almost zero.
- If this has worked, overfit on several datapoints.
- Verify decreasing training loss.
- Verify reasonable training loss at beginning. When you do classification, what is the loss and accuracy of predicting classes uniformly?
- Verify prediction dynamics on a selected batch of data over time.
- Verify shapes. Are you applying linear layers, convolutions or attention on the right dimension? Check whether you use
batch_firstin PyTorch modules. - Check whether you are in train mode during the training loop and in eval mode otherwise. Do not forget
torch.no_gradwhen not training. - First write functions/classes specific to your exact use, only later generalize. Check if general version works identically as the specialized version.
- Check for information leakage: Sometimes you train extremely fast but do not generalize because training data from the dataloader leaks inadvertently the label to compute, e.g. via positional dependencies.
- Add components one at a time and check if everything works. When adding multiple components, amount of bugs grows exponentially and debugging becomes harder.
Architecture
- Whenever possible use residual connections. Check if there is an uninterrupted residual path from the inputs.
- Always use normalization, be it batch, layer, group or instance norm, whatever is applicable and gives best results.
- Normalization before non-linearity.
- GeLU/SwiGLU very often gives a moderate boost in comparison to plain ReLU. You can use it by default.
- Initialize correctly: The standard PyTorch initialization might have too large standard deviations. Proper initialization will help initially with faster convergence and allow you to reach better final loss. Look into Kaiming/Orthogonal/Glorot initialization. You can judge if it worked by looking at loss improvements early on in the training already.
- If you use a wider architecture, you might see significantly greater change per update. Small updates to wide matrices lead to bigger changes than larger updates in thin matrices.
- If you use multiple components in your architecture check parameter count: Does it make sense, is one component much larger than others, should learning rates/weight decay/… be the same for all components or not?
Logging
Logging is the principal way to determine what is going well and where problems might be. If you do not log sufficiently, our discussions will be much less fruitful. Always log the following:
- Training and validation loss.
- Training loss on each training step.
- Validation loss most commonly after each epoch on full validation set.
- Training/validation accuracy/F1-score/precision/recall or whatever quality metric of interest. Log on the same intervals as for the training and validation loss.
- Step size.
- Activations: For each layer or tensor log minimum/maximum/average and median.
- Gradients: For each layer or tensor log minimum/maximum/average/median and norm.
- If you have multiple classes plot a confusion matrix of prediction results.
Losses
- In case your loss has multiple terms:
- A common mistake is not to log the contributions of each individual term to the overall loss. Log the weighted subterms and also log the percentage of their contribution to the overall loss.
- When appropriate, if you see that some term does not contribute much, increase its weight. Similarly decrease the weights of terms that account for too much of the loss. A common mistake is to add loss terms that do not do anything because of too little weight and then wondering why nothing has changed.
- In case your dataset is imbalanced, try to reweight the classes. Upweight rare classes and downweight common classes. One way to do so is inverse frequency. Also try out focal loss.
- If you need to regularize via weight decay, also heed the advice given for multiple loss terms: Regularization term should be in similar range as loss terms.
- During the training, observe whether all loss terms go down. If some loss term goes down less than others, one might want to upweight it, analogously for terms that go down faster.
Generalization
If you do not generalize:
- Try larger weight decay. In some non-standard setting (GNNs, scientific data, …) even an unreasonably large value of [0.1, 10] might be good!
- Try DropOut. Be careful however, this only works if there is redundancy in the data. Sometimes dropout conflicts with other parts of the networks, e.g. it might not play nicely with batch normalization.
- Use early stopping. Take the checkpoint of the model with highest validation accuracy or lowest validation loss.
- Stochastic depth might be helpful.
Optimization
- Use AdamW by default and take a standard learning rate in [1e-3, 1e-4].
- Sometimes SGD with momentum can be better, but it is harder to tune.
- Use warmup and then decay your learning rate, e.g. via a cosine schedule.
- Do not decay your learning rate too fast. You must give the model enough time to learn. It might be wise to turn off learning rate decay at the beginning to see when model converges with the initial large learning rate.
- When fine-tuning use a smaller learning rate, e.g. in [1e-6, 1e-5] and no learning rate decay.
- When you use a pre-trained backbone and train a task head from scratch, use different learning rates for each component: the pre-training learning rate for the backbone and the fine-tuning learning rate for the task head.
- Layer-wise learning rate decay might give additional boost.
- Smaller batch size will give more stochastic updates. It is worth to try small batch sizes for their regularizing effect.
- Use gradient clipping, a sane default clip value is 1.0.
Training Data
- More is better, use as much as you can for final training.
- If you do not have enough training data, generate more:
- Synthetic training data.
- Augmented training data.
- Training in unsupervised or self-supervised way.
- Use pretrained models. They might even generalize to new domains!
- Supervised data is better than unsupervised one. Do not by default buy into the representation learning hype.
Hyperparameters
Do not tune hyperparameters in the beginning. Get some that work decently and work on all other aspects. Only at the end, when everything else works, tune hyperparams to squeeze out the most juice.
- Take hyperparameters from existing literature initially.
- Batch size: Try out the largest batch size that fits on GPU.
- Grid search almost never works. There are typically just too many hyperparams.
- Random search is better and might even work well.
- Try tuning one hyperparameter after the other in order of most important to least important. Hyperparameters might not be strongly interdependent.
Miscellaneous
- Fix all your random seeds to make experiments reproducible.
- Investigating single datapoints and checking exactly what happens is often more insightful than looking at aggregates.
- Use
torch.compileto make model run faster. - Use fused option in optimizer for faster updates.
- Check if training with 16 bits is giving you the same performance. If so, use it.
- Hold hyperparameters in one place, e.g. in a dataclass.
- Automate everything. Running an experiment and evaluation results must be doable with exactly one command.
- For attention check whether you mix tokens and not other dimensions.
Tools
- Use VSCode (or nvim or similar) and develop and debug directly on a GPU machine.
- Use the Python debugger.
- If you do not have a GPU on your development machine, use the SSH extension and develop and debug remotely.
- TensorBoard or Weights & Biases for logging.
- Dataclasses for holding (hyper-)parameters.
- Hydra for storing parameters in configuration files and passing them as command line arguments.
- tmux for running jobs while being disconnected.
General Considerations
- Prefer a simple approach vs. a complicated one.
- Example: Supervised > unsupervised > reinforcement learning. Use complexity only where it is needed.
- When you can encode information through data, this might be preferable to encoding it via the architecture.
- A simpler approach is typically more scalable and will stand the bitter lesson.
- Sometimes, adding complexity is unavoidable. If so, you must go for it.
- If initial hypotheses do not hold up, do not try to make the results fit. Go where results lead you. Revise your initial assumptions, do not cling to them if they turn out not to be helpful. If an approach does not have potential, abandon it.
- Your goal is to iterate as fast as possible. Automate everything. Become productive.
- There are always exceptions to each rule above. Of course do not apply these rules when inapplicable.
- For example, a loss balancing term in mixture of experts should not go down but stay the same during training. Its contribution on the overall loss should be small as possible, while ensuring that load becomes unbalanced.
- Sometimes you cannot start simple with a barebone baseline. Sometimes complicated architectural elements/losses/data augmentations/post-processing are strictly necessary and everything breaks down without them. In this case you must start with an involved baseline.
- Read a lot of papers, work towards scientific maturity, develop your own ideas and try to see if they survive criticism.
Useful Resources
Some material has been taken from these sources, also consult them:
- A Recipe for Training Neural Networks by Andrej Karpathy
- Deep Learning Tuning Playbook by Google Research