MCPcopy Create free account
hub / github.com/apple/ml-pointersect / forward

Method forward

cdslib/core/nn/modules/subspace.py:464–498  ·  view source on GitHub ↗

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",
    )

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls 1

subspace_funMethod · 0.95

Tested by

no test coverage detected