| 171 | |
| 172 | @SUBMODULES.register_module() |
| 173 | class ACTORDecoder(BaseModule): |
| 174 | |
| 175 | def __init__(self, |
| 176 | max_seq_len=16, |
| 177 | njoints=None, |
| 178 | nfeats=None, |
| 179 | input_feats=None, |
| 180 | input_dim=256, |
| 181 | latent_dim=256, |
| 182 | condition_dim=None, |
| 183 | num_heads=4, |
| 184 | ff_size=1024, |
| 185 | num_layers=8, |
| 186 | activation='gelu', |
| 187 | dropout=0.1, |
| 188 | use_condition=False, |
| 189 | num_class=None, |
| 190 | pos_embedding='sinusoidal', |
| 191 | init_cfg=None): |
| 192 | super().__init__(init_cfg=init_cfg) |
| 193 | if input_dim != latent_dim: |
| 194 | self.linear = nn.Linear(input_dim, latent_dim) |
| 195 | else: |
| 196 | self.linear = nn.Identity() |
| 197 | self.njoints = njoints |
| 198 | self.nfeats = nfeats |
| 199 | if input_feats is None: |
| 200 | assert self.njoints is not None and self.nfeats is not None |
| 201 | self.input_feats = njoints * nfeats |
| 202 | else: |
| 203 | self.input_feats = input_feats |
| 204 | self.max_seq_len = max_seq_len |
| 205 | self.input_dim = input_dim |
| 206 | self.latent_dim = latent_dim |
| 207 | self.condition_dim = condition_dim |
| 208 | self.use_condition = use_condition |
| 209 | self.num_class = num_class |
| 210 | if self.use_condition: |
| 211 | if num_class is None: |
| 212 | self.condition_bias = build_MLP(condition_dim, latent_dim) |
| 213 | else: |
| 214 | self.condition_bias = nn.Parameter(torch.randn(num_class, latent_dim)) |
| 215 | if pos_embedding == 'sinusoidal': |
| 216 | self.pos_encoder = SinusoidalPositionalEncoding(latent_dim, dropout) |
| 217 | else: |
| 218 | self.pos_encoder = LearnedPositionalEncoding(latent_dim, dropout, max_len=max_seq_len) |
| 219 | seqTransDecoderLayer = nn.TransformerDecoderLayer( |
| 220 | d_model=self.latent_dim, |
| 221 | nhead=num_heads, |
| 222 | dim_feedforward=ff_size, |
| 223 | dropout=dropout, |
| 224 | activation=activation) |
| 225 | self.seqTransDecoder = nn.TransformerDecoder( |
| 226 | seqTransDecoderLayer, |
| 227 | num_layers=num_layers) |
| 228 | |
| 229 | self.final = nn.Linear(self.latent_dim, self.input_feats) |
| 230 |
nothing calls this directly
no outgoing calls
no test coverage detected