File size: 7,044 Bytes
b025706 | 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 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 | # SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
import ttnn
from models.common.device_utils import get_device_name
# =============================================================================
# CCL tuning defaults - shared across all TTTv2 modules
# =============================================================================
# Default number of chunks per synchronization barrier in CCL operations.
# Higher values reduce sync overhead but increase latency per chunk.
CCL_CHUNKS_PER_SYNC = 10
# Default number of worker threads per Ethernet link for CCL operations.
CCL_NUM_WORKERS_PER_LINK = 2
# Default number of double-buffered channels per CCL link.
CCL_NUM_BUFFERS_PER_CHANNEL = 2
# =============================================================================
# TT_CCL cache - one instance per mesh_device (semaphores are hardware resources)
# =============================================================================
_tt_ccl_cache: dict[int, "TT_CCL"] = {}
def get_tt_ccl(mesh_device: ttnn.MeshDevice) -> "TT_CCL":
"""Get or create TT_CCL for mesh_device (cached per device id)."""
mesh_id = mesh_device.id()
if mesh_id not in _tt_ccl_cache:
_tt_ccl_cache[mesh_id] = TT_CCL(mesh_device)
return _tt_ccl_cache[mesh_id]
def clear_tt_ccl_cache():
"""Clear cache (for testing)."""
_tt_ccl_cache.clear()
def _get_local_num_devices(mesh_device: Optional[ttnn.MeshDevice]) -> int:
"""Return the number of devices visible to the current host process."""
if mesh_device is None:
raise ValueError("mesh_device is required to determine CCL link counts")
try:
local_device_ids = mesh_device.get_device_ids()
except Exception as exc:
raise ValueError("CCL link detection requires at least one host-local device") from exc
if not local_device_ids:
raise ValueError("CCL link detection requires at least one host-local device")
return len(local_device_ids)
# =============================================================================
# TT_CCL class
# =============================================================================
class TT_CCL:
def __init__(
self,
mesh_device,
):
self.mesh_device = mesh_device
self.sub_device_crs = ttnn.CoreRangeSet(
{
ttnn.CoreRange(
ttnn.CoreCoord(0, 0),
ttnn.CoreCoord(
self.mesh_device.compute_with_storage_grid_size().x - 1,
self.mesh_device.compute_with_storage_grid_size().y - 1,
),
)
}
)
self.barrier_semaphore_idx = [0, 0, 0]
self.barrier_semaphore_handles = [[], [], []]
self.ag_semaphores_idx = [0, 0, 0]
self.ag_semaphore_handles = [[], [], []]
self.rs_semaphores_idx = [0, 0, 0]
self.rs_semaphore_handles = [[], [], []]
# cluster-axis-0, cluster-axis-1, no-cluster-axis
for i in range(3):
# double buffered semaphores
for _ in range(2):
self.barrier_semaphore_handles[i].append(
ttnn.create_global_semaphore(self.mesh_device, self.sub_device_crs, 0)
)
self.ag_semaphore_handles[i].append(
[ttnn.create_global_semaphore(self.mesh_device, self.sub_device_crs, 0) for _ in range(2)]
)
self.rs_semaphore_handles[i].append(
[ttnn.create_global_semaphore(self.mesh_device, self.sub_device_crs, 0) for _ in range(3)]
)
def get_and_cycle_barrier_semaphore_handle(self, cluster_axis=None):
semaphore_index = 2 if cluster_axis is None else cluster_axis
current_idx = self.barrier_semaphore_idx[semaphore_index]
self.barrier_semaphore_idx[semaphore_index] = (current_idx + 1) % 2
return self.barrier_semaphore_handles[semaphore_index][current_idx]
def get_and_cycle_ag_semaphore_handles(self, cluster_axis=None):
semaphore_index = 2 if cluster_axis is None else cluster_axis
current_idx = self.ag_semaphores_idx[semaphore_index]
self.ag_semaphores_idx[semaphore_index] = (current_idx + 1) % 2
return self.ag_semaphore_handles[semaphore_index][current_idx]
def get_and_cycle_rs_semaphore_handles(self, cluster_axis=None):
semaphore_index = 2 if cluster_axis is None else cluster_axis
current_idx = self.rs_semaphores_idx[semaphore_index]
self.rs_semaphores_idx[semaphore_index] = (current_idx + 1) % 2
return self.rs_semaphore_handles[semaphore_index][current_idx]
def get_num_links(self, cluster_axis=None):
"""Get the number of available Ethernet links for CCL operations on this mesh device."""
return get_num_links(self.mesh_device, cluster_axis)
# =============================================================================
# Topology auto-detection
# =============================================================================
# todo)) work with the CCL team to find opportunity to simplify this --> e.g., build into TTNN APIs?
def default_topology(mesh_device: ttnn.MeshDevice) -> Optional[ttnn.Topology]:
"""Auto-detect CCL topology based on cluster type and device count."""
num_devices = mesh_device.get_num_devices()
cluster_type = ttnn.cluster.get_cluster_type()
if (num_devices == 8 and cluster_type == ttnn.cluster.ClusterType.T3K) or (
num_devices == 4 and cluster_type == ttnn.cluster.ClusterType.P150_X4
):
# NOTE: we always want to do ring if it is available
return ttnn.Topology.Ring
elif num_devices > 1:
# NOTE: this should be a fallback when the ring is not available
return ttnn.Topology.Linear
return None
def get_num_links(mesh_device: ttnn.MeshDevice, cluster_axis: int | None = None) -> int:
"""
Get the number of available Ethernet links for CCL operations.
Args:
mesh_device: The mesh device to query.
cluster_axis: Optional cluster axis to query links for.
- 0: vertical axis (North-South).
- 1: horizontal axis (East-West).
- None: minimum across all axes.
Returns:
int: The number of available links.
"""
device_name = get_device_name(mesh_device, num_devices=_get_local_num_devices(mesh_device))
link_dict = {
"P100": (0, 0),
"P150": (0, 0),
"N150": (0, 0),
"N300": (1, 1),
"T3K": (1, 1),
"P150x4": (2, 2),
"P150x8": (2, 2),
"P300": (2, 2),
"BHGLX": (2, 2),
"TG": (4, 4),
"N150x4": (1, 1),
}
device_links = link_dict[device_name]
if cluster_axis is None:
return min(device_links)
if cluster_axis in (0, 1):
return device_links[cluster_axis]
raise ValueError(f"Unsupported cluster_axis: {cluster_axis}")
|