clef / code /models /tt_transformers /tt /distributed_norm.py
tt-hous's picture
Add files using upload-large-folder tool
b025706 verified
Raw History Blame Contribute Delete
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