MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / separate_gating

Method separate_gating

sat/dit_video_concat.py:638–672  ·  view source on GitHub ↗
(
        self,
        hidden_states,
        outputs,
        gate, 
        aug_gate,
        cond_inds, 
        pred_inds
    )

Source from the content-addressed store, hash-verified

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,

Callers 1

layer_forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected