MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / __init__

Method __init__

sat/sgm/modules/autoencoding/temporal_ae.py:19–52  ·  view source on GitHub ↗
(
        self,
        out_channels,
        *args,
        dropout=0.0,
        video_kernel_size=3,
        alpha=0.0,
        merge_strategy="learned",
        **kwargs,
    )

Source from the content-addressed store, hash-verified

17
18class VideoResBlock(ResnetBlock):
19 def __init__(
20 self,
21 out_channels,
22 *args,
23 dropout=0.0,
24 video_kernel_size=3,
25 alpha=0.0,
26 merge_strategy="learned",
27 **kwargs,
28 ):
29 super().__init__(out_channels=out_channels, dropout=dropout, *args, **kwargs)
30 if video_kernel_size is None:
31 video_kernel_size = [3, 1, 1]
32 self.time_stack = ResBlock(
33 channels=out_channels,
34 emb_channels=0,
35 dropout=dropout,
36 dims=3,
37 use_scale_shift_norm=False,
38 use_conv=False,
39 up=False,
40 down=False,
41 kernel_size=video_kernel_size,
42 use_checkpoint=False,
43 skip_t_emb=True,
44 )
45
46 self.merge_strategy = merge_strategy
47 if self.merge_strategy == "fixed":
48 self.register_buffer("mix_factor", torch.Tensor([alpha]))
49 elif self.merge_strategy == "learned":
50 self.register_parameter("mix_factor", torch.nn.Parameter(torch.Tensor([alpha])))
51 else:
52 raise ValueError(f"unknown merge strategy {self.merge_strategy}")
53
54 def get_alpha(self, bs):
55 if self.merge_strategy == "fixed":

Callers

nothing calls this directly

Calls 2

ResBlockClass · 0.90
__init__Method · 0.45

Tested by

no test coverage detected