(
cls,
images: torch.Tensor,
mask: torch.Tensor,
)
| 125 | |
| 126 | @classmethod |
| 127 | def execute( |
| 128 | cls, |
| 129 | images: torch.Tensor, |
| 130 | mask: torch.Tensor, |
| 131 | ) -> io.NodeOutput: |
| 132 | if mask.ndim == 4: |
| 133 | mask = mask[:, :, :, 0] |
| 134 | |
| 135 | if mask.shape[0] == 1 and images.shape[0] > 1: |
| 136 | mask = mask.expand(images.shape[0], -1, -1) |
| 137 | |
| 138 | min_frames = min(mask.shape[0], images.shape[0]) |
| 139 | mask = mask[:min_frames] |
| 140 | images = images[:min_frames] |
| 141 | |
| 142 | mask_4d = mask.unsqueeze(-1) # (B, H, W, 1) for broadcasting |
| 143 | bg_color = torch.tensor(_BG_COLOR_RGB).float().to(images.device) / 255 |
| 144 | |
| 145 | result = images * (1 - mask_4d) + bg_color.view(1, 1, 1, 3) * mask_4d |
| 146 | return io.NodeOutput(result) |
nothing calls this directly
no outgoing calls
no test coverage detected