# SPDX-FileCopyrightText: © 2023 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 import torch import ttnn from models.common.lightweightmodule import LightweightModule from models.common.utility_functions import copy_to_buffer from models.tt_transformers.tt.ccl import tt_all_reduce from models.tt_transformers.tt.common import Mode, pad_to_size from models.tt_transformers.tt.model_config import OpGroup, TensorGroup class MLP(LightweightModule): def __init__( self, mesh_device, tt_ccl, args, state_dict, weight_cache_path, layer_num, dtype, model_config, state_dict_prefix=None, prefetcher=None, ): super().__init__() self.mesh_device = mesh_device self.tt_ccl = tt_ccl self.args = args self.dim = args.dim self.model_config = model_config self.layer_num = layer_num # Define the prefetcher object self.prefetcher = prefetcher state_dict_prefix = state_dict_prefix or args.get_state_dict_prefix(self.__class__.__name__, layer_num) torch_weight = lambda name: torch.transpose(state_dict[f"{state_dict_prefix}.{name}.weight"], -2, -1) pad_hidden_dim = lambda tensor, dim: pad_to_size(tensor, dim=dim, size=args.hidden_dim) # If padding was applied (e.g. via env var), add the unpadded hidden dim to the cache name to avoid loading incorrect weights hidden_dim_string = f".hidden_dim_{args.hidden_dim}" if args.hidden_dim != args.unpadded_hidden_dim else "" if args.dummy_weights: cache_name = lambda _: None else: cache_name = lambda name: weight_cache_path / f"{state_dict_prefix}.{name}{hidden_dim_string}" w1_w3_mem_config = args.create_dram_sharded_mem_config(args.dim, args.hidden_dim // args.num_devices) w2_mem_config = args.create_dram_sharded_mem_config(args.hidden_dim // args.num_devices, args.dim) # TODO Clean up this code. With sharding, we load the normal weights and then shard them # Note: unsqueeze(0).unsqueeze(0) makes weights 4D [1, 1, H, W] to match attention weights # This is required for the dram_prefetcher to correctly interpret all weights def as_sharded_tensor(name, type, dims): # First get the raw weight and transpose it raw_weight = torch_weight(name[:2]) # This is 2D: [H, W] # Pad if needed padded_weight = pad_hidden_dim(raw_weight, dims[0] if args.is_galaxy else dims[-1]) # Make 4D: [1, 1, H, W] - CRITICAL for prefetcher to work correctly torch_tensor = padded_weight.unsqueeze(0).unsqueeze(0) result = ttnn.as_tensor( torch_tensor, dtype=type, device=self.mesh_device, mesh_mapper=ttnn.ShardTensor2dMesh(self.mesh_device, dims=dims, mesh_shape=args.cluster_shape), layout=ttnn.TILE_LAYOUT, memory_config=( ttnn.DRAM_MEMORY_CONFIG if args.is_galaxy else w2_mem_config if "w2" in name else w1_w3_mem_config ), cache_file_name=cache_name(name), ) return result # Sharded weights w1_dims = (-1, -2) if args.is_galaxy else (-2, -1) w2_dims = (-2, -1) if args.is_galaxy else (-1, -2) layer_num = max(layer_num, 0) # cross_block uses the configuration of the first decoder # When prefetcher is enabled, use consistent dtypes across all layers to avoid # race conditions caused by different block sizes use_prefetcher = prefetcher is not None self.decoders_optimizations = self.args.decoders_optimizations self.ff1_3_dtype = self.decoders_optimizations.get_tensor_dtype( decoder_id=layer_num, tensor=TensorGroup.FF1_FF3, prefetcher=use_prefetcher ) self.ff2_dtype = self.decoders_optimizations.get_tensor_dtype( decoder_id=layer_num, tensor=TensorGroup.FF2, prefetcher=use_prefetcher ) self.w1 = as_sharded_tensor( "w1_sharded", self.ff1_3_dtype, dims=w1_dims ) # bfp4 normally ok here but sub .99 pcc for llama 3.1 weights self.w2 = as_sharded_tensor("w2_sharded", self.ff2_dtype, dims=w2_dims) self.w3 = as_sharded_tensor("w3_sharded", self.ff1_3_dtype, dims=w1_dims) # Default activation is SILU self.activation_type = ( args.mlp_activation_type if hasattr(args, "mlp_activation_type") else ttnn.UnaryOpType.SILU ) # Insert the tensors into the prefetcher if it is used if self.prefetcher is not None: def register_weights(): self.prefetcher.insert_tensor(self.w1) self.prefetcher.insert_tensor(self.w3) self.prefetcher.insert_tensor(self.w2) self.prefetcher.register_callback(register_weights) def update( self, *, gate_proj: ttnn.Tensor, up_proj: ttnn.Tensor, down_proj: ttnn.Tensor, ) -> None: """In-place replace the on-device MLP weights via ``ttnn.copy``. HF-format input (see ``LLAMA_WEIGHT_TRANSFER.md``): ``gate_proj``, ``up_proj`` are ``(1, 1, I, H)`` and ``down_proj`` is ``(1, 1, H, I)`` (HF Linear wrapped in two unit dims; ``H=args.dim``, ``I=hidden_dim``), bf16, TILE, DRAM-interleaved, replicated. Internal storage is the HF weight transposed; ``update`` transposes each input on device (mirroring the constructor's ``torch.transpose``) then ``copy_to_buffer``s into the existing buffers, preserving addresses (so captured traces and the prefetcher's recorded addresses stay valid). Caveats: hidden-dim padding is not handled (asserted off for Llama-3.2-1B-Instruct); the multi-chip replicated -> 2D-sharded mesh projection is not inserted (a no-op on the 1x1 transfer case). """ assert self.args.num_devices == 1, ( f"MLP.update for num_devices > 1 is not yet implemented " f"(got num_devices={self.args.num_devices}); w1/w2/w3 are " "2D-sharded on a mesh and need a ttnn.mesh_partition into the " "sharded layout before copy." ) assert self.args.hidden_dim == self.args.unpadded_hidden_dim, ( f"MLP.update does not yet support hidden_dim padding " f"(hidden_dim={self.args.hidden_dim}, " f"unpadded_hidden_dim={self.args.unpadded_hidden_dim}); pad on the " "caller side or extend update() with an on-device ttnn.pad." ) w1_internal = ttnn.transpose(gate_proj, -2, -1) w3_internal = ttnn.transpose(up_proj, -2, -1) w2_internal = ttnn.transpose(down_proj, -2, -1) copy_to_buffer(w1_internal, self.w1, self.ff1_3_dtype) copy_to_buffer(w3_internal, self.w3, self.ff1_3_dtype) copy_to_buffer(w2_internal, self.w2, self.ff2_dtype) def forward(self, x: ttnn.Tensor, mode: Mode) -> ttnn.Tensor: """ w1 -> gate_proj w2 -> down_proj w3 -> up_proj HF reference: self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) """ seq_len = x.shape[-2] chunk_size = self.model_config.get("MLP_PREFILL_CHUNK_SIZE", seq_len) if mode == Mode.PREFILL and seq_len > chunk_size: # The MLP is independent across tokens. Bound its gate/up/product intermediates # without splitting attention, which may require a complete sliding window. outputs = [ self.forward(x[..., start : min(start + chunk_size, seq_len), :], mode) for start in range(0, seq_len, chunk_size) ] ttnn.deallocate(x) output = ttnn.concat(outputs, dim=-2, memory_config=outputs[0].memory_config()) for chunk in outputs: ttnn.deallocate(chunk) return output TG = self.args.is_galaxy layer_num = max(self.layer_num, 0) # cross_block uses the configuration of the first decoder activation_dtype = self.decoders_optimizations.get_tensor_dtype( decoder_id=layer_num, tensor=TensorGroup.ACTIVATION ) li_ff1_3_compute_kernel_cfg = self.decoders_optimizations.get_math_fidelity( decoder_id=layer_num, op=OpGroup.LI_FF1_FF3, configuration=self.args ) if mode == Mode.PREFILL and seq_len >= self.args.prefill_len_cutoff: # 512 if Blackhole, 1024 if Wormhole # Reshape input to to fit on device and parallelize computation x = ttnn.reshape(x, [1, seq_len // self.args.prefill_len_cutoff, self.args.prefill_len_cutoff, -1]) # In decode mode (seqlen <= 32) do DRAM sharded matmuls # These use HiFi2; this drops 1 bit of the activations but would be FLOP-bound on 12 cores with HiFi4 pc_1 = self.args.get_mlp_ff1_3_prg_config(mode, seq_len, self.prefetcher) pc_2 = self.args.get_mlp_ff2_prg_config(mode, seq_len, self.prefetcher) pc_3 = self.args.get_mlp_ff1_3_prg_config(mode, seq_len, self.prefetcher) use_tg_decode_no_prefetch = TG and mode == Mode.DECODE and self.prefetcher is None if use_tg_decode_no_prefetch: x = ttnn.to_memory_config(x, ttnn.DRAM_MEMORY_CONFIG) w1_out = ttnn.linear( x, self.w1, dtype=ttnn.bfloat8_b if TG else activation_dtype or ttnn.bfloat16, core_grid=None, # FIXME: validate on TG ttnn.CoreGrid(y=8, x=8) if not pc_1 else None, compute_kernel_config=li_ff1_3_compute_kernel_cfg, program_config=None if use_tg_decode_no_prefetch else pc_1, memory_config=( ttnn.DRAM_MEMORY_CONFIG if use_tg_decode_no_prefetch else self.args.get_mlp_ff1_3_mem_config(mode, self.prefetcher) ), global_cb=self.prefetcher.global_cb if self.prefetcher is not None and mode == Mode.DECODE else None, sub_device_id=self.prefetcher.worker_sub_device_id if self.prefetcher is not None and mode == Mode.DECODE else None, ) w3_out = ttnn.linear( x, self.w3, dtype=ttnn.bfloat8_b if TG else activation_dtype or ttnn.bfloat16, core_grid=None, # FIXME: validate on TG ttnn.CoreGrid(y=8, x=8) if not pc_3 else None, compute_kernel_config=li_ff1_3_compute_kernel_cfg, program_config=None if use_tg_decode_no_prefetch else pc_3, memory_config=( ttnn.DRAM_MEMORY_CONFIG if use_tg_decode_no_prefetch else self.args.get_mlp_ff1_3_mem_config(mode, self.prefetcher) ), global_cb=self.prefetcher.global_cb if self.prefetcher is not None and mode == Mode.DECODE else None, sub_device_id=self.prefetcher.worker_sub_device_id if self.prefetcher is not None and mode == Mode.DECODE else None, ) ttnn.deallocate(x) if TG: # if mode == "decode" and self.dim!=8192: # w1_out = ttnn.to_memory_config(w1_out, ttnn.DRAM_MEMORY_CONFIG) # w3_out = ttnn.to_memory_config(w3_out, ttnn.DRAM_MEMORY_CONFIG) if self.dim == 8192 or mode == Mode.PREFILL: input_mem_cfg = w1_out.memory_config() cluster_axis = 1 w1_out = ttnn.experimental.reduce_scatter_minimal_async( w1_out, persistent_output_buffers=None, dim=3, multi_device_global_semaphore=self.tt_ccl.get_and_cycle_rs_semaphore_handles(cluster_axis), barrier_semaphore=self.tt_ccl.get_and_cycle_barrier_semaphore_handle(cluster_axis), num_links=self.tt_ccl.get_num_links(cluster_axis), cluster_axis=cluster_axis, memory_config=self.model_config["FF1_OUT_REDUCE_SCATTER_MEMCFG"] if mode == Mode.DECODE else None, intermediate_memory_config=ttnn.DRAM_MEMORY_CONFIG, topology=self.args.ccl_topology(), chunks_per_sync=10, num_workers_per_link=2, num_buffers_per_channel=2, ) w3_out = ttnn.experimental.reduce_scatter_minimal_async( w3_out, persistent_output_buffers=None, dim=3, multi_device_global_semaphore=self.tt_ccl.get_and_cycle_rs_semaphore_handles(cluster_axis), barrier_semaphore=self.tt_ccl.get_and_cycle_barrier_semaphore_handle(cluster_axis), num_links=1, cluster_axis=cluster_axis, memory_config=self.model_config["FF1_OUT_REDUCE_SCATTER_MEMCFG"] if mode == Mode.DECODE else None, intermediate_memory_config=ttnn.DRAM_MEMORY_CONFIG, topology=self.args.ccl_topology(), chunks_per_sync=10, num_workers_per_link=2, num_buffers_per_channel=2, ) else: # NOTE: In MLP All-reduce hard codes to 2 links, so we do not get the dynamic link count from the CCL class # to avoid any performance regressions. w1_out = tt_all_reduce( w1_out, self.mesh_device, self.tt_ccl, cluster_axis=1, num_all_gather_links=2, sharded=mode == Mode.DECODE and not use_tg_decode_no_prefetch, topology=self.args.ccl_topology(), memory_config=( ttnn.DRAM_MEMORY_CONFIG if use_tg_decode_no_prefetch else self.model_config["FF1_OUT_GATHERED_MEMCFG"] if mode == Mode.DECODE else None ), ) w3_out = tt_all_reduce( w3_out, self.mesh_device, self.tt_ccl, cluster_axis=1, num_all_gather_links=2, sharded=mode == Mode.DECODE and not use_tg_decode_no_prefetch, topology=self.args.ccl_topology(), memory_config=( ttnn.DRAM_MEMORY_CONFIG if use_tg_decode_no_prefetch else self.model_config["FF1_OUT_GATHERED_MEMCFG"] if mode == Mode.DECODE else None ), ) w2_in = ttnn.mul( w1_out, w3_out, input_tensor_a_activations=[self.activation_type], dtype=activation_dtype or ttnn.bfloat8_b, memory_config=w1_out.memory_config(), ) if mode == Mode.DECODE and not TG and self.prefetcher is None: # w2 may use a different core grid, this is a no-op if they already match w2_in = ttnn.to_memory_config(w2_in, self.args.get_mlp_binary_mult_mem_config(mode)) ttnn.deallocate(w3_out) ttnn.deallocate(w1_out) if TG and (self.dim == 8192 or mode == Mode.PREFILL): cluster_axis = 1 w2_in = ttnn.experimental.all_gather_async( w2_in, persistent_output_buffer=None, dim=3, multi_device_global_semaphore=self.tt_ccl.get_and_cycle_ag_semaphore_handles(cluster_axis), num_links=2, cluster_axis=1, topology=self.args.ccl_topology(), memory_config=input_mem_cfg, barrier_semaphore=self.tt_ccl.get_and_cycle_barrier_semaphore_handle(cluster_axis), chunks_per_sync=10, num_workers_per_link=2, num_buffers_per_channel=2, ) if mode == Mode.DECODE: w2_in = ttnn.to_memory_config(w2_in, ttnn.L1_MEMORY_CONFIG) li_ff2_compute_kernel_cfg = self.decoders_optimizations.get_math_fidelity( decoder_id=layer_num, op=OpGroup.LI_FF2, configuration=self.args ) w2_output_dtype = activation_dtype if w2_output_dtype is None: w2_output_dtype = ttnn.bfloat16 if mode != Mode.DECODE and self.args.use_minimal_prefill_matmul(seq_len): is_qwen3_32b_t3k = self.args.base_model_name == "Qwen3-32B" and self.args.device_name == "T3K" if not is_qwen3_32b_t3k: # None makes minimal_matmul inherit w2_in.dtype. Only the # validated Qwen3-32B/T3K path overrides that default. w2_output_dtype = None w2_out = ttnn.experimental.minimal_matmul( w2_in, self.w2, dtype=w2_output_dtype, compute_kernel_config=li_ff2_compute_kernel_cfg, config=pc_2, ) else: if TG: w2_output_dtype = self.args.ccl_dtype w2_out = ttnn.linear( w2_in, self.w2, compute_kernel_config=li_ff2_compute_kernel_cfg, dtype=w2_output_dtype, program_config=None if use_tg_decode_no_prefetch else pc_2, memory_config=( ttnn.DRAM_MEMORY_CONFIG if use_tg_decode_no_prefetch else self.args.get_mlp_ff2_mem_config(mode, self.prefetcher) ), core_grid=None, # FIXME: validate on TG ttnn.CoreGrid(y=8, x=8) if not pc_2 else None, global_cb=self.prefetcher.global_cb if self.prefetcher is not None and mode == Mode.DECODE else None, sub_device_id=self.prefetcher.worker_sub_device_id if self.prefetcher is not None and mode == Mode.DECODE else None, ) ttnn.deallocate(w2_in) w2_out_reduced = tt_all_reduce( w2_out, self.mesh_device, self.tt_ccl, cluster_axis=0, dim=0 if (TG and self.dim < 8192) else 3, sharded=mode == Mode.DECODE and not use_tg_decode_no_prefetch, memory_config=( ttnn.DRAM_MEMORY_CONFIG if use_tg_decode_no_prefetch else self.args.get_mlp_ff2_all_reduce_mem_config(mode, w2_out) ), rs_memory_config=self.model_config["MLP_RS_CONFIG"]["rs_memory_config"] if mode == Mode.DECODE else ttnn.DRAM_MEMORY_CONFIG, dtype=self.args.ccl_dtype, use_composite=True if self.dim == 8192 else False, topology=self.args.ccl_topology(), chunks_per_sync=self.model_config["MLP_RS_CONFIG"]["chunks_per_sync"] if mode == Mode.DECODE else 10, num_workers_per_link=self.model_config["MLP_RS_CONFIG"]["num_workers_per_link"] if mode == Mode.DECODE else 2, subdevice_id=self.prefetcher.worker_sub_device_id if mode == Mode.DECODE and self.prefetcher is not None else None, ) # Ensure dim 0 and 1 are 1 original_shape = w2_out_reduced.shape w2_out_reduced = ttnn.reshape( w2_out_reduced, (1, 1, original_shape[-4] * original_shape[-3] * original_shape[-2], original_shape[-1]) ) if mode == Mode.DECODE: w2_out_reduced = ttnn.to_memory_config( w2_out_reduced, self.args.get_mlp_output_mem_config(mode, self.prefetcher), ) return w2_out_reduced