stisiTT's picture
Add files using upload-large-folder tool
9aa90e0 verified
Raw History Blame Contribute Delete
35.8 kB
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import math
from typing import TYPE_CHECKING
import torch
from loguru import logger
import ttnn
if TYPE_CHECKING:
from collections.abc import Callable, Sequence
from types import EllipsisType
def prepare_for_fused_swiglu(
tensor: torch.Tensor, ndev: int, tile_width: int = 32, *, gate_is_first: bool = False
) -> torch.Tensor:
"""Reorder a packed [.., 2N] SwiGLU weight/bias for the fused matmul kernel.
The fused kernel emits ``silu(even_tile) * odd_tile`` per pair, so gate (the silu'd half)
must occupy the even tile slot and up (the multiplicand) the odd tile slot within each
device block.
``gate_is_first`` specifies the input column ordering:
- ``gate_is_first=False`` (default): input is ``[up (N) | gate (N)]``, i.e.
``up, gate = chunk(., 2, -1); up * silu(gate)`` — torch/HuggingFace convention.
- ``gate_is_first=True``: input is ``[gate (N) | up (N)]``, i.e.
``gate, up = chunk(., 2, -1); silu(gate) * up``.
In both cases the output is laid out as ``ndev`` contiguous device blocks, each containing
interleaved gate/up tile pairs ``[gate_t0, up_t0, gate_t1, up_t1, ...]`` for that device's shard.
Splitting the output evenly across ``ndev`` devices along the last dimension gives each device
a correctly interleaved local block.
"""
rows = tensor.shape[0]
two_N = tensor.shape[-1]
N = two_N // 2
assert (
N % (ndev * tile_width) == 0
), f"SwiGLU half-width N={N} must be divisible by ndev*tile_width={ndev * tile_width}"
tiles_per_dev = N // ndev // tile_width
# [rows, half=2, d=ndev, t=tiles_per_dev, c=tile_width] -> [rows, d, t, half, c]
t = tensor.reshape(rows, 2, ndev, tiles_per_dev, tile_width)
t = t.permute(0, 2, 3, 1, 4)
if not gate_is_first:
t = t.flip(3)
return t.reshape(rows, two_N)
def typed_tensor(
x: torch.Tensor,
dtype: ttnn.DataType,
device: ttnn.Device | None = None,
mesh_axis=None,
shard_dim=None,
layout=ttnn.TILE_LAYOUT,
on_host=False,
) -> ttnn.Tensor:
"""
Replicates or shards a tensor based on the mesh_axis and shard_dim
"""
assert (mesh_axis is None) == (shard_dim is None)
mesh_mapper = None
if mesh_axis is not None:
mapper_dims = [None, None]
mapper_dims[mesh_axis] = shard_dim
mesh_mapper = ttnn.ShardTensor2dMesh(device, mesh_shape=tuple(device.shape), dims=mapper_dims)
return ttnn.from_torch(
x,
layout=layout,
dtype=dtype,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
device=None if on_host else device,
mesh_mapper=mesh_mapper,
)
def bf16_tensor(
x: torch.Tensor,
device: ttnn.Device | None = None,
mesh_axis=None,
shard_dim=None,
layout=ttnn.TILE_LAYOUT,
on_host=False,
) -> ttnn.Tensor:
return typed_tensor(
x,
ttnn.bfloat16,
device,
mesh_axis,
shard_dim,
layout,
on_host=on_host,
)
def float32_tensor(
x: torch.Tensor, device: ttnn.Device | None = None, mesh_axis=None, shard_dim=None, layout=ttnn.TILE_LAYOUT
) -> ttnn.Tensor:
return typed_tensor(x, ttnn.float32, device, mesh_axis, shard_dim, layout)
def bf16_tensor_host(
x: torch.Tensor, device: ttnn.Device | None = None, mesh_axis=None, shard_dim=None, layout=ttnn.TILE_LAYOUT
) -> ttnn.Tensor:
assert (mesh_axis is None) == (shard_dim is None)
mesh_mapper = None
if mesh_axis is not None:
mapper_dims = [None, None]
mapper_dims[mesh_axis] = shard_dim
mesh_mapper = ttnn.ShardTensor2dMesh(device, mesh_shape=tuple(device.shape), dims=mapper_dims)
return ttnn.from_torch(
x,
layout=layout,
dtype=ttnn.bfloat16,
mesh_mapper=mesh_mapper,
)
def typed_tensor_2dshard(
x: torch.Tensor,
device: ttnn.Device,
shard_mapping: dict[int, int],
layout=ttnn.TILE_LAYOUT,
dtype=ttnn.bfloat16,
on_host=False,
) -> ttnn.Tensor:
assert len(shard_mapping) == 2
assert all(0 <= k <= 1 and 0 <= v < len(x.shape) for k, v in shard_mapping.items())
mapper_dims = [None, None]
for k, v in shard_mapping.items():
mapper_dims[k] = v
mesh_mapper = ttnn.ShardTensor2dMesh(device, mesh_shape=tuple(device.shape), dims=mapper_dims)
return ttnn.from_torch(
x,
layout=layout,
dtype=dtype,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
device=None if on_host else device,
mesh_mapper=mesh_mapper,
)
def bf16_tensor_2dshard(
x: torch.Tensor,
device: ttnn.Device,
shard_mapping: dict[int, int],
layout=ttnn.TILE_LAYOUT,
on_host: bool = False,
) -> ttnn.Tensor:
return typed_tensor_2dshard(x, device, shard_mapping, layout, dtype=ttnn.bfloat16, on_host=on_host)
def from_torch(
x: torch.Tensor,
/,
*,
device: ttnn.MeshDevice | None = None,
layout: ttnn.Layout = ttnn.Layout.TILE,
dtype: ttnn.DataType = ttnn.bfloat16,
memory_config: ttnn.MemoryConfig = ttnn.DRAM_MEMORY_CONFIG,
pad_value: float | None = None,
mesh_axes: Sequence[int | None | EllipsisType] | None = None,
on_host: bool = False,
) -> ttnn.Tensor:
"""Convert a torch.Tensor to a ttnn.Tensor with convenient mesh distribution."""
if mesh_axes is not None:
if device is None:
msg = "device must be specified if mesh_axes is given"
raise ValueError(msg)
mesh_rank = len(list(device.shape))
mesh_axes = canonicalize_tensor_mesh_axes(mesh_axes, tensor_rank=len(x.shape), mesh_rank=mesh_rank)
placements = _invert_placements(mesh_axes, output_rank=mesh_rank)
placements = [ttnn.PlacementShard(p) if p is not None else ttnn.PlacementReplicate() for p in placements]
mesh_mapper = ttnn.create_mesh_mapper(device, ttnn.MeshMapperConfig(placements))
else:
mesh_mapper = None
return ttnn.from_torch(
x,
layout=layout,
dtype=dtype,
memory_config=memory_config,
device=None if on_host else device,
mesh_mapper=mesh_mapper,
pad_value=pad_value,
)
def from_torch_to_devices(
x: torch.Tensor,
*,
devices: Sequence[ttnn.MeshDevice],
mesh_axes: Sequence[int | None] | None = None,
on_host: bool = False,
) -> list[ttnn.Tensor]:
"""Replicate a torch tensor across submesh devices, returning one ttnn.Tensor per device."""
return [from_torch(x, device=d, mesh_axes=mesh_axes, on_host=on_host) for d in devices]
def to_torch(
x: ttnn.Tensor,
/,
*,
mesh_axes: Sequence[int | None | EllipsisType] | None = None,
composer_device: ttnn.MeshDevice | None = None,
) -> torch.Tensor:
"""Converts a ttnn.Tensor to a torch.Tensor.
Strips away redundant data returned by calling ttnn.to_torch on a replicated tensor. If the
tensor is distributed and not on device, composer_device must be provided.
"""
if x.tensor_topology().distribution_shape().mesh_size() == 1:
return ttnn.to_torch(x)
if mesh_axes is None:
mesh_axes = (None,) * len(x.shape)
composer_device = composer_device or x.device()
if composer_device is None:
msg = "composer_device must be specified for distributed host tensors"
raise ValueError(msg)
mesh_rank = len(list(composer_device.shape))
mesh_axes = canonicalize_tensor_mesh_axes(mesh_axes, tensor_rank=len(x.shape), mesh_rank=mesh_rank)
replicated_mesh_axes = list(set(range(mesh_rank)) - {axis for axis in mesh_axes if axis is not None})
mesh_axes = replicated_mesh_axes + list(mesh_axes)
placements = _invert_placements(mesh_axes, output_rank=mesh_rank)
assert all(p is not None for p in placements)
mesh_composer = ttnn.create_mesh_composer(composer_device, ttnn.MeshComposerConfig(placements))
x = x.reshape([1] * len(replicated_mesh_axes) + list(x.shape))
return ttnn.to_torch(x, mesh_composer=mesh_composer)[(0,) * len(replicated_mesh_axes)]
def verify_tensor_mesh_axes(mesh_axes: Sequence[int | None], /, *, tensor_rank: int, mesh_rank: int) -> None:
"""Validates tensor mesh axes specification."""
if len(mesh_axes) != tensor_rank:
msg = f"mesh axis list {tuple(mesh_axes)} should have length {tensor_rank}"
raise ValueError(msg)
for axis in mesh_axes:
if axis is not None and (axis < 0 or axis >= mesh_rank):
msg = f"all mesh axes in mesh axis list {tuple(mesh_axes)} should be positive and smaller than {mesh_rank}"
raise ValueError(msg)
non_none_values = [p for p in mesh_axes if p is not None]
if len(non_none_values) != len(set(non_none_values)):
msg = f"mesh axis list {tuple(mesh_axes)} contains duplicate mesh axis assignments"
raise ValueError(msg)
def canonicalize_tensor_mesh_axes(
mesh_axes: Sequence[int | None | EllipsisType], /, *, tensor_rank: int, mesh_rank: int
) -> tuple[int | None, ...]:
"""Canonicalizes mesh axes specification by expanding Ellipsis and validating."""
mesh_axes = list(mesh_axes)
if Ellipsis in mesh_axes:
if mesh_axes.count(Ellipsis) > 1:
msg = "mesh_axes can contain at most one Ellipsis"
raise ValueError(msg)
ellipsis_index = mesh_axes.index(Ellipsis)
mesh_axes[ellipsis_index : ellipsis_index + 1] = [None] * (tensor_rank - len(mesh_axes) + 1)
verify_tensor_mesh_axes(mesh_axes, tensor_rank=tensor_rank, mesh_rank=mesh_rank)
return tuple(mesh_axes)
def _invert_placements(placements: Sequence[int | None], *, output_rank: int) -> tuple[int | None, ...]:
out = [None] * output_rank
for i, p in enumerate(placements):
if p is not None:
out[p] = i
return tuple(out)
def pad_single(
x: ttnn.Tensor,
/,
*,
dim: int,
front: int = 0,
back: int = 0,
value: float = 0.0,
) -> ttnn.Tensor:
"""Pad a tensor along a single dimension, working around the `ttnn.pad` dimension restriction."""
shape = list(x.shape)
rank = len(shape)
if dim < 0:
dim += rank
if dim < 0 or dim >= rank:
msg = f"padding dimension {dim} is out of bounds for tensor with rank {rank}"
raise ValueError(msg)
# From ttnn: "ttnn::pad only supports padding on the lowest 3 dimensions for tensors with rank > 4."
# With the way we count, the last three dimensions are supported.
if rank <= 4 or dim >= rank - 3:
padding = [(0, 0)] * rank
padding[dim] = (front, back)
return ttnn.pad(x, padding, value=value)
# The reshapes should be fast, since they preserve the last two dimensions.
before = math.prod(shape[:dim])
x = ttnn.reshape(x, [before, -1, *shape[-2:]]) # reshape to 4d
v = math.prod(shape[dim + 1 : -2])
x = ttnn.pad(x, [(0, 0), (front * v, back * v), (0, 0), (0, 0)], value=value)
shape[dim] += front + back
return ttnn.reshape(x, shape)
def local_device_to_torch(tt_tensor: ttnn.Tensor) -> torch.Tensor:
"""Convert a ttnn device tensor to a torch tensor by reading from the local device.
In a distributed environment, iterates over the mesh coordinates to find the
tensor shard that belongs to the local device before calling ``ttnn.to_torch``.
"""
mesh_device = tt_tensor.device()
view = mesh_device.get_view() if ttnn.using_distributed_env() else None
coords = list(tt_tensor.tensor_topology().mesh_coords())
device_tensors = ttnn.get_device_tensors(tt_tensor)
torch_tensor = None
for coord, device_tensor in zip(coords, device_tensors):
if view is None or view.is_local(coord):
torch_tensor = ttnn.to_torch(device_tensor)
break
if torch_tensor is None:
msg = "Failed to find local device tensor"
raise RuntimeError(msg)
return torch_tensor
_to_torch_zero_copy_warned = False
def _to_torch_zero_copy(t: ttnn.Tensor) -> torch.Tensor:
"""Convert a host ttnn tensor to a PyTorch tensor, preferring zero-copy.
Uses ``to_torch_with_padded_shape`` when available — for ROW_MAJOR host
tensors this wraps the existing buffer directly (zero-copy) instead of
copying through ``decode_tensor_data`` as ``to_torch`` always does.
Falls back to ``to_torch`` if the method is removed, with a one-time
warning so the performance regression is visible.
TODO: Once ``to_torch`` supports a ``padded_output`` parameter (or the
zero-copy path becomes the default for ROW_MAJOR), switch to that and
remove this helper.
"""
global _to_torch_zero_copy_warned
try:
return t.to_torch_with_padded_shape()
except AttributeError:
if not _to_torch_zero_copy_warned:
import logging
logging.getLogger(__name__).warning(
"to_torch_with_padded_shape unavailable, falling back to to_torch (slower d2h)"
)
_to_torch_zero_copy_warned = True
return ttnn.to_torch(t)
_TTNN_TO_TORCH_DTYPE = {
ttnn.bfloat16: torch.bfloat16,
ttnn.float32: torch.float32,
ttnn.uint8: torch.uint8,
ttnn.uint16: torch.int16,
ttnn.int32: torch.int32,
}
def _host_buffer_to_torch(buf, padded_shape: list[int], tt_dtype: ttnn.DataType) -> torch.Tensor:
"""Zero-copy conversion of a HostBuffer to a torch tensor.
Uses the DLPack protocol to get a uint8 view of the raw buffer memory,
then reinterprets it as the correct dtype and reshapes.
"""
torch_dtype = _TTNN_TO_TORCH_DTYPE[tt_dtype]
raw = torch.from_dlpack(buf)
return raw.view(torch_dtype).reshape(padded_shape)
def as_bf16(t: ttnn.Tensor) -> ttnn.Tensor:
"""Cast to bf16, skipping the op if the tensor already is bf16.
`ttnn.typecast` has no same-dtype early-out — it always dispatches a full elementwise copy
program — so a redundant cast costs a program launch per call, including on every trace
replay. Use this wherever a tensor may already have been cast by the caller.
"""
return t if t.dtype == ttnn.bfloat16 else ttnn.typecast(t, dtype=ttnn.bfloat16)
def float_to_unit_range(t: ttnn.Tensor) -> ttnn.Tensor:
"""On-device denormalization: map from [-1.0, 1.0] to [0.0, 1.0]."""
t = ttnn.to_layout(t, ttnn.TILE_LAYOUT)
t = ttnn.add(t, 1.0)
t = ttnn.multiply(t, 0.5)
t = ttnn.clamp(t, min=0.0, max=1.0)
t = ttnn.to_layout(t, ttnn.ROW_MAJOR_LAYOUT)
return t
def float_to_uint8(t: ttnn.Tensor) -> ttnn.Tensor:
"""On-device float-to-uint8: map from [-1.0, 1.0] to [0, 255]"""
t = ttnn.to_layout(t, ttnn.TILE_LAYOUT)
t = ttnn.add(t, 1.0) # shift to [0, 2.0]
t = ttnn.multiply(t, 0.5 * 255.0) # scale to [0, 1.0] then [0, 255]
t = ttnn.clamp(t, min=0.0, max=255.0)
t = ttnn.to_layout(t, ttnn.ROW_MAJOR_LAYOUT)
return ttnn.typecast(t, ttnn.uint8)
def _get_inter_host_axis(mesh_device: ttnn.MeshDevice, view, mesh_shape: tuple[int, ...]) -> int:
"""Return the mesh axis that spans multiple hosts (0 or 1).
In a 2D mesh, one axis typically spans hosts (inter-host) while the other
is fully local to each host (intra-host). Finds a local coordinate first,
then checks whether varying each axis stays local.
"""
from ttnn._ttnn.multi_device import MeshCoordinate
# Find any coordinate that is local to this host.
ref_r, ref_c = 0, 0
for r in range(mesh_shape[0]):
for c in range(mesh_shape[1]):
if view.is_local(MeshCoordinate(r, c)):
ref_r, ref_c = r, c
break
else:
continue
break
# Check axis 0: vary row while keeping the local column fixed.
if mesh_shape[0] > 1:
if not all(view.is_local(MeshCoordinate(r, ref_c)) for r in range(mesh_shape[0])):
return 0
# Check axis 1: vary column while keeping the local row fixed.
if mesh_shape[1] > 1:
if not all(view.is_local(MeshCoordinate(ref_r, c)) for c in range(mesh_shape[1])):
return 1
# Both axes are fully local — shouldn't happen in a true distributed env.
return 0
def _reassemble_2d(
mesh_coords: list,
shards: list[torch.Tensor],
shard_shape: list[int],
mesh_shape: tuple[int, ...],
concat_dims: list[int | None],
permute: tuple[int, ...] | None = None,
dtype: torch.dtype | None = None,
) -> torch.Tensor:
"""Reassemble per-device shards into a single tensor for a 2D mesh.
For the common case where both axes are concatenated, writes each shard
directly into a pre-allocated output buffer. When *permute* and/or *dtype*
are given the permutation and type conversion are fused into the scatter
write, halving total memory traffic.
"""
d0, d1 = concat_dims
if d0 is not None and d0 == d1:
# One tensor dim fractured row-major over the *whole* mesh (`ShardTensorToMesh(dim=d)` hands
# shard `i` to device `(i // cols, i % cols)`), so both mesh indices fold into a single
# linearised block offset rather than two independent per-axis concats. The general branch
# writes `slices[d0]` then `slices[d1]` into the same index, losing the row: shard `(r, c)`
# lands at position `c` for every `r`. The C++ equivalent refuses duplicate dims outright
# (`ttnn/core/tensor/xtensor/partition.cpp`).
cols = mesh_shape[1]
span = shard_shape[d0]
full_shape = list(shard_shape)
full_shape[d0] = span * mesh_shape[0] * cols
out_dtype = dtype if dtype is not None else shards[0].dtype
if permute is not None:
out = torch.empty([full_shape[p] for p in permute], dtype=out_dtype)
d0_out = list(permute).index(d0)
else:
out = torch.empty(full_shape, dtype=out_dtype)
d0_out = d0
for coord, shard in zip(mesh_coords, shards):
base = (int(coord[0]) * cols + int(coord[1])) * span
slices = [slice(None)] * len(full_shape)
slices[d0_out] = slice(base, base + span)
out[tuple(slices)] = shard.permute(*permute).contiguous() if permute is not None else shard
return out
if d0 is not None and d1 is not None:
s0, s1 = shard_shape[d0], shard_shape[d1]
full_shape = list(shard_shape)
full_shape[d0] *= mesh_shape[0]
full_shape[d1] *= mesh_shape[1]
ndim = len(full_shape)
if permute is not None:
out_shape = [full_shape[p] for p in permute]
out_dtype = dtype if dtype is not None else shards[0].dtype
perm_list = list(permute)
d0_out = perm_list.index(d0)
d1_out = perm_list.index(d1)
out = torch.empty(out_shape, dtype=out_dtype)
for coord, shard in zip(mesh_coords, shards):
r, c = int(coord[0]), int(coord[1])
slices = [slice(None)] * ndim
slices[d0_out] = slice(r * s0, (r + 1) * s0)
slices[d1_out] = slice(c * s1, (c + 1) * s1)
out[tuple(slices)] = shard.permute(*permute).contiguous()
return out
out_dtype = dtype if dtype is not None else shards[0].dtype
out = torch.empty(full_shape, dtype=out_dtype)
for coord, shard in zip(mesh_coords, shards):
r, c = int(coord[0]), int(coord[1])
slices = [slice(None)] * ndim
slices[d0] = slice(r * s0, (r + 1) * s0)
slices[d1] = slice(c * s1, (c + 1) * s1)
out[tuple(slices)] = shard
return out
if d0 is not None:
by_pos = sorted(zip(mesh_coords, shards), key=lambda x: int(x[0][0]))
return torch.cat([s for _, s in by_pos], dim=d0)
if d1 is not None:
by_pos = sorted(zip(mesh_coords, shards), key=lambda x: int(x[0][1]))
return torch.cat([s for _, s in by_pos], dim=d1)
return shards[0]
def fast_device_to_host(
tt_tensor: ttnn.Tensor,
mesh_device: ttnn.MeshDevice,
concat_dims: list[int | None],
ccl_manager=None,
root: int | None = None,
*,
pre_transfer_fn: Callable[[ttnn.Tensor], ttnn.Tensor] | None = None,
permute: tuple[int, ...] | None = None,
dtype: torch.dtype | None = None,
) -> torch.Tensor | None:
"""Fast D2H transfer using async DMA and zero-copy to_torch.
On a single-host system, this avoids the on-device all_gather by reading
all per-device shards concurrently with async DMA, converting to PyTorch
with zero-copy when possible, and concatenating on host.
On a multi-host (distributed) system, this uses a hybrid approach: an
on-device all_gather for the inter-host axis only, then fast async DMA +
zero-copy for local shards with host-side concatenation for the intra-host
axis.
Args:
tt_tensor: Multi-device ttnn tensor on ``mesh_device``.
mesh_device: The mesh device.
concat_dims: Per mesh axis, the tensor dimension to concatenate along,
or ``None`` to skip that axis. E.g. ``[3, 4]`` means concatenate
along dim 3 for mesh axis 0 and dim 4 for mesh axis 1.
ccl_manager: Optional :class:`CCLManager` instance. Required for
multi-host environments where only local devices are accessible.
root: If set, only the host with this MPI rank performs the D2H
transfer and returns the assembled tensor; all other ranks return
``None``. If ``None`` (default), all ranks perform D2H.
pre_transfer_fn: Optional on-device transformation applied just before
the DMA read. For multi-host, runs after ``mesh_partition``; for
single-host, runs right before ``.cpu()``. Typical use is
``float_to_uint8`` to shrink data before the PCIe transfer.
permute: If set, each shard is permuted before being written into the
output. Fuses the permutation into the scatter write so that no
intermediate tensor in the original layout is ever materialised.
dtype: Output dtype. When combined with ``permute``, the dtype
conversion is fused into the scatter write (single-pass copy).
"""
mesh_shape = tuple(mesh_device.shape)
if len(mesh_shape) != 2:
raise ValueError(
f"fast_device_to_host only supports 2D meshes, got mesh shape {mesh_shape} (ndim={len(mesh_shape)})"
)
if len(concat_dims) != 2:
raise ValueError(
f"concat_dims must have exactly 2 elements for a 2D mesh, got {len(concat_dims)}: {concat_dims}"
)
# --- Multi-host: hybrid on-device collective + fast local DMA -----------
if ttnn.using_distributed_env():
if ccl_manager is None:
msg = "fast_device_to_host requires ccl_manager in a distributed (multi-host) environment"
raise ValueError(msg)
view = mesh_device.get_view()
rank = int(ttnn.distributed_context_get_rank())
inter_host_axis = _get_inter_host_axis(mesh_device, view, mesh_shape)
intra_host_axis = 1 - inter_host_axis
# Step 1: On-device all_gather + repeat + mesh_partition.
# All_gather replicates the inter-host axis. Repeat + mesh_partition
# then re-shard it so every local device holds *unique* data,
# maximising PCIe bandwidth during the DMA read.
gathered_tensor = tt_tensor
inter_dim = concat_dims[inter_host_axis]
if inter_dim is not None and mesh_shape[inter_host_axis] > 1:
gathered_tensor = ttnn.to_layout(gathered_tensor, ttnn.TILE_LAYOUT)
gathered_tensor = ccl_manager.all_gather(
gathered_tensor,
dim=inter_dim,
mesh_axis=inter_host_axis,
use_hyperparams=True,
use_persistent_buffer=True,
)
n_hosts = int(ttnn.distributed_context_get_size())
if n_hosts > 1:
# A TILE-layout all_gather/repeat/mesh_partition is only safe when
# BOTH of the last two dims are tile-aligned. If the second-to-last
# dim (H) is not a multiple of TILE_SIZE, its tile padding makes the
# TILE repeat/mesh_partition interleave H rows across the shards ->
# horizontal-band noise corruption on multi-host (single-host skips
# this block, hence 4x8 was clean but 4x32 corrupted). Predict the
# per-chip shape after repeat (× n_hosts) and mesh_partition
# (÷ inter_axis_size) on inter_dim, and drop to ROW_MAJOR unless
# BOTH last dims are tile-aligned (matches the known-good behavior).
post_shape = list(gathered_tensor.shape)
post_shape[inter_dim] = post_shape[inter_dim] * n_hosts // mesh_shape[inter_host_axis]
if post_shape[-1] % ttnn.TILE_SIZE != 0 or post_shape[-2] % ttnn.TILE_SIZE != 0:
gathered_tensor = ttnn.to_layout(gathered_tensor, ttnn.ROW_MAJOR_LAYOUT)
repeat_dims = [1] * len(gathered_tensor.shape)
repeat_dims[inter_dim] = n_hosts
gathered_tensor = ttnn.repeat(gathered_tensor, repeat_dims)
gathered_tensor = ttnn.mesh_partition(gathered_tensor, dim=inter_dim, cluster_axis=inter_host_axis)
if pre_transfer_fn is not None:
gathered_tensor = pre_transfer_fn(gathered_tensor)
else:
gathered_tensor = ttnn.to_layout(gathered_tensor, ttnn.ROW_MAJOR_LAYOUT)
elif pre_transfer_fn is not None:
gathered_tensor = pre_transfer_fn(gathered_tensor)
# Step 2: Only root rank (if specified) does D2H.
if root is not None and rank != root:
return None
# Step 3: DMA all local shards and reassemble on host.
# Single .cpu() on the mesh tensor batches all local DMA reads into
# one C++ dispatch — the reader thread pool processes all device
# completion queues in parallel.
host_tensor = gathered_tensor.cpu(blocking=False)
ttnn.synchronize_device(mesh_device)
# Extract local shard buffers via get_shard (zero-copy, no MPI).
host_mesh_coords = list(host_tensor.tensor_topology().mesh_coords())
distributed_buf = host_tensor.host_buffer()
tt_dtype = host_tensor.dtype
padded_shape = list(host_tensor.padded_shape)
logical_shape = list(host_tensor.shape)
trim = tuple(slice(0, d) for d in logical_shape)
local_coords_and_bufs = []
for c in host_mesh_coords:
if not view.is_local(c):
continue
buf = distributed_buf.get_shard(c)
if buf is not None:
local_coords_and_bufs.append((c, buf))
shards = [_host_buffer_to_torch(buf, padded_shape, tt_dtype)[trim] for _, buf in local_coords_and_bufs]
# Build local mesh shape and 0-based coordinates for _reassemble_2d.
local_inter_positions = sorted({int(c[inter_host_axis]) for c, _ in local_coords_and_bufs})
local_intra_positions = sorted({int(c[intra_host_axis]) for c, _ in local_coords_and_bufs})
local_mesh_shape = [0, 0]
local_mesh_shape[inter_host_axis] = len(local_inter_positions)
local_mesh_shape[intra_host_axis] = len(local_intra_positions)
local_mesh_shape = tuple(local_mesh_shape)
inter_remap = {pos: i for i, pos in enumerate(local_inter_positions)}
intra_remap = {pos: i for i, pos in enumerate(local_intra_positions)}
local_coords = []
for c, _ in local_coords_and_bufs:
coord = [0, 0]
coord[inter_host_axis] = inter_remap[int(c[inter_host_axis])]
coord[intra_host_axis] = intra_remap[int(c[intra_host_axis])]
local_coords.append(tuple(coord))
return _reassemble_2d(local_coords, shards, logical_shape, local_mesh_shape, concat_dims, permute, dtype)
# --- Single-host: async DMA on all devices + host-side concat -----------
# Grab mesh coordinates from the device tensor before DMA.
mesh_coords = list(tt_tensor.tensor_topology().mesh_coords())
if pre_transfer_fn is not None:
tt_tensor = pre_transfer_fn(tt_tensor)
# Single .cpu() on the mesh tensor batches all DMA reads into one C++
# dispatch — host buffers are allocated in parallel and the reader thread
# pool processes all completion-queue reads concurrently.
host_tensor = tt_tensor.cpu(blocking=False)
ttnn.synchronize_device(mesh_device)
# Extract per-shard host tensors (single-host: just wraps each shard).
host_shard_tensors = ttnn.get_device_tensors(host_tensor)
# Zero-copy to_torch, trimmed to logical (un-padded) shape.
logical_shape = list(host_shard_tensors[0].shape)
trim = tuple(slice(0, d) for d in logical_shape)
shards = [_to_torch_zero_copy(s)[trim] for s in host_shard_tensors]
return _reassemble_2d(mesh_coords, shards, logical_shape, mesh_shape, concat_dims, permute, dtype)
def upsample(
x: ttnn.Tensor,
/,
*,
scale_factor: int,
memory_config: ttnn.MemoryConfig | None = None,
) -> ttnn.Tensor:
"""Wrapper around ttnn.upsample that allows for padded tensors."""
change_layout = x.layout == ttnn.TILE_LAYOUT and x.shape != x.padded_shape
if change_layout:
x = ttnn.to_layout(x, ttnn.ROW_MAJOR_LAYOUT)
x = ttnn.upsample(x, scale_factor=scale_factor, memory_config=memory_config)
if change_layout:
x = ttnn.to_layout(x, ttnn.TILE_LAYOUT)
return x
def unflatten(x: ttnn.Tensor, dim: int, sizes: Sequence[int]) -> ttnn.Tensor:
"""Expands a dimension of the input tensor over multiple dimensions.
Args:
x: ttnn.Tensor
The input tensor to unflatten
dim: int
Dimension to be unflattened, specified as an index into input.shape.
sizes: Sequence[int]
New shape of the unflattened dimension. One of its elements can be -1 in which case the corresponding output dimension is inferred. Otherwise, the product of sizes must equal
Returns:
ttnn.Tensor
The unflattened tensor.
"""
assert (
x.shape[dim] % abs(math.prod(sizes)) == 0
), f"The total number of elements in the new shape {sizes} must be equal or a factor of the number of elements (when using inferred dimensions) in the original shape {x.shape[dim]}"
new_shape = list(x.shape)
if dim == -1:
new_shape[-1:] = sizes
else:
new_shape[dim : dim + 1] = sizes
return ttnn.reshape(x, new_shape)
def full(
size: ttnn.Shape | Sequence[int],
fill_value: float,
*,
dtype: ttnn.DataType,
layout: ttnn.Layout = ttnn.TILE_LAYOUT,
device: ttnn.MeshDevice,
memory_config: ttnn.MemoryConfig | None = None,
) -> ttnn.Tensor:
"""Alternative to `ttnn.full` that supports tracing."""
if not isinstance(size, ttnn.Shape):
size = ttnn.Shape(size)
result = ttnn.allocate_tensor_on_device(size, dtype, layout, device, memory_config)
ttnn.fill(result, fill_value, output_tensor=result)
return result
def arange(
start: float,
end: float,
step: float = 1.0,
*,
dtype: ttnn.DataType,
layout: ttnn.Layout = ttnn.TILE_LAYOUT,
device: ttnn.MeshDevice,
memory_config: ttnn.MemoryConfig | None = None,
) -> ttnn.Tensor:
"""Alternative to `ttnn.arange` that supports tracing."""
x = full(
[math.ceil((end - start) / step)],
fill_value=step,
dtype=dtype,
layout=layout,
device=device,
memory_config=memory_config,
)
return ttnn.cumsum(x, 0) + (start - step)
_tril_cache: dict[tuple, ttnn.Tensor] = {}
_triu_cache: dict[tuple, ttnn.Tensor] = {}
def tril(
x: ttnn.Tensor,
/,
diagonal: int = 0,
*,
memory_config: ttnn.MemoryConfig | None = None,
output_tensor: ttnn.Tensor | None = None,
) -> ttnn.Tensor:
"""Alternative to `ttnn.tril` that supports tracing."""
device = x.device()
if device is None:
msg = "tril is not supported for host tensors"
raise ValueError(msg)
mask_shape = tuple(x.shape)[-2:]
cache_key = (mask_shape, device.id(), diagonal)
if cache_key in _tril_cache:
mask = _tril_cache[cache_key]
else:
mask = full(mask_shape, 1.0, dtype=ttnn.bfloat4_b, layout=ttnn.TILE_LAYOUT, device=device)
mask = ttnn.tril(mask, diagonal=diagonal)
_tril_cache[cache_key] = mask
return ttnn.mul(x, mask, memory_config=memory_config, output_tensor=output_tensor)
def triu(
x: ttnn.Tensor,
/,
diagonal: int = 0,
*,
memory_config: ttnn.MemoryConfig | None = None,
output_tensor: ttnn.Tensor | None = None,
) -> ttnn.Tensor:
"""Alternative to `ttnn.triu` that supports tracing."""
device = x.device()
if device is None:
msg = "triu is not supported for host tensors"
raise ValueError(msg)
mask_shape = tuple(x.shape)[-2:]
cache_key = (mask_shape, device.id(), diagonal)
if cache_key in _triu_cache:
mask = _triu_cache[cache_key]
else:
mask = full(mask_shape, 1.0, dtype=ttnn.bfloat4_b, layout=ttnn.TILE_LAYOUT, device=device)
mask = ttnn.triu(mask, diagonal=diagonal)
_triu_cache[cache_key] = mask
return ttnn.mul(x, mask, memory_config=memory_config, output_tensor=output_tensor)
def print_tensor_mem_info(tt: ttnn.Tensor):
logger.info("storage_type:", tt.storage_type())
logger.info("global shape:", tt.shape)
logger.info("global padded:", tt.padded_shape)
logger.info("global layout:", tt.layout)
logger.info("global dtype:", tt.dtype)
logger.info("global memory_config:", tt.memory_config())
topo = tt.tensor_topology()
logger.info("distribution_shape:", list(topo.distribution_shape()))
logger.info("placements:", [str(p) for p in topo.placements()]) # Replicate vs Shard(dim)
logger.info("mesh_coords:", list(topo.mesh_coords()))
# Per-device shard view
shards = ttnn.get_device_tensors(tt)
for i, s in enumerate(shards):
logger.info(f"\nshard[{i}]")
logger.info(" shape:", s.shape)
logger.info(" padded_shape:", s.padded_shape)
logger.info(" layout:", s.layout)
logger.info(" dtype:", s.dtype)
logger.info(" memory_config:", s.memory_config())
def prepare_weight_for_concatenated_input(
weight: torch.Tensor,
sizes: Sequence[int],
*,
device_count: int,
tile_pad_segments: bool = True,
) -> torch.Tensor:
"""Shard weight by device_count per segment and stack.
tile_pad_segments=True (fused concat): zero-pad each per-device segment K to a tile boundary
so minimal_matmul([prefix, suffix], weight) works for any channel count.
tile_pad_segments=False (materialized concat): contiguous stack, for use with ttnn.concat.
"""
segments = weight.split(sizes, dim=1)
padded_segments = []
for seg in segments:
unf = seg.unflatten(1, [device_count, -1]) # [out, device_count, K_seg/dev]
if tile_pad_segments:
k_per_dev = unf.shape[2]
k_padded = ((k_per_dev + ttnn.TILE_SIZE - 1) // ttnn.TILE_SIZE) * ttnn.TILE_SIZE
if k_padded != k_per_dev:
pad = torch.zeros(*unf.shape[:2], k_padded - k_per_dev, dtype=unf.dtype)
unf = torch.cat([unf, pad], dim=2)
padded_segments.append(unf)
return torch.cat(padded_segments, dim=2).flatten(1, 2)