# 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