@@ -423,8 +423,9 @@ def __init__(self, d_in, d_out, bias: bool = True, fast_init: bool = False):
423423 super ().__init__ (d_in , d_out , bias = bias )
424424 if not fast_init :
425425 nn .init .trunc_normal_ (self .weight , std = 0.02 )
426- u = torch .linalg .svd (self .weight .T , full_matrices = False )[- 1 ][0 ].detach ()
427- v = torch .linalg .svd (self .weight , full_matrices = False )[- 1 ][0 ].detach ()
426+ U , _S , Vh = torch .linalg .svd (self .weight , full_matrices = False )
427+ u = U [:, 0 ].detach ()
428+ v = Vh [0 ].detach ()
428429 else :
429430 # initialization from legacy version used in the original AMAGO paper.
430431 # This was a guess based on the sigma reparam pseudocode before the code was released,
@@ -469,10 +470,12 @@ def __init__(
469470 dropout_qkv : float = 0.0 ,
470471 head_scaling : bool = True ,
471472 sigma_reparam : bool = True ,
473+ use_rope : bool = False ,
472474 ):
473475 super ().__init__ ()
474476 assert isinstance (self_attention , SelfAttention )
475477 self .self_attention = self_attention
478+ self .rope = RotaryPosEmb (d_qkv ) if use_rope else None
476479 FF = SigmaReparam if sigma_reparam else nn .Linear
477480 self .qkv_projection = FF (d_model , 3 * d_qkv * n_heads , bias = False )
478481 self .dropout_qkv = nn .Dropout (dropout_qkv )
@@ -482,14 +485,23 @@ def __init__(
482485 )
483486 self .n_heads = n_heads
484487
485- def forward (self , sequence , key_cache = None , val_cache = None , cache_seqlens = None ):
488+ def forward (
489+ self ,
490+ sequence ,
491+ key_cache = None ,
492+ val_cache = None ,
493+ cache_seqlens = None ,
494+ pos_idxs = None ,
495+ ):
486496 qkv = self .dropout_qkv (self .qkv_projection (sequence ))
487497 qkv = rearrange (
488498 qkv ,
489499 "batch len (three d_qkv heads) -> batch len three heads d_qkv" ,
490500 heads = self .n_heads ,
491501 three = 3 ,
492502 )
503+ if self .rope is not None :
504+ qkv = self .rope (qkv , pos_idxs )
493505 out = self .head_scaler * self .self_attention (
494506 qkv = qkv ,
495507 key_cache = key_cache ,
@@ -538,10 +550,21 @@ def __init__(
538550 self .d_model = d_model
539551
540552 @torch .compile
541- def forward (self , self_seq , key_cache = None , val_cache = None , cache_seqlens = None ):
553+ def forward (
554+ self ,
555+ self_seq ,
556+ key_cache = None ,
557+ val_cache = None ,
558+ cache_seqlens = None ,
559+ pos_idxs = None ,
560+ ):
542561 q1 = self .norm1 (self_seq ) # pre-norm
543562 q1 = self .attention_layer (
544- q1 , key_cache = key_cache , val_cache = val_cache , cache_seqlens = cache_seqlens
563+ q1 ,
564+ key_cache = key_cache ,
565+ val_cache = val_cache ,
566+ cache_seqlens = cache_seqlens ,
567+ pos_idxs = pos_idxs ,
545568 )
546569 q1 = self .norm2 (q1 ) # normformer extra norm 1
547570 self_seq = self_seq + q1
@@ -667,6 +690,39 @@ def forward(self, pos_idxs: torch.LongTensor):
667690 return self .embeddings (pos_idxs )
668691
669692
693+ class RotaryPosEmb (nn .Module ):
694+ """Rotary Position Embedding (RoPE). Half-split (Llama-style) convention."""
695+
696+ def __init__ (self , head_dim : int , base : float = 10000.0 ):
697+ super ().__init__ ()
698+ assert head_dim % 2 == 0 , f"RoPE requires even head_dim, got { head_dim } "
699+ inv_freq = 1.0 / (base ** (torch .arange (0 , head_dim , 2 ).float () / head_dim ))
700+ self .register_buffer ("inv_freq" , inv_freq , persistent = False )
701+ self .head_dim = head_dim
702+
703+ def forward (self , qkv : torch .Tensor , pos_idxs : torch .Tensor ) -> torch .Tensor :
704+ """Rotate Q and K in packed (B, L, 3, H, D) QKV; V is unchanged."""
705+ pos = pos_idxs .squeeze (- 1 ).to (self .inv_freq .dtype )
706+ freqs = torch .einsum ("bl,d->bld" , pos , self .inv_freq )
707+ cos = freqs .cos ().unsqueeze (2 )
708+ sin = freqs .sin ().unsqueeze (2 )
709+
710+ q , k , v = qkv .unbind (dim = 2 )
711+ q = self ._apply_rotary (q , cos , sin ).to (q .dtype )
712+ k = self ._apply_rotary (k , cos , sin ).to (k .dtype )
713+ return torch .stack ([q , k , v ], dim = 2 )
714+
715+ @staticmethod
716+ def _apply_rotary (
717+ x : torch .Tensor , cos : torch .Tensor , sin : torch .Tensor
718+ ) -> torch .Tensor :
719+ """x: (B, L, H, D), cos/sin: (B, L, 1, D/2)."""
720+ d_half = x .shape [- 1 ] // 2
721+ x1 = x [..., :d_half ]
722+ x2 = x [..., d_half :]
723+ return torch .cat ([x1 * cos - x2 * sin , x1 * sin + x2 * cos ], dim = - 1 )
724+
725+
670726class Transformer (nn .Module ):
671727 """Build a full Transformer model from a list of layers."""
672728
@@ -680,13 +736,16 @@ def __init__(
680736 pos_emb : str = "fixed" ,
681737 ):
682738 super ().__init__ ()
739+ self .use_rope = pos_emb == "rope"
683740 if pos_emb == "fixed" :
684741 self .position_embedding = FixedPosEmb (d_model )
685742 elif pos_emb == "learnable" :
686743 self .position_embedding = LearnablePosEmb (d_model )
744+ elif pos_emb == "rope" :
745+ self .position_embedding = None
687746 else :
688747 raise ValueError (
689- f"Unrecognized pos_emb: { pos_emb } . Options are 'fixed' or 'learnable '."
748+ f"Unrecognized pos_emb: { pos_emb } . Options are 'fixed', 'learnable', or 'rope '."
690749 )
691750 self .inp = nn .Linear (inp_dim , d_model )
692751 self .dropout = nn .Dropout (dropout_emb )
@@ -701,20 +760,22 @@ def emb_dim(self):
701760 return self .d_model
702761
703762 def preprocess_seq (self , seq , pos_idxs ):
704- pos_emb = self .position_embedding (pos_idxs .squeeze (- 1 ))
705763 traj_emb = self .inp (seq )
706- traj_emb = self .dropout (traj_emb + pos_emb )
764+ if self .position_embedding is not None :
765+ pos_emb = self .position_embedding (pos_idxs .squeeze (- 1 ))
766+ traj_emb = traj_emb + pos_emb
767+ traj_emb = self .dropout (traj_emb )
707768 return traj_emb
708769
709770 @torch .compile
710- def training_forward (self , seq ):
771+ def training_forward (self , seq , pos_idxs = None ):
711772 for layer in self .layers :
712- seq = layer (seq )
773+ seq = layer (seq , pos_idxs = pos_idxs )
713774 return self .norm (seq )
714775
715- def inference_forward (self , seq , hidden_state ):
776+ def inference_forward (self , seq , hidden_state , pos_idxs = None ):
716777 for i , layer in enumerate (self .layers ):
717- seq = layer (seq , * hidden_state [i ])
778+ seq = layer (seq , * hidden_state [i ], pos_idxs = pos_idxs )
718779 return self .norm (seq )
719780
720781 def forward (self , seq , pos_idxs , hidden_state : Optional [TformerHiddenState ] = None ):
@@ -731,10 +792,11 @@ def forward(self, seq, pos_idxs, hidden_state: Optional[TformerHiddenState] = No
731792 """
732793
733794 traj_emb = self .preprocess_seq (seq , pos_idxs )
795+ rope_pos = pos_idxs if self .use_rope else None
734796 if hidden_state is not None :
735797 assert not self .training
736- traj_emb = self .inference_forward (traj_emb , hidden_state )
798+ traj_emb = self .inference_forward (traj_emb , hidden_state , pos_idxs = rope_pos )
737799 hidden_state .update ()
738800 else :
739- traj_emb = self .training_forward (traj_emb )
801+ traj_emb = self .training_forward (traj_emb , pos_idxs = rope_pos )
740802 return traj_emb , hidden_state
0 commit comments