File size: 6,649 Bytes
48b5986
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
import importlib
import numbers
import os
import sys
import time

import torch
from torch.nn.parameter import Parameter

sys.path.append(os.path.dirname(__file__))

try:
    fastfold_layer_norm_cuda = importlib.import_module("fastfold_layer_norm_cuda")
except ImportError:
    from model.protenix.layer_norm.torch_ext_compile import compile

    current_dir = os.path.dirname(__file__)
    fastfold_layer_norm_cuda = compile(
        name="fastfold_layer_norm_cuda",
        sources=[
            os.path.join(f"{current_dir}/kernel", file)
            for file in ["layer_norm_cuda.cpp", "layer_norm_cuda_kernel.cu"]
        ],
        extra_include_paths=[f"{current_dir}/kernel"],
        build_directory=current_dir,
    )


class FusedLayerNormAffineFunction(torch.autograd.Function):

    @staticmethod
    def forward(ctx, input, weight, bias, normalized_shape, eps):
        d = input.dtype

        ctx.normalized_shape = normalized_shape
        ctx.eps = eps
        input_ = input.contiguous()

        if weight is None:
            if bias is None:
                output, mean, invvar = fastfold_layer_norm_cuda.forward_none_affine(
                    input_, ctx.normalized_shape, ctx.eps
                )
            else:
                output, mean, invvar = (
                    fastfold_layer_norm_cuda.forward_with_bias_affine(
                        input_, ctx.normalized_shape, bias.to(d), ctx.eps
                    )
                )
        else:
            if bias is None:
                output, mean, invvar = (
                    fastfold_layer_norm_cuda.forward_with_weight_affine(
                        input_, ctx.normalized_shape, weight.to(d), ctx.eps
                    )
                )
            else:
                output, mean, invvar = (
                    fastfold_layer_norm_cuda.forward_with_both_affine(
                        input_, ctx.normalized_shape, weight.to(d), bias.to(d), ctx.eps
                    )
                )
        ctx.save_for_backward(input_, weight, bias, mean, invvar)
        return output

    @staticmethod
    def backward(ctx, grad_output):
        d = grad_output.dtype
        input_, weight_, bias_, mean, invvar = ctx.saved_tensors
        grad_input = grad_weight = grad_bias = None

        if weight_ is None:
            if bias_ is None:
                grad_input, grad_weight, grad_bias = (
                    fastfold_layer_norm_cuda.backward_none_affine(
                        grad_output.contiguous(),
                        mean,
                        invvar,
                        input_,
                        ctx.normalized_shape,
                        ctx.eps,
                    )
                )
            else:
                grad_input, grad_weight, grad_bias = (
                    fastfold_layer_norm_cuda.backward_with_bias_affine(
                        grad_output.contiguous(),
                        mean,
                        invvar,
                        input_,
                        ctx.normalized_shape,
                        bias_.to(dtype=d),
                        ctx.eps,
                    )
                )
        else:
            if bias_ is None:
                grad_input, grad_weight, grad_bias = (
                    fastfold_layer_norm_cuda.backward_with_weight_affine(
                        grad_output.contiguous(),
                        mean,
                        invvar,
                        input_,
                        ctx.normalized_shape,
                        weight_.to(dtype=d),
                        ctx.eps,
                    )
                )
            else:
                grad_input, grad_weight, grad_bias = (
                    fastfold_layer_norm_cuda.backward_with_both_affine(
                        grad_output.contiguous(),
                        mean,
                        invvar,
                        input_,
                        ctx.normalized_shape,
                        weight_.to(dtype=d),
                        bias_.to(dtype=d),
                        ctx.eps,
                    )
                )
        return (
            grad_input,
            None if weight_ is None else grad_weight,
            None if bias_ is None else grad_bias,
            None,
            None,
        )


class FusedLayerNorm(torch.nn.Module):

    def __init__(
        self,
        normalized_shape,
        create_scale=True,
        create_offset=True,
        eps=1e-5,
    ):
        """
        Args:
            normalized_shape (int or list or torch.Size) input shape from an expected input of size
            create_scale (bool) If set to False, the layer will not learn an additive weight, Default: True
            create_offset (bool) If set to False, the layer will not learn an additive bias, Default: True
            eps (float) a value added to the denominator for numerical stability. Default: 1e-5
        """
        super(FusedLayerNorm, self).__init__()

        if isinstance(normalized_shape, numbers.Integral):
            normalized_shape = (normalized_shape,)
        self.normalized_shape = torch.Size(normalized_shape)
        self.eps = eps
        if create_scale:
            self.weight = Parameter(torch.ones(*normalized_shape))
        else:
            self.weight = None

        if create_offset:
            self.bias = Parameter(torch.zeros(*normalized_shape))
        else:
            self.bias = None

        self.reset_parameters()

    def reset_parameters(self):
        if self.weight is not None:
            torch.nn.init.ones_(self.weight)
        if self.bias is not None:
            torch.nn.init.zeros_(self.bias)

    def forward(self, input):
        return FusedLayerNormAffineFunction.apply(
            input, self.weight, self.bias, self.normalized_shape, self.eps
        )


if __name__ == "__main__":
    dtype = torch.float32
    data = torch.rand(10, 10).cuda().to(dtype=dtype)
    data1 = data * 1
    data.requires_grad = True
    data1.requires_grad = True
    layer_norm = (
        FusedLayerNorm(10, create_scale=True, create_offset=True).cuda().to(dtype=dtype)
    )
    layer_norm_torch = torch.nn.LayerNorm(10).cuda().to(dtype=dtype)
    out = layer_norm(data)
    out1 = layer_norm_torch(data1)
    # print(out - out1)
    loss = out.sum()
    loss.backward()
    loss1 = out1.sum()
    loss1.backward()
    print(data.grad - data1.grad)
    print(layer_norm.weight.grad - layer_norm_torch.weight.grad)
    print(layer_norm.bias.grad - layer_norm_torch.bias.grad)
    print(layer_norm.weight.grad, layer_norm.bias.grad)