How NanoChat Manages Computation Precision (dtype) on Different Hardware
NanoChat automatically selects the optimal floating-point precision for matrix operations by detecting CUDA compute capabilities at runtime, defaulting to bfloat16 on Ampere-or-newer GPUs while allowing manual override via the NANOCHAT_DTYPE environment variable.
The karpathy/nanochat repository handles heavy numeric computation—matrix multiplications, activations, and attention kernels—through a centralized dtype management strategy that adapts to underlying hardware capabilities. Understanding how nanochat manages computation precision (dtype) on different hardware ensures you can optimize inference speed on modern GPUs or maintain stability on older hardware and CPUs.
Automatic dtype Detection Strategy
At import time, nanochat determines the global computation dtype through a cascading decision tree implemented in nanochat/common.py::_detect_compute_dtype. The function evaluates three factors in order: user preference, GPU architecture, and safe fallbacks.
Environment Variable Override
The system checks for the NANOCHAT_DTYPE environment variable first. If set to bfloat16, float16, or float32, nanochat uses this value verbatim without further hardware inspection.
export NANOCHAT_DTYPE=bfloat16
python -m scripts.chat_cli.py
This override is processed in nanochat/common.py lines 18–21, ensuring all downstream components see the same dtype immediately.
CUDA Compute Capability Check
When no override exists and CUDA is available, the code queries the GPU's compute capability. Devices with SM ≥ 8.0 (Ampere architecture or newer) automatically select bfloat16 for optimal tensor core utilization. Older GPUs fall back to float32—the code explicitly avoids float16 on these devices because nanochat does not implement a GradScaler for loss scaling (lines 22–29 in nanochat/common.py).
Non-CUDA Fallbacks
For CPU or MPS (Metal Performance Shaders) backends, the default is strictly float32 (line 30 in nanochat/common.py). This conservative choice prevents numerical instability on hardware without native bfloat16 support.
Once detected, nanochat stores the dtype in the module-level constant COMPUTE_DTYPE alongside a human-readable explanation in COMPUTE_DTYPE_REASON (lines 31–32).
dtype Propagation in the Inference Stack
After detection, the chosen dtype flows through the inference pipeline, affecting memory allocation, kernel selection, and checkpoint loading.
Engine KVCache Allocation
When the Engine class initializes generation, it creates a KVCache with precision matching the device type. In nanochat/engine.py lines 173–181, the code explicitly selects:
dtype = torch.bfloat16 if device.type == "cuda" else torch.float32
This mirrors the repository-wide assumption that CUDA devices handle bfloat16 efficiently while CPU inference requires full precision.
Flash Attention Kernel Selection
The optimized attention kernels in nanochat/flash_attention.py reference COMPUTE_DTYPE to select implementation paths. At line 58, the code checks this global constant to determine whether to invoke bfloat16-specific flash-attention kernels or standard precision alternatives, ensuring hardware-appropriate execution without runtime overhead.
Checkpoint Loading Safety
When loading models on non-CUDA devices, nanochat/checkpoint_manager.py converts any bfloat16 tensors to float32 automatically (lines 88–90). This conversion prevents runtime errors on CPUs that lack native bfloat16 support, maintaining inference stability across hardware platforms.
Configuring and Inspecting dtype at Runtime
You can verify or manipulate the computation precision using the public API.
Inspect Detected dtype
from nanochat.common import COMPUTE_DTYPE, COMPUTE_DTYPE_REASON
print(f"Using compute dtype: {COMPUTE_DTYPE} ({COMPUTE_DTYPE_REASON})")
On an Ampere GPU, this outputs:
Using compute dtype: torch.bfloat16 (auto-detected: CUDA SM 80 (bf16 supported))
Force CPU-Safe Loading
from nanochat.checkpoint_manager import load_model
# On CPU, automatically converts bfloat16 weights to float32
model, tokenizer, _ = load_model("base", device, phase="eval")
Create Engine with Hardware-Aware Caching
from nanochat.engine import Engine
# Engine automatically selects dtype based on device.type
engine = Engine(model, tokenizer)
# KVCache internally uses bfloat16 on CUDA, float32 on CPU
Summary
- User override: Set
NANOCHAT_DTYPEenvironment variable to force a specific precision regardless of hardware. - Ampere GPUs: Automatically use bfloat16 via detection in
nanochat/common.py::_detect_compute_dtype(SM ≥ 8.0). - Legacy GPUs and CPUs: Default to float32 for compatibility and numerical stability.
- Global constants:
COMPUTE_DTYPEandCOMPUTE_DTYPE_REASONexport the decision for use in attention kernels and caching. - Checkpoint safety:
nanochat/checkpoint_manager.pyconverts bfloat16 tensors to float32 when loading on non-CUDA devices.
Frequently Asked Questions
Why does nanochat avoid float16 on older GPUs?
The codebase intentionally falls back to float32 rather than float16 on pre-Ampere GPUs because mixed-precision training requires a GradScaler to handle gradient underflow, which is not implemented in nanochat (as noted in nanochat/common.py lines 22–29). Using float32 avoids silent numerical errors during inference on older hardware.
How can I force float32 on a modern GPU?
Set the environment variable before importing nanochat:
export NANOCHAT_DTYPE=float32
python your_script.py
This bypasses the automatic Ampere detection and forces full precision throughout the stack.
Does the dtype affect checkpoint compatibility?
No. Checkpoints are saved in their native precision, but nanochat/checkpoint_manager.py handles automatic conversion during loading. If you load a bfloat16-trained model on CPU, the tensors convert to float32 at line 88–90, allowing seamless cross-hardware deployment without manual conversion.
Where is the dtype detection logic located?
The core detection logic resides in nanochat/common.py within the _detect_compute_dtype function (lines 18–32). This module exports COMPUTE_DTYPE and COMPUTE_DTYPE_REASON as global constants that nanochat/engine.py, nanochat/flash_attention.py, and nanochat/checkpoint_manager.py all reference to maintain consistent precision across the inference pipeline.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →