File size: 2,127 Bytes
742c169
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
import torch.nn as nn
import torch.nn.functional as F

from .utils import compile_wrapper


class TemporalPatch(nn.Module):
    def __init__(
        self,
        input_channels,
        output_channels,
        patch_size=2,
        overlap=1,
        init_gamma=1.1,
    ):
        super(TemporalPatch, self).__init__()
        self.patch_size = patch_size
        self.overlap = overlap
        self.input_channels = input_channels
        self.output_channels = output_channels
        if input_channels != output_channels:
            self.fproj = nn.Conv1d(input_channels, output_channels, 1, 1, 0)
        else:
            self.fproj = nn.Identity()
        self.tproj = nn.Conv1d(
            output_channels, output_channels, patch_size + overlap, stride=patch_size
        )
        self.init_gamma = init_gamma
        self.init_weight()

    def init_weight(self):
        if isinstance(self.fproj, nn.Conv1d):
            nn.init.constant_(self.fproj.bias, 0)
            nn.init.normal_(self.fproj.weight, std=1 / self.input_channels**0.5)

        nn.init.constant_(self.tproj.bias, 0)
        tproj_patch = self.tproj.weight.size(-1)
        I = torch.eye(self.tproj.weight.size(0)).to(self.tproj.weight)[..., None]
        weight = I.repeat(1, 1, tproj_patch) * self.init_gamma
        scale = [self.init_gamma**i for i in range(tproj_patch)]
        scale = torch.tensor(scale).to(self.tproj.weight) / sum(scale)
        self.tproj.weight.data.copy_(weight * scale[None, None, :])

    @compile_wrapper
    def forward(self, x):
        """
        x: [B, T, L, C]
        """
        b, t, l, c = x.shape
        x = x.permute(0, 2, 3, 1).contiguous().view(b * l, c, t)

        x = self.fproj(x)
        first_frame = x[..., 0:1]
        x = self.tproj(x)
        x = torch.concat([first_frame, x], dim=-1)

        _, oc, ot = x.shape
        x = x.view(b, l, oc, ot).permute(0, 3, 1, 2).contiguous()
        return x


if __name__ == "__main__":
    x = torch.randn(2, 1 + 15 * 8, 3, 224, 224)
    m = TemporalPatch(3, 4, 8, 1)
    y = m(x)  # [B, 1 + (T-1)//p, 4, H, W]
    print(y.shape)