NVIDIA / NVIDIA/apex

Questions about numeric precision of FusedRMSNorm

Open
#1,652 5 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
9k
Forks
1.5k
Avg merge
2d 4h
Merged PRs (30d)
3

Description

Hello, I have tested the numeric precision of FusedRMSNorm and MixFusedRMSNorm in two different version respectively. Finally, I found that the gradient of model weights can not keep the same and there are a quiet large difference in weight gradient. Therefore, can you help me to solve it or give some suggestions?

The following is my test implementation:

import copy

import torch
from torch import nn

from apex.normalization.fused_layer_norm import MixedFusedRMSNorm, FusedRMSNorm


def manual_rms_norm(input, normalized_shape, weight, eps):
    # layer norm should always be calculated in float32
    
    dims = tuple(i for i in range(-1, -len(normalized_shape)-1, -1))
    
    variance = input.to(torch.float32).pow(2).mean(dims, keepdim=True)
    input = input * torch.rsqrt(variance + eps)
    if weight is None:
        return input
    # convert into half-precision if necessary
    if weight.dtype in [torch.float16, torch.bfloat16]:
        input = input.to(weight.dtype)
    return weight * input


class Manual_RMSNorm(nn.Module):
    
    def __init__(self, dim, eps, device=None):
        super().__init__()
        self.eps = eps
        self.dim = dim
        self.weight = self.weight = nn.Parameter(torch.ones(dim, device=device))
        
    def forward(self, x):
        return manual_rms_norm(x, self.weight.shape, self.weight, self.eps)


def main():
    
    dtype = torch.float16
    device = 'cuda'
    hidden_size = 10240
    input_size = 10240
    weight_type = torch.float16

    repeats = 1000
    
    x_pt = torch.randn(1, input_size, hidden_size, dtype=dtype, device=device).requires_grad_()
    x = x_pt.detach().clone().requires_grad_()
    x_mixed = x_pt.detach().clone().requires_grad_()
    x_func = x_pt.detach().clone().requires_grad_()

    model_pt = RMSNorm(dim=hidden_size, eps=1e-6).to(device=device, dtype=dtype)
    model = FusedRMSNorm(hidden_size, eps=1e-6).to(device=device, dtype=dtype)
    model_mixed = MixedFusedRMSNorm(hidden_size, eps=1e-6).to(device=device, dtype=dtype) 
    model_manual = Manual_RMSNorm(hidden_size, eps=1e-6).to(device=device, dtype=dtype)
    
    with torch.no_grad():
        model.weight.copy_(model_pt.weight)
        model_mixed.weight.copy_(model_pt.weight)
        model_manual.weight.copy_(model_pt.weight)
    
    output_pt = model_pt(x_pt)
    output = model(x)
    output_mixed = model_mixed(x_mixed)
    output_man = model_manual(x_func)
    
    loss = torch.rand_like(output) / 32
    
    output_pt.backward(loss)
    output.backward(loss)
    output_mixed.backward(loss)
    output_man.backward(loss)
    
    
    print("pytorch = ", model_pt.weight.grad)
    print("fused = ", model.weight.grad)
    print("mixed = ", model_mixed.weight.grad)
    print("man = ", model_manual.weight.grad)

The following is my test results:

pytorch =  tensor([-0.1270, -0.2615, -3.3340,  ..., -0.0049, -2.6680, -3.2207],
       device='cuda:0', dtype=torch.float16)
fused =  tensor([-0.1283, -0.2607, -3.3340,  ..., -0.0045, -2.6680, -3.2207],
       device='cuda:0', dtype=torch.float16)
mixed =  tensor([-0.1283, -0.2607, -3.3340,  ..., -0.0045, -2.6680, -3.2207],
       device='cuda:0', dtype=torch.float16)
man =  tensor([-0.1270, -0.2615, -3.3340,  ..., -0.0049, -2.6680, -3.2207],
       device='cuda:0', dtype=torch.float16)

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with the reproduction in the issue and the apex.normalization.fused_layer_norm module, comparing FusedRMSNorm and MixedFusedRMSNorm against the manual and PyTorch implementations. Check the forward and backward precision paths and reproduce the weight-gradient differences. Done means identifying whether the discrepancy is expected numerical variation or a bug, with tests or guidance covering the result.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning, performance
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.