Download code/models/tt_dit/utils/tensor.py from stisiTT/flux2-dev-qb2: direct link, hf CLI and curl.
- Browser
- Download file 35.8 kB
-
https://huggingface.co/stisiTT/flux2-dev-qb2/resolve/main/code/models/tt_dit/utils/tensor.py
- Command line
-
hf download hf://stisiTT/flux2-dev-qb2/code/models/tt_dit/utils/tensor.py
-
curl -L -o tensor.py https://huggingface.co/stisiTT/flux2-dev-qb2/resolve/main/code/models/tt_dit/utils/tensor.py
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) | |