# SPDX-FileCopyrightText: © 2024 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 import ttnn from models.common.lightweightmodule import LightweightModule from models.common.rmsnorm import RMSNorm from models.tt_transformers.tt.attention import Attention as DefaultAttention from models.tt_transformers.tt.common import Mode from models.tt_transformers.tt.distributed_norm import DistributedNorm from models.tt_transformers.tt.mixtral_mlp import TtMixtralMLP from models.tt_transformers.tt.mixtral_moe import TtMoeLayer from models.tt_transformers.tt.mlp import MLP from models.tt_transformers.tt.model_config import TensorGroup class TransformerBlock(LightweightModule): def __init__( self, args, mesh_device, tt_ccl, dtype, state_dict, layer_num, weight_cache_path, transformation_mats, paged_attention_config=None, use_paged_kv_cache=False, attention_class=None, prefetcher=None, ): super().__init__() self.mesh_device = mesh_device self.tt_ccl = tt_ccl self.prefetcher = prefetcher self.num_devices = args.num_devices self.args = args self.hidden_size = args.dim self.n_heads = args.n_heads self.head_dim = self.hidden_size // self.n_heads self.max_seq_len = args.max_seq_len self.dim = args.dim self.max_batch_size = args.max_batch_size self.n_kv_heads = args.n_kv_heads self.current = 0 self.model_config = args.get_model_config() self.is_mixture_of_experts = False self.layer_num = layer_num ActualAttentionClass = attention_class if attention_class is not None else DefaultAttention self.attention = ActualAttentionClass( mesh_device=mesh_device, tt_ccl=self.tt_ccl, args=args, state_dict=state_dict, weight_cache_path=weight_cache_path, layer_num=layer_num, dtype=dtype, transformation_mats=transformation_mats, configuration=args, paged_attention_config=paged_attention_config, use_paged_kv_cache=use_paged_kv_cache, prefetcher=prefetcher, ) if getattr(self.args, "is_mixture_of_experts", False): self.feed_forward = TtMoeLayer( mesh_device=mesh_device, state_dict=state_dict, experts=TtMixtralMLP( mesh_device=mesh_device, state_dict=state_dict, args=args, layer_num=layer_num, dtypes={ "w1": dtype, "w2": dtype, "w3": dtype, }, ), args=args, layer_num=layer_num, dtype=dtype, tt_ccl=self.tt_ccl, ) else: self.feed_forward = MLP( mesh_device=mesh_device, tt_ccl=self.tt_ccl, args=args, state_dict=state_dict, weight_cache_path=weight_cache_path, layer_num=layer_num, dtype=dtype, model_config=self.model_config, prefetcher=prefetcher, ) # TODO: remove after https://github.com/tenstorrent/tt-metal/issues/35650 is fixed extra_rmsnorm_kwargs = {} # Llama 8B on a Galaxy DP4 row submesh runs out of L1 with fp32 RMSNorm # accumulation, matching the existing Qwen workaround below. use_galaxy_row_submesh_rmsnorm_l1_workaround = ( args.base_model_name == "Llama-3.1-8B" and args.num_devices == 8 and args.mesh_device is not None and tuple(args.mesh_device.shape) == (1, 8) and ttnn.cluster.get_cluster_type() == ttnn.cluster.ClusterType.GALAXY ) if ( args.base_model_name in ( "Qwen2.5-7B", "Qwen2.5-VL-7B", ) or use_galaxy_row_submesh_rmsnorm_l1_workaround ): extra_rmsnorm_kwargs["fp32_dest_acc_en"] = False # Post-norm decoders (EXAONE-4.x) have no input_layernorm: attention reads # the raw residual stream (h = x + post_attention_layernorm(attn(x))). The # attention_norm slot degenerates to its gather role (norm=None), keeping # the fractured->replicated all-gather the norm normally provides. self.use_post_norm = getattr(args, "use_post_norm", False) self.attention_norm = DistributedNorm( None if self.use_post_norm else RMSNorm( device=mesh_device, dim=args.dim, eps=args.norm_eps, state_dict=state_dict, state_dict_prefix=args.get_state_dict_prefix("", layer_num), weight_cache_path=None if args.dummy_weights else weight_cache_path, weight_dtype=ttnn.bfloat16, weight_key="attention_norm", is_distributed=self.args.is_distributed_norm, add_unit_offset=self.args.rms_norm_add_unit_offset, ccl_topology=self.args.ccl_topology(), tt_ccl=self.tt_ccl, **extra_rmsnorm_kwargs, ), args, tt_ccl=self.tt_ccl, prefetcher=self.prefetcher, TG=args.is_galaxy, ag_config_key="ATTN_LN_AG_CONFIG", ) self.ff_norm = DistributedNorm( RMSNorm( device=mesh_device, dim=args.dim, eps=args.norm_eps, state_dict=state_dict, state_dict_prefix=args.get_state_dict_prefix("", layer_num), weight_cache_path=None if args.dummy_weights else weight_cache_path, weight_dtype=ttnn.bfloat16, weight_key="ffn_norm", is_distributed=self.args.is_distributed_norm, add_unit_offset=self.args.rms_norm_add_unit_offset, ccl_topology=self.args.ccl_topology(), tt_ccl=self.tt_ccl, **extra_rmsnorm_kwargs, ), args, tt_ccl=self.tt_ccl, prefetcher=self.prefetcher, TG=args.is_galaxy, ag_config_key="FFN_LN_AG_CONFIG", ) if f"layers.{layer_num}.pre_feedforward_layernorm.weight" in state_dict: self.pre_ff_norm = DistributedNorm( # pre_feedforward_layernorm RMSNorm( device=mesh_device, dim=args.dim, eps=args.norm_eps, state_dict=state_dict, add_unit_offset=self.args.rms_norm_add_unit_offset, state_dict_prefix=args.get_state_dict_prefix("", layer_num), weight_cache_path=None if args.dummy_weights else weight_cache_path, weight_dtype=ttnn.bfloat16, weight_key="pre_feedforward_layernorm", is_distributed=self.args.is_distributed_norm, ccl_topology=self.args.ccl_topology(), tt_ccl=self.tt_ccl, ), args, tt_ccl=self.tt_ccl, prefetcher=self.prefetcher, TG=args.is_galaxy, ) self.ff_norm.enable_all_gather = ( False # output of ff_norm should be sharded if model uses pre_ff_norm, so skip all_gather ) elif self.use_post_norm: # Post-norm (EXAONE-4.x): route forward() through the sandwich branch so # ff_norm (HF post_attention_layernorm) is applied to the attention OUTPUT # before the residual add, with a gather-only identity in the pre-MLP slot # (MLP reads the raw residual stream: out = h + post_ff_norm(mlp(h))). self.pre_ff_norm = DistributedNorm( None, args, tt_ccl=self.tt_ccl, prefetcher=self.prefetcher, TG=args.is_galaxy, ) self.ff_norm.enable_all_gather = ( False # output of ff_norm should be sharded if model uses pre_ff_norm, so skip all_gather ) else: # If pre_feedforward_layernorm is not in state_dict, we do not use it self.pre_ff_norm = None if f"layers.{layer_num}.post_feedforward_layernorm.weight" in state_dict: self.post_ff_norm = DistributedNorm( # post_feedforward_layernorm RMSNorm( device=mesh_device, dim=args.dim, eps=args.norm_eps, add_unit_offset=self.args.rms_norm_add_unit_offset, state_dict=state_dict, state_dict_prefix=args.get_state_dict_prefix("", layer_num), weight_cache_path=None if args.dummy_weights else weight_cache_path, weight_dtype=ttnn.bfloat16, weight_key="post_feedforward_layernorm", is_distributed=self.args.is_distributed_norm, ccl_topology=self.args.ccl_topology(), tt_ccl=self.tt_ccl, ), args, tt_ccl=self.tt_ccl, prefetcher=self.prefetcher, TG=args.is_galaxy, enable_all_gather=False, ) else: # If post_feedforward_layernorm is not in state_dict, we do not use it self.post_ff_norm = None def update_weights( self, layer_hf_state_dict: dict[str, ttnn.Tensor], *, hf_rope: bool = False, ) -> None: """Strict layer-local weight update from an HF-keyed dict of on-device 4D ttnn tensors. Keys are the suffix after ``model.layers.{i}.`` (e.g. ``self_attn.q_proj.weight``); ``Transformer.update_weights`` strips the layer prefix and routes each layer's slice here. Strict: missing required keys raise ``KeyError``; any unconsumed key (e.g. an extra bias, q/k-norm, or Gemma-style pre/post-FF norm not yet wired into the leaf ``update()`` methods) raises ``ValueError``. ``hf_rope`` is forwarded to ``Attention.update``. """ unconsumed = set(layer_hf_state_dict.keys()) def consume(key: str) -> ttnn.Tensor: if key not in layer_hf_state_dict: raise KeyError( f"TransformerBlock.update_weights (layer {self.layer_num}): " f"missing required HF key {key!r}" ) unconsumed.discard(key) return layer_hf_state_dict[key] self.attention.update( q_proj=consume("self_attn.q_proj.weight"), k_proj=consume("self_attn.k_proj.weight"), v_proj=consume("self_attn.v_proj.weight"), o_proj=consume("self_attn.o_proj.weight"), hf_rope=hf_rope, ) self.feed_forward.update( gate_proj=consume("mlp.gate_proj.weight"), up_proj=consume("mlp.up_proj.weight"), down_proj=consume("mlp.down_proj.weight"), ) self.attention_norm.update(weight=consume("input_layernorm.weight")) self.ff_norm.update(weight=consume("post_attention_layernorm.weight")) if unconsumed: sample = sorted(unconsumed)[:10] raise ValueError( f"TransformerBlock.update_weights (layer {self.layer_num}): " f"{len(unconsumed)} HF key(s) not consumed within this layer. " f"This usually means an extra weight (a bias, q/k norm, or " f"Gemma-style pre/post-FF norm) that the leaf .update() " f"doesn't yet support. Showing up to 10: {sample}" ) def forward( self, x: ttnn.Tensor, current_pos, rot_mats_global=None, rot_mats_local=None, user_id=0, mode="decode", page_table=None, chunk_page_table=None, chunk_start_idx=None, kv_cache=None, batch_size=1, ) -> ttnn.Tensor: TG = self.args.is_galaxy residual = x # x is fractured across devices and interleaved in DRAM (for prefill) and sharded in L1 (for decode) skip_mem_cfg = self.args.get_residual_mem_config(mode, self.prefetcher) assert ( x.memory_config() == skip_mem_cfg ), f"decoder input memcfg mismatch: {x.memory_config()} != {skip_mem_cfg}" # Choose the correct rotation matrices based on the mode rot_mats = ( rot_mats_local if (hasattr(self.attention, "is_sliding") and self.attention.is_sliding) else rot_mats_global ) # Norms take fractured inputs and output replicated across devices attn_norm_config = self.args.get_norm_config("attn", mode, self.prefetcher) attn_in = self.attention_norm(x, mode, norm_config=attn_norm_config) # Reshape to [B, 1, S_per_user, H] so attention infers batch_size from shape[0] if batch_size > 1: attn_in = ttnn.reshape(attn_in, [batch_size, 1, attn_in.shape[-2] // batch_size, -1]) attn_out = self.attention.forward( attn_in, current_pos, rot_mats, user_id, mode, page_table=page_table, chunk_page_table=chunk_page_table, chunk_start_idx=chunk_start_idx, kv_cache=kv_cache, ) # To match the batch-related reshape inside the attention module # Use the batch_size parameter instead of inferring from shape[-3] # because for [32, 1, S, H] tensors, shape[-3] is 1, not 32 # This reshape is only applicable in prefill mode with batched prefill if mode == Mode.PREFILL and batch_size > 1: residual = ttnn.reshape(residual, [1, 1, residual.shape[-2] * residual.shape[-3] * residual.shape[0], -1]) # TODO: create correct memory config in RopeSetup (issue is in ttnn.add op because of different shape in memory config for residual and rot_mats) attn_out = ttnn.to_memory_config(attn_out, skip_mem_cfg) if self.pre_ff_norm is None: hidden_states = ttnn.add( residual, attn_out, memory_config=skip_mem_cfg, dtype=ttnn.bfloat16 if TG else None ) residual = hidden_states if mode == "prefill": x.deallocate(True) else: hidden_states = attn_out ff_norm_config = self.args.get_norm_config("ff", mode, self.prefetcher) hidden_states = self.ff_norm(hidden_states, mode, norm_config=ff_norm_config) if self.pre_ff_norm is not None: # Mesh partition ff_norm output to match residual sharding, skip if using distributed norm, because output is already sharded if self.num_devices > 1 and not self.args.is_distributed_norm(mode): hidden_states = ttnn.mesh_partition( hidden_states, memory_config=hidden_states.memory_config(), dim=3, cluster_axis=1, ) hidden_states = ttnn.add( residual, hidden_states, memory_config=skip_mem_cfg, dtype=ttnn.bfloat16 if TG else None ) residual = hidden_states pre_ff_norm_config = self.args.get_norm_config("ff", mode, self.prefetcher) hidden_states = self.pre_ff_norm(hidden_states, mode, norm_config=pre_ff_norm_config) ttnn.deallocate(attn_out) if TG and mode == "decode": hidden_states = ttnn.to_memory_config(hidden_states, memory_config=self.args.get_mlp_act_mem_config(mode)) # MLP takes replicated inputs and produces fractured outputs hidden_states = self.feed_forward.forward(hidden_states, mode) activation_dtype = self.args.decoders_optimizations.get_tensor_dtype( decoder_id=self.layer_num, tensor=TensorGroup.ACTIVATION ) if self.post_ff_norm is not None: post_ff_norm_config = self.args.get_norm_config("ff", mode, self.prefetcher) hidden_states = self.post_ff_norm(hidden_states, mode, norm_config=post_ff_norm_config) # Gathered if self.num_devices > 1 and not self.args.is_distributed_norm(mode): hidden_states = ttnn.mesh_partition( hidden_states, memory_config=hidden_states.memory_config(), dim=3, cluster_axis=1, ) out = ttnn.add( residual, hidden_states, memory_config=skip_mem_cfg, dtype=self.args.ccl_dtype if TG and not self.args.is_distributed_norm(mode) else activation_dtype or ttnn.bfloat16, ) return out # fractured across devices