tt-hous's picture
Add files using upload-large-folder tool
b025706 verified
Raw History Blame Contribute Delete
20 kB
# 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