Spaces:
				
			
			
	
			
			
		Runtime error
		
	
	
	
			
			
	
	
	
	
		
		
		Runtime error
		
	File size: 1,633 Bytes
			
			| 8b54513 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 | import warnings
import torch
import torch.nn as nn
try:
    from apex.normalization import FusedRMSNorm as RMSNorm
except ImportError:
    warnings.warn("Cannot import apex RMSNorm, switch to vanilla implementation")
    class RMSNorm(torch.nn.Module):
        def __init__(self, dim: int, eps: float = 1e-6):
            """
            Initialize the RMSNorm normalization layer.
            Args:
                dim (int): The dimension of the input tensor.
                eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
            Attributes:
                eps (float): A small value added to the denominator for numerical stability.
                weight (nn.Parameter): Learnable scaling parameter.
            """
            super().__init__()
            self.eps = eps
            self.weight = nn.Parameter(torch.ones(dim))
        def _norm(self, x):
            """
            Apply the RMSNorm normalization to the input tensor.
            Args:
                x (torch.Tensor): The input tensor.
            Returns:
                torch.Tensor: The normalized tensor.
            """
            return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
        def forward(self, x):
            """
            Forward pass through the RMSNorm layer.
            Args:
                x (torch.Tensor): The input tensor.
            Returns:
                torch.Tensor: The output tensor after applying RMSNorm.
            """
            output = self._norm(x.float()).type_as(x)
            return output * self.weight
 | 
