mirror of
https://github.com/vinta/awesome-python.git
synced 2026-10-02 08:23:10 +08:00
docs: add deep learning category intro
The Deep Learning page had no intro, leaving readers with no pick among PyTorch, Keras, JAX, Lightning, and the reinforcement learning libraries. Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,24 @@
|
||||
Three Python deep learning frameworks, three strengths: PyTorch for research and new architectures, Keras for a high-level API, JAX for compiled code on TPUs.
|
||||
|
||||
How to choose:
|
||||
|
||||
- Research and new model architectures: PyTorch
|
||||
- High-level model building on JAX or PyTorch: Keras
|
||||
- NumPy-style array code compiled for GPUs and TPUs: JAX
|
||||
- PyTorch training without the loop boilerplate: PyTorch Lightning
|
||||
- Reinforcement learning environments: Gymnasium
|
||||
- Reinforcement learning algorithms: Stable-Baselines3
|
||||
|
||||
PyTorch models are `nn.Module` subclasses, and autograd [builds the computational graph as your code runs](https://docs.pytorch.org/docs/stable/user_guide/pytorch_main_components.html). Install it with [the command from its selector](https://pytorch.org/get-started/locally/), which matches your OS and GPU. To keep a model, [save its `state_dict`](https://docs.pytorch.org/tutorials/beginner/saving_loading_models.html#save-load-state-dict-recommended) instead of pickling the whole module. Wrap your model in [`torch.compile`](https://docs.pytorch.org/tutorials/intermediate/torch_compile_tutorial.html) to speed it up with minimal code changes. To train on more than one GPU, use [DistributedDataParallel](https://docs.pytorch.org/docs/stable/notes/cuda.html#use-nn-parallel-distributeddataparallel-instead-of-multiprocessing-or-nn-dataparallel).
|
||||
|
||||
Keras is a [multi-framework API](https://keras.io/getting_started/about/): a Keras model can run as a PyTorch Module or as a JAX function. Install a backend next to it and [set `KERAS_BACKEND`](https://keras.io/getting_started/#configuring-your-backend) before you import Keras. Build simple models as a Sequential stack of layers and anything more complex with the functional API. Save the whole model to [a `.keras` file](https://keras.io/getting_started/faq/#what-are-my-options-for-saving-models) with `model.save()` instead of pickling it. The file [reloads with any backend](https://keras.io/keras_3/).
|
||||
|
||||
JAX does [accelerator-oriented array computation](https://docs.jax.dev/en/latest/) with a NumPy-style API and composable transformations: `jax.grad` for derivatives, `jax.jit` for compilation, and `jax.vmap` for batching. The transformations only work on [functionally pure functions](https://docs.jax.dev/en/latest/notebooks/Common_Gotchas_in_JAX.html#pure-functions), so pass all data in as arguments and return every result. JAX itself stays narrow: to train neural networks, use the [JAX AI Stack](https://docs.jaxstack.ai/en/latest/getting_started.html), with Flax NNX for models and Optax for optimizers. On NVIDIA GPUs, the JAX team strongly recommends [installing CUDA and cuDNN from pip wheels](https://docs.jax.dev/en/latest/installation.html#pip-installation-nvidia-gpu-cuda-installed-via-pip-easier).
|
||||
|
||||
PyTorch Lightning [organizes PyTorch code to remove boilerplate](https://lightning.ai/docs/pytorch/stable/home/introduction): you write the model logic in a LightningModule, and the Trainer handles devices, precision, and distributed training. Install it as [the `lightning` package](https://lightning.ai/docs/pytorch/stable/home/installation). Its [style guide](https://lightning.ai/docs/pytorch/stable/reference/starter/style_guide) recommends keeping each LightningModule self-contained, the model separate from the system that trains it, and data loading in a LightningDataModule. To keep your own training loop, use [Lightning Fabric](https://lightning.ai/docs/fabric/stable), which scales a plain PyTorch script after you change a few lines.
|
||||
|
||||
Gymnasium is [an API standard for reinforcement learning](https://gymnasium.farama.org/), with a collection of reference environments. It's the maintained fork of OpenAI's Gym, and many older tutorials still use Gym's old API, so follow its [migration guide](https://gymnasium.farama.org/introduction/migration_guide/) when you port one. [Register your own environment](https://gymnasium.farama.org/introduction/create_custom_env/#registering-and-making-the-environment) so `gymnasium.make()` creates it like a built-in one, and run [`check_env`](https://gymnasium.farama.org/introduction/create_custom_env/#check-environment-validity) on it to catch common issues.
|
||||
|
||||
Stable-Baselines3 is a set of [reliable implementations of reinforcement learning algorithms in PyTorch](https://stable-baselines3.readthedocs.io/en/master/), and it trains on any environment that [follows the Gymnasium interface](https://stable-baselines3.readthedocs.io/en/master/guide/custom_env.html). It [assumes you know some reinforcement learning](https://github.com/DLR-RM/stable-baselines3). Its tips page recommends [starting from the RL Zoo's tuned hyperparameters](https://stable-baselines3.readthedocs.io/en/master/guide/rl_tips.html#general-advice-when-using-reinforcement-learning) and normalizing the agent's input. Evaluate the agent on [a separate test environment](https://stable-baselines3.readthedocs.io/en/master/guide/rl_tips.html#how-to-evaluate-an-rl-algorithm), since training adds exploration noise. [Pick an algorithm](https://stable-baselines3.readthedocs.io/en/master/guide/rl_tips.html#which-algorithm-should-i-use) by your action space first: DQN handles only discrete actions, and SAC only continuous ones.
|
||||
|
||||
PyTorch's security policy says [running untrusted models is equivalent to running untrusted code](https://github.com/pytorch/pytorch/blob/main/SECURITY.md), so run untrusted ones in a sandbox. Load checkpoints with [`weights_only=True`](https://docs.pytorch.org/docs/stable/notes/serialization.html#weights-only-security) in `torch.load`, and leave [`safe_mode`](https://keras.io/api/models/model_saving_apis/model_saving_and_loading/) on when Keras loads a model.
|
||||
Reference in New Issue
Block a user