File size: 2,772 Bytes
740d966
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch,os,sys
import torch.nn as nn
import torch.nn.functional as F
code_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.append(f'{code_dir}/../')
from core.submodule import Conv2x_IN
import timm



class ContextNetSharedBackbone(nn.Module):
  def __init__(self, args, c04, c08, c16, output_dim=[(128,128,128), (128,128,128)], norm_fn='batch', downsample=3):
    super().__init__()
    self.args = args
    self.conv04 = nn.ModuleList([
      nn.Conv2d(c04, output_dim[0][0], kernel_size=3, padding=1),
      nn.Conv2d(c04, output_dim[1][0], kernel_size=3, padding=1),
    ])

  def forward(self, x4, x8, x16):
    outputs04 = []
    for i in range(len(self.conv04)):
      outputs04.append(self.conv04[i](x4))
    return (outputs04,)



class DepthAnythingFeature:
    model_configs = {
        'vitl': {'encoder': 'vitl', 'features': 256, 'out_channels': [256, 512, 1024, 1024]},
        'vitb': {'encoder': 'vitb', 'features': 128, 'out_channels': [96, 192, 384, 768]},
        'vits': {'encoder': 'vits', 'features': 64, 'out_channels': [48, 96, 192, 384]}
    }



class Feature(nn.Module):
    def __init__(self, args):
        super(Feature, self).__init__()
        self.args = args
        model = timm.create_model('edgenext_small', pretrained=True, features_only=False)
        self.stem = model.stem
        self.stages = model.stages
        chans = [48, 96, 160, 304]
        self.chans = chans
        vit_feat_dim = DepthAnythingFeature.model_configs[self.args.vit_size]['features']//2

        self.deconv32_16 = Conv2x_IN(chans[3], chans[2], deconv=True, concat=True)
        self.deconv16_8 = Conv2x_IN(chans[2]*2, chans[1], deconv=True, concat=True)
        self.deconv8_4 = Conv2x_IN(chans[1]*2, chans[0], deconv=True, concat=True)

        self.conv4 = nn.Conv2d(chans[0]*2, self.chans[0]*2+vit_feat_dim, kernel_size=1, stride=1, padding=0)

        self.d_out = [self.chans[0]*2+vit_feat_dim, self.chans[1]*2, self.chans[2]*2, self.chans[3]]


    def forward(self, x):
        B,C,H,W = x.shape
        if hasattr(self, 'stem'):
          x = self.stem(x)
          x4 = self.stages[0](x)
          x8 = self.stages[1](x4)
          x16 = self.stages[2](x8)
          x32 = self.stages[3](x16)
        else:
          intermediates = self.model.forward_intermediates(x, intermediates_only=True)
          x4, x8, x16, x32 = intermediates[-4:]

        with torch.profiler.record_function("feature_deconv"):
          x16 = self.deconv32_16(x32, x16)
          x8 = self.deconv16_8(x16, x8)
          x4 = self.deconv8_4(x8, x4)
          x4 = self.conv4(x4)
          if hasattr(self, 'conv8'):
            x8 = self.conv8(x8)
            x16 = self.conv16(x16)
            x32 = self.conv32(x32)
        return [x4, x8, x16, x32]