| 226 | |
| 227 | @SUBMODULES.register_module() |
| 228 | class FineMoGenTransformer(DiffusionTransformer): |
| 229 | |
| 230 | def __init__(self, |
| 231 | scale_func_cfg=None, |
| 232 | pose_encoder_cfg=None, |
| 233 | pose_decoder_cfg=None, |
| 234 | moe_route_loss_weight=1.0, |
| 235 | template_kl_loss_weight=0.0001, |
| 236 | **kwargs): |
| 237 | super().__init__(**kwargs) |
| 238 | self.scale_func_cfg = scale_func_cfg |
| 239 | self.joint_embed = PoseEncoder(**pose_encoder_cfg) |
| 240 | self.out = zero_module(PoseDecoder(**pose_decoder_cfg)) |
| 241 | self.moe_route_loss_weight = moe_route_loss_weight |
| 242 | self.template_kl_loss_weight = template_kl_loss_weight |
| 243 | |
| 244 | def build_temporal_blocks(self, sa_block_cfg, ca_block_cfg, ffn_cfg): |
| 245 | self.temporal_decoder_blocks = nn.ModuleList() |
| 246 | for i in range(self.num_layers): |
| 247 | if isinstance(ffn_cfg, list): |
| 248 | ffn_cfg_block = ffn_cfg[i] |
| 249 | else: |
| 250 | ffn_cfg_block = ffn_cfg |
| 251 | self.temporal_decoder_blocks.append( |
| 252 | DecoderLayer(ca_block_cfg=ca_block_cfg, ffn_cfg=ffn_cfg_block)) |
| 253 | |
| 254 | def scale_func(self, timestep): |
| 255 | scale = self.scale_func_cfg['scale'] |
| 256 | w = (1 - (1000 - timestep) / 1000) * scale + 1 |
| 257 | output = {'text_coef': w, 'none_coef': 1 - w} |
| 258 | return output |
| 259 | |
| 260 | def aux_loss(self): |
| 261 | aux_loss = 0 |
| 262 | kl_loss = 0 |
| 263 | for module in self.temporal_decoder_blocks: |
| 264 | if hasattr(module.ca_block, 'aux_loss'): |
| 265 | aux_loss = aux_loss + module.ca_block.aux_loss |
| 266 | if hasattr(module.ca_block, 'kl_loss'): |
| 267 | kl_loss = kl_loss + module.ca_block.kl_loss |
| 268 | losses = {} |
| 269 | if aux_loss > 0: |
| 270 | losses['moe_route_loss'] = aux_loss * self.moe_route_loss_weight |
| 271 | if kl_loss > 0: |
| 272 | losses['template_kl_loss'] = kl_loss * self.template_kl_loss_weight |
| 273 | return losses |
| 274 | |
| 275 | def get_precompute_condition(self, |
| 276 | text=None, |
| 277 | motion_length=None, |
| 278 | xf_out=None, |
| 279 | re_dict=None, |
| 280 | device=None, |
| 281 | sample_idx=None, |
| 282 | clip_feat=None, |
| 283 | **kwargs): |
| 284 | if xf_out is None: |
| 285 | xf_out = self.encode_text(text, clip_feat, device) |
nothing calls this directly
no outgoing calls
no test coverage detected