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

Method __init__

sat/sgm/modules/autoencoding/temporal_ae.py:109–134  ·  view source on GitHub ↗
(self, in_channels: int, alpha: float = 0, merge_strategy: str = "learned")

Source from the content-addressed store, hash-verified

107
108class VideoBlock(AttnBlock):
109 def __init__(self, in_channels: int, alpha: float = 0, merge_strategy: str = "learned"):
110 super().__init__(in_channels)
111 # no context, single headed, as in base class
112 self.time_mix_block = VideoTransformerBlock(
113 dim=in_channels,
114 n_heads=1,
115 d_head=in_channels,
116 checkpoint=False,
117 ff_in=True,
118 attn_mode="softmax",
119 )
120
121 time_embed_dim = self.in_channels * 4
122 self.video_time_embed = torch.nn.Sequential(
123 torch.nn.Linear(self.in_channels, time_embed_dim),
124 torch.nn.SiLU(),
125 torch.nn.Linear(time_embed_dim, self.in_channels),
126 )
127
128 self.merge_strategy = merge_strategy
129 if self.merge_strategy == "fixed":
130 self.register_buffer("mix_factor", torch.Tensor([alpha]))
131 elif self.merge_strategy == "learned":
132 self.register_parameter("mix_factor", torch.nn.Parameter(torch.Tensor([alpha])))
133 else:
134 raise ValueError(f"unknown merge strategy {self.merge_strategy}")
135
136 def forward(self, x, timesteps, skip_video=False):
137 if skip_video:

Callers

nothing calls this directly

Calls 2

__init__Method · 0.45

Tested by

no test coverage detected