@@ -215,7 +215,11 @@ def __init__(
215215 ):
216216 self .feed_forward = LoRAFeedForward (args .dim , args .hidden_dim , args )
217217 else :
218- self .feed_forward = FeedForward (dim = args .dim , hidden_dim = args .hidden_dim )
218+ self .feed_forward = FeedForward (
219+ dim = args .dim ,
220+ hidden_dim = args .hidden_dim ,
221+ act_fn = args .act_fn .get_function (),
222+ )
219223
220224 if isinstance (self .attention , AttentionSkip ):
221225 self .attention_norm = nn .Identity ()
@@ -357,6 +361,7 @@ def __init__(self, params: ModelArgs, layers: nn.ModuleList, rope: Rope):
357361 self .output_prune_map = params .output_prune_map
358362 # YOCO (You Only Cache Once) KV sharing configuration.
359363 self .num_kv_shared_layers = params .num_kv_shared_layers
364+ self .layer_types = params .layer_types
360365
361366 def _forward_layers (
362367 self ,
@@ -365,6 +370,7 @@ def _forward_layers(
365370 freqs_sin : torch .Tensor ,
366371 attn_options_ : Dict ,
367372 seqlen : int ,
373+ freqs_by_type : Optional [Dict [str , Tuple [torch .Tensor , torch .Tensor ]]] = None ,
368374 ) -> Tuple [torch .Tensor , Optional [Any ]]:
369375 """Run transformer layers with YOCO KV sharing support."""
370376 attn_options_update = None
@@ -383,7 +389,14 @@ def _forward_layers(
383389 if donor_idx in shared_kv :
384390 attn_options_ ["shared_kv" ] = shared_kv [donor_idx ]
385391
386- h , attn_options_update = layer (h , freqs_cos , freqs_sin , attn_options_ )
392+ # Per-layer-type RoPE: select freqs based on layer type when available.
393+ l_cos , l_sin = freqs_cos , freqs_sin
394+ if freqs_by_type is not None and self .layer_types is not None :
395+ layer_type = self .layer_types [layer_idx ]
396+ if layer_type in freqs_by_type :
397+ l_cos , l_sin = freqs_by_type [layer_type ]
398+
399+ h , attn_options_update = layer (h , l_cos , l_sin , attn_options_ )
387400
388401 if _is_kv_donor_layer (layer_idx , self .n_layers , self .num_kv_shared_layers ):
389402 assert (
@@ -425,10 +438,23 @@ def forward(
425438 attn_options .get ("input_pos" ), seqlen
426439 )
427440
441+ # Compute per-layer-type freqs when per-layer RoPE is configured.
442+ freqs_by_type = None
443+ if hasattr (self , "ropes" ):
444+ input_pos = attn_options .get ("input_pos" )
445+ freqs_by_type = {
446+ lt : r .get_freqs (input_pos , seqlen ) for lt , r in self .ropes .items ()
447+ }
448+
428449 attn_options_ = attn_options .copy () if attn_options is not None else {}
429450
430451 h , attn_options_update = self ._forward_layers (
431- h , freqs_cos , freqs_sin , attn_options_ , seqlen
452+ h ,
453+ freqs_cos ,
454+ freqs_sin ,
455+ attn_options_ ,
456+ seqlen ,
457+ freqs_by_type = freqs_by_type ,
432458 )
433459
434460 if not self .generate_full_logits :
@@ -467,11 +493,35 @@ def forward(
467493 return logits
468494
469495
496+ def _build_ropes (model_args : ModelArgs ) -> Tuple [Rope , Dict [str , Rope ]]:
497+ """Build Rope instances, creating per-layer-type ropes when rope_parameters is set.
498+
499+ Returns (default_rope, ropes_by_type). ropes_by_type is empty when no
500+ per-layer-type configuration is provided.
501+ """
502+ import copy as _copy
503+
504+ if not model_args .rope_parameters :
505+ return Rope (model_args ), {}
506+
507+ ropes : Dict [str , Rope ] = {}
508+ for layer_type , rope_params in model_args .rope_parameters .items ():
509+ rope_args = _copy .copy (model_args )
510+ if "rope_theta" in rope_params :
511+ rope_args .rope_theta = rope_params ["rope_theta" ]
512+ rope_args .rope_freq_base = rope_params ["rope_theta" ]
513+ if "partial_rotary_factor" in rope_params :
514+ rope_args .partial_rotary_factor = rope_params ["partial_rotary_factor" ]
515+ ropes [layer_type ] = Rope (rope_args )
516+ return next (iter (ropes .values ())), ropes
517+
518+
470519def construct_transformer (model_args : ModelArgs ) -> Transformer :
471520 """
472521 Construct a Transformer model from the given model arguments.
473522 """
474- rope = Rope (model_args )
523+ rope , ropes = _build_ropes (model_args )
524+
475525 if model_args .attention_type not in ATTENTION_REGISTRY :
476526 raise ValueError (
477527 f"Unknown attention type: { model_args .attention_type } . "
@@ -521,12 +571,20 @@ def construct_transformer(model_args: ModelArgs) -> Transformer:
521571 )
522572 layers .append (transformer_block )
523573 else :
574+ # Select per-layer-type RoPE when available.
575+ layer_rope = rope
576+ if ropes and model_args .layer_types :
577+ layer_type = model_args .layer_types [layer_id ]
578+ layer_rope = ropes .get (layer_type , rope )
524579 attention = cls (
525- model_args , layer_id , rope , ** model_args .attention_kwargs
580+ model_args , layer_id , layer_rope , ** model_args .attention_kwargs
526581 ) # pyre-ignore[45]
527582 transformer_block = TransformerBlock (
528583 model_args , attention , layer_id = layer_id
529584 )
530585 layers .append (transformer_block )
531586
532- return Transformer (model_args , layers , rope )
587+ transformer = Transformer (model_args , layers , rope )
588+ if ropes :
589+ transformer .ropes = torch .nn .ModuleDict (ropes )
590+ return transformer
0 commit comments