Flow Matching Transformer (FMT)
| 192 | |
| 193 | |
| 194 | class FlowMatchingTransformer(BaseModel): |
| 195 | """ |
| 196 | Flow Matching Transformer (FMT) |
| 197 | """ |
| 198 | def __init__(self, opt) -> None: |
| 199 | super().__init__() |
| 200 | self.opt = opt |
| 201 | |
| 202 | self.num_frames_for_clip = int(self.opt.wav2vec_sec * self.opt.fps) |
| 203 | self.num_prev_frames = int(opt.num_prev_frames) |
| 204 | self.num_total_frames = self.num_prev_frames + self.num_frames_for_clip |
| 205 | |
| 206 | self.hidden_size = opt.dim_h |
| 207 | self.mlp_ratio = opt.mlp_ratio |
| 208 | self.fmt_depth = opt.fmt_depth |
| 209 | self.num_heads = opt.num_heads |
| 210 | |
| 211 | self.x_embedder = SequenceEmbed(opt.dim_w, self.hidden_size) |
| 212 | |
| 213 | # video time position encoding |
| 214 | self.pos_embed = nn.Parameter(torch.zeros(1, self.num_total_frames, self.hidden_size), requires_grad=False) |
| 215 | |
| 216 | # flow trajectory time encoding |
| 217 | self.t_embedder = TimestepEmbedder(self.hidden_size) |
| 218 | self.c_embedder = nn.Linear(opt.dim_w + opt.dim_a + opt.dim_e, self.hidden_size) |
| 219 | |
| 220 | # define FMT blocks |
| 221 | self.blocks = nn.ModuleList([FMTBlock(self.hidden_size, self.num_heads, mlp_ratio=self.mlp_ratio) for _ in range(self.fmt_depth)]) |
| 222 | self.decoder = Decoder(self.hidden_size, self.opt.dim_w) |
| 223 | self.initialize_weights() |
| 224 | |
| 225 | # define alignment mask |
| 226 | alignment_mask = enc_dec_mask(self.num_total_frames, self.num_total_frames, 1, expansion=opt.attention_window).to(opt.rank) |
| 227 | self.register_buffer('alignment_mask', alignment_mask) |
| 228 | |
| 229 | |
| 230 | def initialize_weights(self) -> None: |
| 231 | def _basic_init(module): |
| 232 | if isinstance(module, nn.Linear): |
| 233 | torch.nn.init.xavier_uniform_(module.weight) |
| 234 | if module.bias is not None: |
| 235 | nn.init.constant_(module.bias, 0) |
| 236 | |
| 237 | self.apply(_basic_init) |
| 238 | |
| 239 | pos_embed = get_sinusoid_encoding_table(self.num_total_frames, self.hidden_size) |
| 240 | self.pos_embed.data.copy_(pos_embed.unsqueeze(0)) |
| 241 | |
| 242 | w = self.x_embedder.proj.weight.data |
| 243 | nn.init.xavier_uniform_(w.view([w.shape[0], -1])) |
| 244 | nn.init.constant_(self.x_embedder.proj.bias, 0) |
| 245 | |
| 246 | # Initialize timestep embedding MLP: |
| 247 | nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02) |
| 248 | nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02) |
| 249 | |
| 250 | # Zero-out adaLN modulation layers in FMT blocks: |
| 251 | for block in self.blocks: |