r""" Args: x `(*, in_features)`: input tensor styles `(*, style_features)` or list of tensors `(*, style_features)` broadcastable to x: style tensor Returns: `(*, out_features)`
(
self,
x: torch.Tensor,
styles: T.Union[T.Sequence[torch.Tensor], torch.Tensor],
use_subspace_fun: bool = False,
xls: T.Union[T.Sequence[torch.Tensor], torch.Tensor] = None,
xrs: T.Union[T.Sequence[torch.Tensor], torch.Tensor] = None,
mode: str = "A",
)
| 462 | self.addons.append(None) |
| 463 | |
| 464 | def forward( |
| 465 | self, |
| 466 | x: torch.Tensor, |
| 467 | styles: T.Union[T.Sequence[torch.Tensor], torch.Tensor], |
| 468 | use_subspace_fun: bool = False, |
| 469 | xls: T.Union[T.Sequence[torch.Tensor], torch.Tensor] = None, |
| 470 | xrs: T.Union[T.Sequence[torch.Tensor], torch.Tensor] = None, |
| 471 | mode: str = "A", |
| 472 | ): |
| 473 | r""" |
| 474 | Args: |
| 475 | x `(*, in_features)`: |
| 476 | input tensor |
| 477 | styles `(*, style_features)` or list of tensors `(*, style_features)` broadcastable to x: |
| 478 | style tensor |
| 479 | |
| 480 | Returns: |
| 481 | `(*, out_features)` |
| 482 | """ |
| 483 | |
| 484 | if use_subspace_fun: |
| 485 | return self.subspace_fun(xls=xls, xrs=xrs, mode=mode) |
| 486 | |
| 487 | else: |
| 488 | if isinstance(styles, torch.Tensor): |
| 489 | styles = [styles] * self.num_layers |
| 490 | assert len(styles) == self.num_layers |
| 491 | for i in range(self.num_layers): |
| 492 | assert styles[i].ndim == x.ndim |
| 493 | |
| 494 | for layer_idx in range(self.num_layers): |
| 495 | x = self.main[layer_idx](x, styles[layer_idx]) |
| 496 | x = self.addons[layer_idx](x) |
| 497 | |
| 498 | return x |
| 499 | |
| 500 | def subspace_fun( |
| 501 | self, |
nothing calls this directly
no test coverage detected