| 49 | |
| 50 | |
| 51 | class DiffusionTransformer(BaseModule, metaclass=ABCMeta): |
| 52 | |
| 53 | def __init__(self, |
| 54 | input_feats, |
| 55 | max_seq_len=240, |
| 56 | latent_dim=512, |
| 57 | time_embed_dim=2048, |
| 58 | num_layers=8, |
| 59 | sa_block_cfg=None, |
| 60 | ca_block_cfg=None, |
| 61 | ffn_cfg=None, |
| 62 | text_encoder=None, |
| 63 | use_pos_embedding=True, |
| 64 | use_residual_connection=False, |
| 65 | time_embedding_type='sinusoidal', |
| 66 | post_process_cfg=None, |
| 67 | init_cfg=None): |
| 68 | super().__init__(init_cfg=init_cfg) |
| 69 | self.input_feats = input_feats |
| 70 | self.max_seq_len = max_seq_len |
| 71 | self.latent_dim = latent_dim |
| 72 | self.num_layers = num_layers |
| 73 | self.time_embed_dim = time_embed_dim |
| 74 | self.use_pos_embedding = use_pos_embedding |
| 75 | if self.use_pos_embedding: |
| 76 | self.sequence_embedding = nn.Parameter( |
| 77 | torch.randn(max_seq_len, latent_dim)) |
| 78 | self.build_text_encoder(text_encoder) |
| 79 | |
| 80 | # Input Embedding |
| 81 | self.joint_embed = nn.Linear(self.input_feats, self.latent_dim) |
| 82 | |
| 83 | self.time_embedding_type = time_embedding_type |
| 84 | if time_embedding_type == 'learnable': |
| 85 | self.time_tokens = nn.Embedding(1000, self.latent_dim) |
| 86 | self.time_embed = nn.Sequential( |
| 87 | nn.Linear(self.latent_dim, self.time_embed_dim), |
| 88 | nn.SiLU(), |
| 89 | nn.Linear(self.time_embed_dim, self.time_embed_dim), |
| 90 | ) |
| 91 | self.build_temporal_blocks(sa_block_cfg, ca_block_cfg, ffn_cfg) |
| 92 | |
| 93 | # Output Module |
| 94 | self.out = zero_module(nn.Linear(self.latent_dim, self.input_feats)) |
| 95 | self.use_residual_connection = use_residual_connection |
| 96 | self.post_process_cfg = post_process_cfg |
| 97 | |
| 98 | def build_temporal_blocks(self, sa_block_cfg, ca_block_cfg, ffn_cfg): |
| 99 | self.temporal_decoder_blocks = nn.ModuleList() |
| 100 | for i in range(self.num_layers): |
| 101 | self.temporal_decoder_blocks.append( |
| 102 | DecoderLayer(sa_block_cfg=sa_block_cfg, |
| 103 | ca_block_cfg=ca_block_cfg, |
| 104 | ffn_cfg=ffn_cfg)) |
| 105 | |
| 106 | def build_text_encoder(self, text_encoder): |
| 107 | |
| 108 | text_latent_dim = text_encoder['latent_dim'] |
nothing calls this directly
no outgoing calls
no test coverage detected