Download code/models/tt_transformers/tt/distributed_norm.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 8.49 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/distributed_norm.py
- Command line
-
hf download hf://tt-hous/clef/code/models/tt_transformers/tt/distributed_norm.py
-
curl -L -o distributed_norm.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/distributed_norm.py
8.49 kB
| # 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 | |