MCPcopy Create free account
hub / github.com/deepbrainai-research/float / FlowMatchingTransformer

Class FlowMatchingTransformer

models/float/FMT.py:194–336  ·  view source on GitHub ↗

Flow Matching Transformer (FMT)

Source from the content-addressed store, hash-verified

192
193
194class 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:

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected