(
self,
*,
device: torch.device,
param_shapes: Dict[str, Tuple[int]],
params_proj: Dict[str, Any],
d_latent: int,
latent_warp: Optional[Dict[str, Any]] = None,
renderer: Renderer,
)
| 144 | |
| 145 | class VectorDecoder(nn.Module): |
| 146 | def __init__( |
| 147 | self, |
| 148 | *, |
| 149 | device: torch.device, |
| 150 | param_shapes: Dict[str, Tuple[int]], |
| 151 | params_proj: Dict[str, Any], |
| 152 | d_latent: int, |
| 153 | latent_warp: Optional[Dict[str, Any]] = None, |
| 154 | renderer: Renderer, |
| 155 | ): |
| 156 | super().__init__() |
| 157 | self.device = device |
| 158 | self.param_shapes = param_shapes |
| 159 | |
| 160 | if latent_warp is None: |
| 161 | latent_warp = dict(name="identity") |
| 162 | self.d_latent = d_latent |
| 163 | self.params_proj = params_proj_from_config( |
| 164 | params_proj, device=device, param_shapes=param_shapes, d_latent=d_latent |
| 165 | ) |
| 166 | self.latent_warp = latent_warp_from_config(latent_warp, device=device) |
| 167 | self.renderer = renderer |
| 168 | |
| 169 | def bottleneck_to_params( |
| 170 | self, vector: torch.Tensor, options: Optional[AttrDict] = None |
nothing calls this directly
no test coverage detected