| 121 | |
| 122 | |
| 123 | def forward(self, sample): |
| 124 | # 1. pre-process |
| 125 | hidden_states = rearrange(sample, "C T H W -> T C H W") |
| 126 | hidden_states = hidden_states / self.scaling_factor |
| 127 | hidden_states = self.conv_in(hidden_states) |
| 128 | time_emb, text_emb, res_stack = None, None, None |
| 129 | |
| 130 | # 2. blocks |
| 131 | for i, block in enumerate(self.blocks): |
| 132 | hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack) |
| 133 | |
| 134 | # 3. output |
| 135 | hidden_states = self.conv_norm_out(hidden_states) |
| 136 | hidden_states = self.conv_act(hidden_states) |
| 137 | hidden_states = self.conv_out(hidden_states) |
| 138 | hidden_states = rearrange(hidden_states, "T C H W -> C T H W") |
| 139 | hidden_states = self.time_conv_out(hidden_states) |
| 140 | |
| 141 | return hidden_states |
| 142 | |
| 143 | |
| 144 | def build_mask(self, data, is_bound): |