| name | fp32-reproducible-llm-inference |
| title | Give Me FP32 or Give Me Death? Challenges and Solutions for Reproducible Reasoning |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.09501 |
| keywords | ["LLM reproducibility","floating-point precision","numerical stability","BF16","FP32","inference determinism"] |
| description | Diagnose and solve LLM reproducibility failures caused by floating-point precision across hardware configurations using LayerCast optimization for deterministic inference with minimal memory overhead. |
Give Me FP32 or Give Me Death?
Core Concept
LLM inference reproducibility fails dramatically across hardware configurations due to non-associative floating-point arithmetic. Even with fixed random seeds and greedy decoding, changing batch size, GPU count, or GPU type produces divergent outputs with up to 9% accuracy variance and 9,000-token length differences in reasoning models. The root cause: limited precision in BF16 (7 mantissa bits) creates rounding error accumulation that varies by kernel execution order.
Architecture Overview
- Precision Hierarchy: FP32 (23 bits) achieves near-perfect reproducibility; FP16 (10 bits) shows moderate variability; BF16 (7 bits) fails dramatically
- Non-Associativity Problem: Floating-point addition violates associativity—kernel scheduling and GPU memory layout change computation order, producing different accumulated rounding errors
- LayerCast Solution: Hybrid approach storing weights in memory-efficient BF16 while performing all computations in FP32, achieving deterministic results with 34% memory savings
- Configuration Impact: Divergence occurs predictably: different batch sizes→different padding patterns→different GPU kernel launches→different floating-point operation ordering
Implementation
Step 1: Diagnose Precision-Related Nondeterminism
import torch
import numpy as np
def measure_reproducibility_drift(model, input_ids, configs):
"""
Test model outputs across different hardware configurations.
Configs: list of dicts with 'batch_size', 'num_gpus', 'gpu_type'
"""
results = {}
for config in configs:
outputs = []
for run in range(3):
torch.manual_seed(42)
with torch.no_grad():
output = model.generate(
input_ids,
max_length=512,
do_sample=False,
num_beams=1
)
outputs.append(output)
divergence_positions = []
run (, (outputs)):
first_diff = (outputs[] != outputs[run]).nonzero(as_tuple=)
(first_diff[]) > :
divergence_positions.append(first_diff[][].item())
results[(config)] = {
: np.mean(divergence_positions) divergence_positions -,
: np.var([(o) o outputs])
}
results