Skip to content
Back to skills

Fp32 Reproducible Llm Inference

ASecurity

Diagnose and solve LLM reproducibility failures caused by floating-point precision across hardware configurations using LayerCast optimization for deterministic inference with minimal memory overhead.

  • 6 stars
  • 0 votes
  • 0 copies
  • 2 views
  • Added September 9, 2026
developmentpythongobackendperformance

Security analysis

A100/100

Scanned September 9, 2026

npx -y skills add ADu2021/skillXiv --skill fp32-reproducible-llm-inference --agent claude-code

Installs into .claude/skills of the current project.

Are you the author of Fp32 Reproducible Llm Inference?

Add the live security badge to your README. It updates with every re-scan.

Security grade badge for Fp32 Reproducible Llm Inference
[![Security: A — Skills Directory](https://www.skillsdirectory.com/api/skills/adu2021-fp32-reproducible-llm-inference/badge)](https://www.skillsdirectory.com/skills/adu2021-fp32-reproducible-llm-inference)

More formats (shields.io, HTML) on the badges page. Keep it an A: scan every change in CI with Pro.

Download with Pro
SKILL.md
---
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

```python
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,  # greedy decoding
                    num_beams=1
                )
            outputs.append(output)

        # Measure variance across runs
        divergence_positions = []
        for run in range(1, len(outputs)):
            first_diff = (outputs[0] != outputs[run]).nonzero(as_tuple=True)
            if len(first_diff[0]) > 0:
                divergence_positions.append(first_diff[1][0].item())

        results[str(config)] = {
            'divergence_index': np.mean(divergence_positions) if divergence_positions else -1,
            'output_length_var': np.var([len(o) for o in outputs])
        }

    return results
```

### Step 2: Implement LayerCast Optimization

```python
import torch
import torch.nn as nn

class LayerCastOptimizer:
    """
    Hybrid precision wrapper: store weights in BF16, compute in FP32.
    Preserves memory efficiency while ensuring numerical stability.
    """

    def __init__(self, model, compute_dtype=torch.float32):
        self.model = model
        self.compute_dtype = compute_dtype
        self.original_forward = {}

    def patch_matmul_layers(self):
        """Replace linear layer forward passes with LayerCast computation."""
        for name, module in self.model.named_modules():
            if isinstance(module, nn.Linear):
                self.original_forward[name] = module.forward
                module.forward = self._create_cast_forward(module)

    def _create_cast_forward(self, linear_module):
        """Create a forward function that upcasts weights to FP32 for computation."""
        original_forward = linear_module.forward

        def cast_forward(x):
            # Upcast weights from BF16 to FP32 just-in-time
            weight_fp32 = linear_module.weight.to(self.compute_dtype)
            bias_fp32 = linear_module.bias.to(self.compute_dtype) if linear_module.bias is not None else None
            x_fp32 = x.to(self.compute_dtype)

            # Perform FP32 computation
            output = torch.nn.functional.linear(x_fp32, weight_fp32, bias_fp32)

            # Cast back to BF16 if needed for downstream layers
            return output.to(linear_module.weight.dtype)

        return cast_forward
```

### Step 3: Configure Inference for Determinism

```python
def setup_deterministic_inference(model, batch_size=8, use_fp32=True):
    """
    Configure model and environment for reproducible inference.
    """
    # Set random seeds globally
    torch.manual_seed(42)
    torch.cuda.manual_seed_all(42)
    np.random.seed(42)

    # Enable deterministic algorithms (may impact performance)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False

    # Move model to consistent precision
    if use_fp32:
        model = model.to(torch.float32)
    else:
        # If using BF16, apply LayerCast
        optimizer = LayerCastOptimizer(model, compute_dtype=torch.float32)
        optimizer.patch_matmul_layers()

    return model

def generate_with_reproducibility(model, input_ids, max_length=512):
    """
    Generate outputs with maximum reproducibility guarantees.
    """
    with torch.no_grad():
        output = model.generate(
            input_ids,
            max_length=max_length,
            do_sample=False,  # Greedy decoding—no sampling variance
            num_beams=1,
            use_cache=False,  # Disable kv-cache if reproducibility critical
            temperature=1.0,
            top_p=1.0,
            top_k=None
        )
    return output
```

## Practical Guidance

**When to Use FP32 for Reproducibility**:
- Reasoning tasks generating long chains of thought (>2000 tokens)
- Mathematical problem-solving requiring consistency across runs
- Applications where inference output variance exceeds tolerance

**When LayerCast is Sufficient**:
- Production deployments prioritizing memory efficiency
- Batch inference where 34% memory savings outweighs minor non-determinism
- Models under 13B parameters where BF16 drift is modest (<2% accuracy variance)

**Trade-offs to Consider**:
- Full FP32: Perfect reproducibility but 50% memory increase
- LayerCast: 34% memory savings with FP32-level determinism (best balance)
- Pure BF16: 50% memory savings but up to 9% accuracy variance across configs

**Configuration Best Practices**:
- Standardize batch sizes within teams (one batch size per deployment)
- Pin GPU types when reproducibility is contractually required
- Report uncertainty quantiles (mean ± std) for downstream model outputs when using lower precision
- Log exact hardware configurations alongside inference results for auditing

## Reference

- Non-associativity of floating-point arithmetic: fundamental constraint of IEEE 754 standard
- Rounding error accumulation: proportional to operation count, controlled by mantissa precision
- Dynamic shapes and padding: creating uncontrollable kernel scheduling variations in operator fusion

Attribution

Is this your skill, or is something wrong with this listing? Request removal or report an issue. Author removals are honored within 72 hours.

Comments

Loading comments…