(
v0: torch.Tensor, # [B, C, T, H, W]
v1: torch.Tensor, # [B, C, T, H, W]
)
| 325 | |
| 326 | |
| 327 | def project( |
| 328 | v0: torch.Tensor, # [B, C, T, H, W] |
| 329 | v1: torch.Tensor, # [B, C, T, H, W] |
| 330 | ): |
| 331 | dtype = v0.dtype |
| 332 | v0, v1 = v0.double(), v1.double() |
| 333 | v1 = torch.nn.functional.normalize(v1, dim=[-1, -2, -3, -4]) |
| 334 | v0_parallel = (v0 * v1).sum(dim=[-1, -2, -3, -4], keepdim=True) * v1 |
| 335 | v0_orthogonal = v0 - v0_parallel |
| 336 | return v0_parallel.to(dtype), v0_orthogonal.to(dtype) |
| 337 | |
| 338 | |
| 339 | def adaptive_projected_guidance( |
no outgoing calls
no test coverage detected