File size: 5,863 Bytes
be3ecc8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0

"""
Automatic composition of multi-device sharded tensors using TensorTopology.

This module provides utilities to infer the correct MeshToTensor composer from a
sharded ttnn.Tensor's topology metadata and use it to compose shards on host.
"""

from typing import Optional

import torch
from loguru import logger

import ttnn

# ======================================================================================
# Public API
# ======================================================================================


def to_torch_auto_compose(tensor: ttnn.Tensor, device: Optional[ttnn.MeshDevice] = None) -> torch.Tensor:
    """
    Convert a (possibly multi-device) TTNN tensor to torch, automatically
    composing shards based on the tensor's topology.

    Args:
        tensor: The distributed tensor to convert
        device: Optional MeshDevice to use when the tensor lives on host

    Returns:
        PyTorch tensor with shards composed
    """
    composer = _infer_mesh_composer_from_topology(tensor, device=device)
    try:
        return ttnn.to_torch(tensor, mesh_composer=composer)
    except Exception as e:
        logger.error(f"Failed to convert tensor to torch with mesh_composer: {e}")
        raise


def extract_tensor_topology_info(
    tensor: ttnn.Tensor,
) -> tuple[list[object], list[int]]:
    """
    Extract placements and distribution shape from a tensor's topology.

    Returns:
        (placements, dist_shape)
    """
    topology = tensor.tensor_topology()
    placements = topology.placements()
    dist_shape = list(topology.distribution_shape())
    return placements, dist_shape


def get_device_from_tensor(tensor: ttnn.Tensor) -> Optional[ttnn.MeshDevice]:
    """Get device from tensor or fallback to provided mesh_device."""
    device = tensor.device()
    # tensor.device() returns None if the tensor is on the host (ttnn/core/tensor/tensor.cpp --> Tensor::device())
    if device is None:
        logger.debug("tensor.device() returns None, tensor is on the host")
    else:
        logger.debug(f"tensor.device() returns {device}")

    return device


# ======================================================================================
# Private Implementation
# ======================================================================================


def _infer_mesh_composer_from_topology(
    tensor: ttnn.Tensor, *, device: Optional[ttnn.MeshDevice] = None
) -> Optional[ttnn.CppMeshToTensor]:
    """
    Return a MeshToTensor composer inferred from the tensor's TensorTopology,
    or None if no composition is needed (fully replicated, single-device).

    Note: For ND meshes with replicated dimensions, the composer will concatenate
    all replicas, resulting in duplicated data. Callers may want to slice the
    result if only one copy is desired.

    Args:
        tensor: The distributed tensor to infer composer for

    Returns:
        MeshToTensor composer or None if no composition needed
    """
    placements, dist_shape = extract_tensor_topology_info(tensor)

    # No distribution or trivial 1-device case
    if len(dist_shape) == 0 or (len(dist_shape) == 1 and dist_shape[0] == 1):
        return None

    tensor_device = get_device_from_tensor(tensor)
    mesh_device = tensor_device or device
    if mesh_device is None:
        # As a last resort, try default device for backward-compatibility
        mesh_device = ttnn.GetDefaultDevice()
        if mesh_device is None:
            raise RuntimeError(
                "Tensor is on host and no mesh_device provided. "
                "Pass device=... to to_torch_auto_compose or set a default via ttnn.SetDefaultDevice(...)."
            )

    # Must match length (should be guaranteed by C++ TT_FATAL in ttnn/core/distributed/distributed_tensor.cpp)
    assert len(dist_shape) == len(placements)

    if len(dist_shape) == 1 and mesh_device.shape.dims() == 1:
        return _compose_1d_sharded(mesh_device, placements, dist_shape)
    else:
        # N >= 2 dimensions
        return _compose_nd_sharded(mesh_device, placements, dist_shape)


def _compose_1d_sharded(
    device: ttnn.MeshDevice,
    placements: list[object],
    dist_shape: list[int],
) -> Optional[ttnn.CppMeshToTensor]:
    """Handle 1D case - returns None if fully replicated."""
    p = placements[0]
    if isinstance(p, ttnn.PlacementShard):
        # Use ND composer with shape override to match the tensor's distribution
        composer_cfg = ttnn.MeshComposerConfig(dims=[p.dim], mesh_shape_override=ttnn.MeshShape(dist_shape))
        return ttnn.create_mesh_composer(device, composer_cfg)
    # Fully replicated - no composition needed
    return None


def _compose_nd_sharded(
    device: ttnn.MeshDevice,
    placements: list[object],
    dist_shape: list[int],
) -> ttnn.CppMeshToTensor:
    """
    Handle ND (N>=2) case.

    For replicated mesh dims, we use dim 0 as convention (the composed result
    will include all replicas concatenated, which is typically not desired but
    is how the C++ API works).
    """
    dims = []
    shape_override = []
    for i, p in enumerate(placements):
        if isinstance(p, ttnn.PlacementShard):
            dims.append(p.dim)
            shape_override.append(dist_shape[i])
        else:
            assert isinstance(p, ttnn.PlacementReplicate)
            # [INFO] steal from TensorDistribution2x4Test test case in test_distributed_tensor.cpp
            # Replicated: use dim 0 as convention
            dims.append(0)
            # Replicated: use shape 1 to skip concatenation
            shape_override.append(1)

    composer_cfg = ttnn.MeshComposerConfig(dims=dims, mesh_shape_override=ttnn.MeshShape(shape_override))
    return ttnn.create_mesh_composer(device, composer_cfg)