MCPcopy Create free account
hub / github.com/drinkingcoder/NeuralMarker / __init__

Method __init__

core/raft.py:26–61  ·  view source on GitHub ↗
(self, args)

Source from the content-addressed store, hash-verified

24
25class RAFT(nn.Module):
26 def __init__(self, args):
27 super(RAFT, self).__init__()
28 self.args = args
29
30 if args.small:
31 self.hidden_dim = hdim = 96
32 self.context_dim = cdim = 64
33 args.corr_levels = 4
34 args.corr_radius = 3
35
36 else:
37 self.hidden_dim = hdim = 128
38 self.context_dim = cdim = 128
39 args.corr_levels = 4
40 args.corr_radius = 4
41
42 if 'dropout' not in self.args:
43 self.args.dropout = 0
44
45 if 'alternate_corr' not in self.args:
46 self.args.alternate_corr = False
47
48 # feature network, context network, and update block
49 if args.small:
50 self.fnet = SmallEncoder(output_dim=128, norm_fn='instance', dropout=args.dropout)
51 self.cnet = SmallEncoder(output_dim=hdim+cdim, norm_fn='none', dropout=args.dropout)
52 self.update_block = SmallUpdateBlock(self.args, hidden_dim=hdim)
53
54 else:
55 if self.args.fnet == 'CNN':
56 self.fnet = BasicEncoder(output_dim=256, norm_fn='instance', dropout=args.dropout)
57 self.cnet = BasicEncoder(output_dim=hdim+cdim, norm_fn='batch', dropout=args.dropout)
58 elif self.args.fnet == 'twins':
59 self.fnet = twins_svt_large(pretrained=True)
60 self.cnet = twins_svt_large(pretrained=True)
61 self.update_block = BasicUpdateBlock(self.args, hidden_dim=hdim)
62
63 def freeze_bn(self):
64 for m in self.modules():

Callers

nothing calls this directly

Calls 5

SmallEncoderClass · 0.90
SmallUpdateBlockClass · 0.90
BasicEncoderClass · 0.90
twins_svt_largeClass · 0.90
BasicUpdateBlockClass · 0.90

Tested by

no test coverage detected