File size: 10,682 Bytes
be3ecc8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 | # SPDX-FileCopyrightText: © 2023 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
import ttnn
from models.common.lightweightmodule import LightweightModule
from models.common.utility_functions import copy_to_buffer
from models.tt_transformers.tt.common import Mode
TILE = 32
SHARD_HEIGHT = TILE # Current ttnn.rms_norm implementation requires shard height to be a single tile
class RMSNorm(LightweightModule):
"""
RMSNorm supporting replication over a MeshDevice and sharding within devices.
This class implements a Root Mean Square Normalization (RMSNorm) that can be
distributed across multiple devices and cores. If the `device` parameter is a
MeshDevice, the weights and computations are replicated across all devices in
the mesh. Expects an interleaved input tensor, can optionally output a sharded tensor.
Args:
device: The device or MeshDevice on which to perform the computations.
state_dict: The state dictionary containing the model parameters.
dim: Input dimension (e.g. model hidden dimension size).
layer_num: The layer number to determine the weight key in the state dictionary.
weight_key: The key for retrieving the weight from the state dictionary.
weight_cache_path: Optional path for caching the tilized weights.
weight_memory_config: Configuration for the weight memory, default is DRAM_MEMORY_CONFIG.
weight_dtype: The data type for the tensors, bfp8_b hits >0.999 PCC in the models we tested.
model_config: Optional configuration dictionary for the model.
eps (float): Small value to avoid division by zero in normalization, default is 1e-05.
If model_config is provided, it must specify SHARDED_NORM_INPUT_MEMCFG, SHARDED_NORM_PRGM_CFG
and SHARDED_NORM_OUTPUT_MEMCFG. If not provided, default configurations will be generated.
"""
def __init__(
self,
device,
dim,
state_dict,
weight_key,
layer_num=None,
state_dict_prefix=None,
weight_cache_path=None,
weight_memory_config=ttnn.DRAM_MEMORY_CONFIG,
weight_dtype=ttnn.bfloat16,
is_distributed=None,
eps: float = 1e-05,
add_unit_offset=False,
sharded_program_config=None,
sharded_output_config=None,
output_mem_config=None,
ccl_topology=ttnn.Topology.Ring,
tt_ccl=None,
fp32_dest_acc_en=True,
):
super().__init__()
self.device = device
self.eps = eps
self.is_distributed = is_distributed
self.ccl_topology = ccl_topology
self.tt_ccl = tt_ccl
self.add_unit_offset = add_unit_offset
if state_dict_prefix:
weight_name = f"{state_dict_prefix}{weight_key}.weight"
else:
if layer_num is None:
weight_name = f"{weight_key}.weight"
else:
weight_name = f"layers.{layer_num}.{weight_key}.weight"
torch_weight = (
state_dict[weight_name].unsqueeze(0).view(1, 1, dim).reshape([1, 1, dim // SHARD_HEIGHT, SHARD_HEIGHT])
)
# Add offset before caching
if add_unit_offset:
torch_weight = torch_weight + 1.0
# Compatibility with models that don't use mesh devices (e.g. single-chip Mistral-7b)
is_mesh_device = device.__class__.__name__ == "MeshDevice"
self.weight = ttnn.as_tensor(
torch_weight,
device=device,
dtype=weight_dtype,
layout=ttnn.ROW_MAJOR_LAYOUT,
memory_config=weight_memory_config,
cache_file_name=None if weight_cache_path is None else weight_cache_path / weight_name,
mesh_mapper=ttnn.ReplicateTensorToMesh(device) if is_mesh_device else None,
)
if self.is_distributed:
self.weight_distributed = ttnn.as_tensor(
torch_weight,
device=device,
dtype=weight_dtype,
layout=ttnn.ROW_MAJOR_LAYOUT,
memory_config=weight_memory_config,
cache_file_name=(
None if weight_cache_path is None else weight_cache_path / (weight_name + "_distributed")
),
mesh_mapper=(
ttnn.ShardTensor2dMesh(device, dims=(None, 2), mesh_shape=list(device.shape))
if is_mesh_device
else None
),
)
self.sharded_output_config = sharded_output_config
self.sharded_program_config = sharded_program_config
self.output_mem_config = output_mem_config
self.compute_kernel_config_hifi2 = ttnn.WormholeComputeKernelConfig(
math_fidelity=ttnn.MathFidelity.HiFi2,
math_approx_mode=False,
fp32_dest_acc_en=fp32_dest_acc_en,
packer_l1_acc=True,
)
def update(self, *, weight: ttnn.Tensor) -> None:
"""In-place replace the RMSNorm gamma via ``ttnn.copy``.
HF-format input: ``weight`` is HF ``...norm.weight``, shape
``(1, 1, 1, dim)``, bf16, TILE, DRAM-interleaved, replicated.
``copy_to_buffer`` reshapes to the storage shape
``(1, 1, dim // SHARD_HEIGHT, SHARD_HEIGHT)`` and TILE -> ROW_MAJOR to
match ``self.weight``. ``add_unit_offset`` is not supported (see
assert): the caller must ship a gamma that already includes the +1.
When ``self.weight_distributed`` (the column-sharded mirror) exists it's
kept in sync on device: project ``self.weight`` into the sharded layout
via ``ttnn.mesh_partition`` (the inverse of the constructor's
``ShardTensor2dMesh(dims=(None, 2))``, hence ``dim=2, cluster_axis=1``)
and ``ttnn.copy`` into it. Both buffers keep their address, so captured
traces and the prefetcher's recorded addresses stay valid.
"""
assert not self.add_unit_offset, "RMSNorm.update does not support add_unit_offset=True"
copy_to_buffer(weight, self.weight, self.weight.dtype)
if getattr(self, "weight_distributed", None) is not None:
partitioned = ttnn.mesh_partition(
self.weight,
memory_config=self.weight_distributed.memory_config(),
dim=2,
cluster_axis=1,
)
copy_to_buffer(partitioned, self.weight_distributed, self.weight_distributed.dtype)
def forward(
self,
x: ttnn.Tensor,
mode: Mode | str,
in_sharded=False,
out_sharded=False,
norm_config=None,
) -> ttnn.Tensor:
if isinstance(mode, str):
try:
mode = Mode(mode)
except ValueError:
raise ValueError(f"Invalid mode: {mode}")
elif not isinstance(mode, Mode):
raise ValueError(f"Invalid mode: {mode}")
sharded_program_config = norm_config.get("sharded_program_config") if norm_config else None
sharded_output_config = norm_config.get("sharded_output_config") if norm_config else None
output_mem_config = norm_config.get("output_mem_config") if norm_config else None
# Optional L1 placement for the distributed 3-op outputs (pre/gather/post); None -> DRAM default.
distributed_out_mc = norm_config.get("distributed_output_mem_config") if norm_config else None
# If input is sharded do sharded RMSNorm and optionally return sharded output
program_config = sharded_program_config if in_sharded else None
memory_config = sharded_output_config if out_sharded else None
distributed = self.is_distributed and self.is_distributed(mode)
weight = self.weight_distributed if distributed else self.weight
if in_sharded:
assert not distributed, "Distributed RMSNorm does not support sharded inputs"
else:
assert not out_sharded, "Non-sharded version of RMSNorm cannot output a sharded tensor"
if distributed:
x = self._distributed_rmsnorm(
x,
epsilon=self.eps,
weight=weight,
compute_kernel_config=self.compute_kernel_config_hifi2,
output_memory_config=distributed_out_mc,
)
else:
x = ttnn.rms_norm(
x,
epsilon=self.eps,
weight=weight,
program_config=program_config,
memory_config=memory_config,
compute_kernel_config=self.compute_kernel_config_hifi2,
)
if in_sharded and not out_sharded:
return ttnn.sharded_to_interleaved(x)
else:
if output_mem_config is not None:
x = ttnn.to_memory_config(x, output_mem_config)
return x
def _distributed_rmsnorm(
self,
inp,
epsilon=None,
weight=None,
program_config=None,
memory_config=None,
compute_kernel_config=None,
output_memory_config=None,
):
assert program_config is None, "Distributed RMSNorm does not support sharded inputs"
assert memory_config is None, "Distributed RMSNorm does not support sharded outputs"
assert self.tt_ccl is not None, "Distributed RMSNorm requires tt_ccl"
# Interleaved output placement for the 3 ops; default DRAM (matches the prior hardcoded behavior).
mc = output_memory_config if output_memory_config is not None else ttnn.DRAM_MEMORY_CONFIG
# Run distributed rmsnorm part 1
tt_stats = ttnn.rms_norm_pre_all_gather(
inp, compute_kernel_config=compute_kernel_config, dtype=ttnn.bfloat16, memory_config=mc
)
# AllGather stats
tt_stats = ttnn.experimental.all_gather_async(
tt_stats,
persistent_output_buffer=None,
dim=3,
multi_device_global_semaphore=self.tt_ccl.get_and_cycle_ag_semaphore_handles(),
num_links=1,
topology=self.ccl_topology,
memory_config=mc,
barrier_semaphore=self.tt_ccl.get_and_cycle_barrier_semaphore_handle(),
chunks_per_sync=10,
num_workers_per_link=2,
num_buffers_per_channel=2,
)
# Run distributed rmsnorm part 2
tt_out = ttnn.rms_norm_post_all_gather(
inp,
tt_stats,
epsilon=epsilon,
weight=weight,
compute_kernel_config=compute_kernel_config,
memory_config=mc,
)
tt_stats.deallocate(True)
return tt_out
|