MegaDLM Numerical Precision Formats: FP32, FP16, BF16, FP8, and Quantized Inference

MegaDLM supports FP32, FP16, BF16, FP8, and INT4/INT8 quantized formats for training and inference, controlled via command-line arguments in megatron/training/arguments.py and megatron/inference/arguments.py.

Built on Megatron-LM and NVIDIA Transformer Engine, the MegaDLM framework (from jinjieni/megadlms) exposes granular numerical precision formats to optimize memory usage and computational throughput across large-scale model training and deployment pipelines.

Training Precision Formats

MegaDLM defines its training precision formats through flags in megatron/training/arguments.py, supporting both standard floating-point types and the latest FP8 specification.

FP32, FP16, and BF16

By default, MegaDLM operates in FP32 (torch.float32) when no precision flag is specified. To enable mixed-precision training, use either --fp16 for torch.float16 (defined at line 1659) or --bf16 for torch.bfloat16 (line 1661). These flags trigger automatic type casting in the forward and backward passes while maintaining FP32 master weights for numerical stability.

In megatron/core/transformer/transformer_config.py, the framework enforces mutual exclusivity between these flags with assertions like assert not (fp16 and bf16), preventing conflicting precision configurations.

FP8 Support via Transformer Engine

For hardware-accelerated FP8 training, MegaDLM integrates with NVIDIA Transformer Engine. The --fp8-format flag (line 833) accepts e4m3 or e5m2 formats, storing tensors as torch.uint8 FP8 representations. Activating --fp8-param-gather keeps parameters in FP8 precision throughout the all-gather communication step, reducing bandwidth overhead.

Additional FP8-related options include --fp32-residual-connection and --attention-softmax-in-fp32, which cast specific operations to FP32 to preserve numerical stability where needed.

Per-Tensor Dtype Overrides

Beyond global precision settings, MegaDLM allows fine-grained control over individual tensor types through arguments defined around lines 2213–2217 in megatron/training/arguments.py:

  • --main-grads-dtype: Select fp32 or bf16 for gradient storage
  • --main-params-dtype: Choose fp32 or fp16 for parameter storage
  • --exp-avg-dtype / --exp-avg-sq-dtype: Configure Adam optimizer states as fp32, fp16, or fp8

These overrides propagate through megatron/training/initialize.py, which constructs the model provider and casts tensors accordingly using .to(dtype) operations.

Inference Quantization Formats

For deployment scenarios, MegaDLM supports aggressive quantization through megatron/inference/arguments.py (line 22). The --quantization flag accepts:

  • int8 and int8_sq: 8-bit integer quantization
  • fp8: 8-bit floating point inference
  • int4_awq, w4a8_awq, int4: 4-bit quantization schemes including Activation-aware Weight Quantization (AWQ)

These modes engage specialized kernels that reduce memory footprint and accelerate token generation in megatron/inference/text_generation/server.py.

Configuration Examples

BF16 Mixed Precision Training

python pretrain_difflm.py \
    --config-path configs/gpt2_1b.yaml \
    --bf16 \
    --main-grads-dtype bf16 \
    --main-params-dtype fp16 \
    --exp-avg-dtype fp16 \
    --exp-avg-sq-dtype fp16

This configuration activates bfloat16 for forward/backward computations while storing parameters and optimizer states in FP16, balancing memory efficiency and numerical range.

FP8 Training with Transformer Engine

python pretrain_difflm.py \
    --config-path configs/gpt2_1b.yaml \
    --fp8-format e4m3 \
    --fp8-param-gather \
    --main-grads-dtype fp8 \
    --exp-avg-dtype fp8 \
    --exp-avg-sq-dtype fp8

The --fp8-format e4m3 argument selects the 4-bit exponent/3-bit mantissa layout, while --fp8-param-gather maintains FP8 precision during distributed parameter synchronization.

INT8 Quantized Inference

python megatron/inference/text_generation/server.py \
    --model-path checkpoints/gpt2_1b \
    --quantization int8 \
    --batch-size 8 \
    --max-new-tokens 32

Setting --quantization int8 triggers the INT8 kernel path, reducing model memory by approximately 50% compared to FP16 baselines.

Mixed Precision Inference (FP16)

python megatron/inference/text_generation/server.py \
    --model-path checkpoints/gpt2_1b \
    --dtype fp16 \
    --batch-size 4

The --dtype fp16 flag executes the forward pass in half precision, mapped internally to the same parser handling --fp16 in training contexts.

Summary

  • MegaDLM supports five core numerical precision formats: FP32 (default), FP16, BF16, FP8, and INT4/INT8 quantization for training and inference.
  • Training precision is controlled via --fp16, --bf16, and --fp8-format flags in megatron/training/arguments.py, with per-tensor overrides for gradients, parameters, and optimizer states.
  • Inference quantization offers INT8, INT4 AWQ, and FP8 modes through megatron/inference/arguments.py, enabling reduced memory footprints for deployment.
  • Transformer Engine integration provides FP8-aware kernels and automatic mixed-precision handling for residual connections and attention softmax operations.

Frequently Asked Questions

Does MegaDLM support FP8 training?

Yes. MegaDLM supports FP8 training through integration with NVIDIA Transformer Engine. Use the --fp8-format flag (defined at line 833 in megatron/training/arguments.py) with values e4m3 or e5m2, and enable --fp8-param-gather to maintain FP8 precision during distributed parameter collection.

What is the default numerical precision format in MegaDLM?

The default format is FP32 (torch.float32), used when neither --fp16, --bf16, nor --fp8-format is specified. This provides maximum numerical stability at the cost of increased memory usage and computational requirements compared to lower-precision formats.

How do I enable INT8 quantization for inference?

Pass --quantization int8 to the inference server script located in megatron/inference/text_generation/server.py. This option is defined in megatron/inference/arguments.py (line 22) and activates INT8 kernels that reduce model memory consumption by approximately half compared to FP16 baselines.

Can I mix different precision formats for gradients and parameters?

Yes. MegaDLM provides independent dtype controls via --main-grads-dtype (fp32 or bf16), --main-params-dtype (fp32 or fp16), and optimizer state flags like --exp-avg-dtype (supporting fp32, fp16, or fp8). These per-tensor overrides are processed in megatron/training/initialize.py and allow fine-tuning the memory-accuracy tradeoff for specific components.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →