Spaces:
Running on Zero
Running on Zero
Download grn/models/fused_op.py from hanjian/GRN: direct link, hf CLI and curl.
- Browser
- Download file 893 Bytes
-
https://huggingface.co/spaces/hanjian/GRN/resolve/main/grn/models/fused_op.py
- Command line
-
hf download hf://spaces/hanjian/GRN/grn/models/fused_op.py
-
curl -L -o fused_op.py https://huggingface.co/spaces/hanjian/GRN/resolve/main/grn/models/fused_op.py
893 Bytes
| import gc | |
| from copy import deepcopy | |
| from typing import Union | |
| import torch | |
| from torch import nn as nn | |
| from torch.nn import functional as F | |
| def fused_rms_norm(x: torch.Tensor, weight: nn.Parameter, eps: float): | |
| x = x.float() | |
| return (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(eps))) * weight | |
| def fused_ada_layer_norm(C: int, eps: float, x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor): | |
| x = x.float() | |
| x = F.layer_norm(input=x, normalized_shape=(C,), weight=None, bias=None, eps=eps) | |
| return x.mul(scale.add(1)).add_(shift) | |
| def fused_ada_rms_norm(C: int, eps: float, x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor): | |
| x = x.float() | |
| x = (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(eps))) | |
| return x.mul(scale.add(1)).add_(shift) | |