File size: 8,487 Bytes
b025706 | 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 | # SPDX-FileCopyrightText: © 2023 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
from loguru import logger
import ttnn
from models.common.lightweightmodule import LightweightModule
from models.tt_transformers.tt.ccl import tt_distributed_rmsnorm, tt_sharded_distributed_rmsnorm
from models.tt_transformers.tt.common import Mode
def galaxy_distributed_norm_core_grid(dim: int) -> tuple[int, int]:
"""Choose the Galaxy decode RMSNorm grid as (y, x).
Returns the legacy ``(min(4, dim // 4 // 32 // 8), 8)`` grid for every dim it
already handled. Only dims that grid cannot tile-align get an override, and a
dim with no tile-aligned override keeps the legacy grid with a warning rather
than raising -- ``gather_in_mem_cfg`` / ``ln_prg_cfg`` are consumed only on the
sharded decode path, so raising here would break prefill-only Galaxy runs that
never read them (e.g. every Gemma variant).
"""
hidden_size_per_device = dim // 4
legacy_core_grid = (min(4, hidden_size_per_device // 32 // 8), 8)
if hidden_size_per_device == 1280:
# Qwen's silicon-validated Galaxy layout: 10 cores with four tiles per shard.
# The legacy grid would be (4, 8) = 32 cores, which cannot tile-align 1280.
core_grid = (5, 2)
else:
core_grid = legacy_core_grid
num_cores = core_grid[0] * core_grid[1]
if num_cores == 0 or hidden_size_per_device % (num_cores * 32) != 0:
logger.warning(
f"Galaxy distributed norm hidden size {hidden_size_per_device} is not tile-shardable "
f"across grid {core_grid}; keeping the legacy grid {legacy_core_grid}. The sharded "
"decode norm config will be misaligned if it is used."
)
return legacy_core_grid
return core_grid
class DistributedNorm(LightweightModule):
"""Wraps a norm with the TP gather it implies.
``norm=None`` turns this into a gather-only identity: the input is gathered
(or left fractured + gathered after, in distributed-norm mode) exactly as it
would be around a real norm, but no normalization is applied. Post-norm
decoders (EXAONE-4.x) use this where pre-norm models have input_layernorm /
pre_feedforward_layernorm, since those slots double as the fractured->
replicated gather points. Note an RMSNorm with all-ones gamma is NOT an
identity (it still divides by RMS), hence this explicit mode.
"""
def __init__(self, norm, args, tt_ccl, prefetcher=None, TG=False, ag_config_key=None, enable_all_gather=True):
if norm is None and TG:
raise NotImplementedError("Gather-only DistributedNorm (norm=None) is not supported on Galaxy (TG)")
self.norm = norm
self.args = args
self.tt_ccl = tt_ccl
self.prefetcher = prefetcher
self.ag_config_key = ag_config_key
# Flag to control whether all_gather is performed after distributed norm (can be disabled when output should remain sharded)
self.enable_all_gather = enable_all_gather
if TG:
core_grid_ln = galaxy_distributed_norm_core_grid(args.dim)
num_cores_ln = core_grid_ln[0] * core_grid_ln[1]
hidden_size_per_device_distributed_ln = args.dim // 4
self.gather_in_mem_cfg = ttnn.create_sharded_memory_config(
shape=(1, 1, 32, hidden_size_per_device_distributed_ln),
core_grid=ttnn.CoreGrid(y=core_grid_ln[0], x=core_grid_ln[1]),
strategy=ttnn.ShardStrategy.WIDTH,
)
self.ln_prg_cfg = ttnn.LayerNormShardedMultiCoreProgramConfig(
compute_with_storage_grid_size=(core_grid_ln[1], core_grid_ln[0]),
subblock_w=(hidden_size_per_device_distributed_ln // num_cores_ln) // 32,
block_h=1,
block_w=(hidden_size_per_device_distributed_ln // num_cores_ln) // 32,
inplace=False,
)
self.ln_sharded_stats_memcfg = ttnn.create_sharded_memory_config(
shape=[1, 1, 32, 32 * 4],
core_grid=ttnn.CoreGrid(y=1, x=1),
strategy=ttnn.ShardStrategy.WIDTH,
)
self.ln_cfg = ttnn.WormholeComputeKernelConfig(
math_fidelity=ttnn.MathFidelity.HiFi2,
math_approx_mode=False,
fp32_dest_acc_en=False,
packer_l1_acc=False,
)
self.TG = TG
def update(self, *, weight: ttnn.Tensor) -> None:
"""Pass-through to the wrapped ``RMSNorm.update`` (``DistributedNorm``
owns no weights of its own). Same HF-format contract: ``(1, 1, 1, dim)``,
TILE, bf16, DRAM-interleaved, replicated.
"""
if self.norm is None:
raise ValueError("Gather-only DistributedNorm (norm=None) owns no weights to update")
self.norm.update(weight=weight)
def forward(self, x, mode: Mode, norm_config=None):
"""Apply a norm, possibly gathering inputs if required."""
sharded_output_config = norm_config.get("sharded_output_config") if norm_config else None
if self.TG:
if mode == Mode.DECODE:
return tt_sharded_distributed_rmsnorm(
x,
epsilon=self.norm.eps,
gamma=self.norm.weight_distributed,
mesh_device=self.args.mesh_device,
tt_ccl=self.tt_ccl,
ln_sharded_input_memcfg=self.gather_in_mem_cfg,
ln_sharded_progcfg=self.ln_prg_cfg,
ln_sharded_stats_memcfg=self.ln_sharded_stats_memcfg,
)
else:
return tt_distributed_rmsnorm(
x,
epsilon=self.norm.eps,
gamma=self.norm.weight_distributed,
mesh_device=self.args.mesh_device,
tt_ccl=self.tt_ccl,
compute_kernel_config=self.ln_cfg,
)
input_mem_cfg = sharded_output_config if mode == Mode.DECODE else ttnn.DRAM_MEMORY_CONFIG
# Distributed norm already performs a gather
if self.args.is_multichip and not self.args.is_distributed_norm(mode):
x = ttnn.experimental.all_gather_async(
x,
persistent_output_buffer=None,
dim=3,
multi_device_global_semaphore=self.tt_ccl.get_and_cycle_ag_semaphore_handles(),
num_links=self.args.model_config[self.ag_config_key]["num_links"]
if self.ag_config_key and mode == "decode"
else self.tt_ccl.get_num_links(1),
topology=self.args.ccl_topology(),
memory_config=input_mem_cfg,
barrier_semaphore=self.tt_ccl.get_and_cycle_barrier_semaphore_handle(),
chunks_per_sync=self.args.model_config[self.ag_config_key]["chunks_per_sync"]
if self.ag_config_key and mode == "decode"
else 10,
num_workers_per_link=self.args.model_config[self.ag_config_key]["num_workers_per_link"]
if self.ag_config_key and mode == "decode"
else 2,
num_buffers_per_channel=2,
subdevice_id=self.prefetcher.worker_sub_device_id if self.prefetcher is not None else None,
)
else:
x = ttnn.to_memory_config(x, input_mem_cfg)
if self.norm is not None:
x = self.norm(
x,
mode=mode,
in_sharded=(mode == Mode.DECODE),
out_sharded=(mode == Mode.DECODE),
norm_config=norm_config,
)
# Distributed norm requires a gather
if self.args.is_distributed_norm(mode) and self.enable_all_gather:
x = ttnn.experimental.all_gather_async(
x,
persistent_output_buffer=None,
dim=3,
multi_device_global_semaphore=self.tt_ccl.get_and_cycle_ag_semaphore_handles(),
num_links=self.tt_ccl.get_num_links(1),
topology=self.args.ccl_topology(),
memory_config=x.memory_config(),
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,
)
return x
|