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

Class AlphaBlender

sat/sgm/modules/diffusionmodules/util.py:281–328  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

279
280
281class AlphaBlender(nn.Module):
282 strategies = ["learned", "fixed", "learned_with_images"]
283
284 def __init__(
285 self,
286 alpha: float,
287 merge_strategy: str = "learned_with_images",
288 rearrange_pattern: str = "b t -> (b t) 1 1",
289 ):
290 super().__init__()
291 self.merge_strategy = merge_strategy
292 self.rearrange_pattern = rearrange_pattern
293
294 assert merge_strategy in self.strategies, f"merge_strategy needs to be in {self.strategies}"
295
296 if self.merge_strategy == "fixed":
297 self.register_buffer("mix_factor", torch.Tensor([alpha]))
298 elif self.merge_strategy == "learned" or self.merge_strategy == "learned_with_images":
299 self.register_parameter("mix_factor", torch.nn.Parameter(torch.Tensor([alpha])))
300 else:
301 raise ValueError(f"unknown merge strategy {self.merge_strategy}")
302
303 def get_alpha(self, image_only_indicator: torch.Tensor) -> torch.Tensor:
304 if self.merge_strategy == "fixed":
305 alpha = self.mix_factor
306 elif self.merge_strategy == "learned":
307 alpha = torch.sigmoid(self.mix_factor)
308 elif self.merge_strategy == "learned_with_images":
309 assert image_only_indicator is not None, "need image_only_indicator ..."
310 alpha = torch.where(
311 image_only_indicator.bool(),
312 torch.ones(1, 1, device=image_only_indicator.device),
313 rearrange(torch.sigmoid(self.mix_factor), "... -> ... 1"),
314 )
315 alpha = rearrange(alpha, self.rearrange_pattern)
316 else:
317 raise NotImplementedError
318 return alpha
319
320 def forward(
321 self,
322 x_spatial: torch.Tensor,
323 x_temporal: torch.Tensor,
324 image_only_indicator: Optional[torch.Tensor] = None,
325 ) -> torch.Tensor:
326 alpha = self.get_alpha(image_only_indicator)
327 x = alpha.to(x_spatial.dtype) * x_spatial + (1.0 - alpha).to(x_spatial.dtype) * x_temporal
328 return x

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected