clef / code /models /common /modules /tt_ccl.py
tt-hous's picture
Add files using upload-large-folder tool
b025706 verified
Raw History Blame Contribute Delete
7.04 kB
# 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}")