clef / code /models /tt_transformers /tt /decoder.py
tt-hous's picture
Add files using upload-large-folder tool
b025706 verified
Raw History Blame Contribute Delete
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