Download code/models/tt_transformers/tt/decoder.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 17.3 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/decoder.py
- Command line
-
hf download hf://tt-hous/clef/code/models/tt_transformers/tt/decoder.py
-
curl -L -o decoder.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/decoder.py
17.3 kB
| # 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 | |