| 636 | return video_latent |
| 637 | |
| 638 | def separate_gating( |
| 639 | self, |
| 640 | hidden_states, |
| 641 | outputs, |
| 642 | gate, |
| 643 | aug_gate, |
| 644 | cond_inds, |
| 645 | pred_inds |
| 646 | ): |
| 647 | hidden_states, outputs = map( |
| 648 | lambda x: rearrange(x, 'b (t n) d -> b t n d', t=self.compressed_num_frames), |
| 649 | (hidden_states, outputs) |
| 650 | ) |
| 651 | cond_hidden_states, pred_hidden_states = hidden_states[:, cond_inds], hidden_states[:, pred_inds] |
| 652 | cond_outputs, pred_outputs = outputs[:, cond_inds], outputs[:, pred_inds] |
| 653 | |
| 654 | cond_hidden_states, pred_hidden_states, cond_outputs, pred_outputs = map( |
| 655 | lambda x: rearrange(x, 'b t n d -> b (t n) d'), |
| 656 | (cond_hidden_states, pred_hidden_states, cond_outputs, pred_outputs) |
| 657 | ) |
| 658 | |
| 659 | cond_hidden_states = cond_hidden_states + aug_gate * cond_outputs |
| 660 | pred_hidden_states = pred_hidden_states + gate * pred_outputs |
| 661 | |
| 662 | # cond_hidden_states, pred_hidden_states = map( |
| 663 | # lambda x: rearrange(x, 'b (t n) d -> b t n d', t=self.compressed_num_frames), |
| 664 | # (cond_hidden_states, pred_hidden_states) |
| 665 | # ) |
| 666 | cond_hidden_states = rearrange(cond_hidden_states, 'b (t n) d -> b t n d', t=len(cond_inds)) |
| 667 | pred_hidden_states = rearrange(pred_hidden_states, 'b (t n) d -> b t n d', t=len(pred_inds)) |
| 668 | |
| 669 | hidden_states = torch.cat([cond_hidden_states, pred_hidden_states], dim=1) |
| 670 | hidden_states = rearrange(hidden_states, 'b t n d -> b (t n) d') |
| 671 | |
| 672 | return hidden_states |
| 673 | |
| 674 | def layer_forward( |
| 675 | self, |