With CP, the function assume that the input tensor is already split. It delegates the embedding generation to generate_embeddings function.
(self, x_B_T_H_W_C: torch.Tensor, fps: Optional[torch.Tensor])
| 470 | return 1 |
| 471 | |
| 472 | def forward(self, x_B_T_H_W_C: torch.Tensor, fps: Optional[torch.Tensor]) -> torch.Tensor: |
| 473 | """ |
| 474 | With CP, the function assume that the input tensor is already split. |
| 475 | It delegates the embedding generation to generate_embeddings function. |
| 476 | """ |
| 477 | B_T_H_W_C = x_B_T_H_W_C.shape |
| 478 | if self._cp_group is not None: |
| 479 | cp_ranks = get_process_group_ranks(self._cp_group) |
| 480 | cp_size = len(cp_ranks) |
| 481 | B, T, H, W, C = B_T_H_W_C |
| 482 | B_T_H_W_C = torch.Size((B, T * cp_size, H, W, C)) |
| 483 | embeddings = self.generate_embeddings(B_T_H_W_C, fps=fps) |
| 484 | |
| 485 | return embeddings |
| 486 | |
| 487 | def generate_embeddings(self, B_T_H_W_C: torch.Size, fps: Optional[torch.Tensor]) -> Any: |
| 488 | raise NotImplementedError |
nothing calls this directly
no test coverage detected