Download mindxtrain/train/dispatch.py from PYTHAI/mindXtrain: direct link, hf CLI and curl.
- Browser
- Download file 2.27 kB
-
https://huggingface.co/PYTHAI/mindXtrain/resolve/refs%2Fpr%2F1/mindxtrain/train/dispatch.py
- Command line
-
hf download hf://PYTHAI/mindXtrain@refs/pr/1/mindxtrain/train/dispatch.py
-
curl -L -o dispatch.py https://huggingface.co/PYTHAI/mindXtrain/resolve/refs%2Fpr%2F1/mindxtrain/train/dispatch.py
2.27 kB
| """Training backend dispatch. | |
| Lane selection happens at `cfg.train.backend`: | |
| - `axolotl` β GPU SFT/LoRA via Axolotl subprocess (default for MI300X recipes). | |
| - `unsloth` β GPU SFT via Unsloth. | |
| - `torchtune` β GPU SFT via torchtune. | |
| - `primus` β AMD's training stack. | |
| - `trl_cpu` β CPU SFT/LoRA via TRL in-process. Real checkpoints, slow. | |
| Use for: mindX self-training, smoke-testing a recipe before burning AMD | |
| credits, anywhere a MI300X droplet isn't available. | |
| - `trl_local` β same in-process TRL trainer, but device-aware: uses a local | |
| consumer GPU (CUDA or ROCm Radeon, bf16/fp16) when one is visible, else falls | |
| back to CPU. One recipe runs on a laptop or a gaming GPU unchanged. | |
| The CPU lane is paired with `hardware.gpus: 0` in the recipe. The dispatcher | |
| itself does not enforce that pairing β the recipe is the source of truth β | |
| but the schema's `Literal[0, 1, 8]` constrains the GPU count. | |
| """ | |
| from __future__ import annotations | |
| from pathlib import Path | |
| from mindxtrain.autotune.plan import AutotunePlan | |
| from mindxtrain.config.schema import XTrainConfig | |
| def dispatch_training( | |
| cfg: XTrainConfig, | |
| plan: AutotunePlan, | |
| out_dir: Path, | |
| ) -> Path: | |
| """Dispatch a training run to the configured backend. | |
| Returns the path to the produced checkpoint directory. | |
| """ | |
| backend = cfg.train.backend | |
| if backend == "axolotl": | |
| from mindxtrain.train.sft import run_axolotl | |
| return run_axolotl(cfg, plan, out_dir) | |
| if backend == "unsloth": | |
| from mindxtrain.train.backend_unsloth import run_unsloth | |
| return run_unsloth(cfg, plan, out_dir) | |
| if backend == "torchtune": | |
| from mindxtrain.train.backend_torchtune import run_torchtune | |
| return run_torchtune(cfg, plan, out_dir) | |
| if backend == "primus": | |
| from mindxtrain.train.backend_primus import run_primus | |
| return run_primus(cfg, plan, out_dir) | |
| if backend == "trl_cpu": | |
| from mindxtrain.train.backend_trl_cpu import run_trl_cpu | |
| return run_trl_cpu(cfg, plan, out_dir) | |
| if backend == "trl_local": | |
| from mindxtrain.train.backend_trl_cpu import run_trl_local | |
| return run_trl_local(cfg, plan, out_dir) | |
| msg = f"unknown backend {backend!r}" | |
| raise ValueError(msg) | |