File size: 8,142 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
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0

"""
Tensor utility functions for TTTv2 modules.
"""

import json
import math
import re

import torch

import ttnn

# Standard tile size - hardware constant
TILE_SIZE = ttnn.TILE_SIZE  # 32


def nearest_multiple(value: int, multiple: int) -> int:
    return math.ceil(value / multiple) * multiple


def nearest_32(value: int) -> int:
    return nearest_multiple(value, TILE_SIZE)


def num_to_core_range_set(num_cores: int):
    assert num_cores < 8 or num_cores % 8 == 0
    num_x = min(num_cores, 8)
    num_y = num_cores // num_x
    assert num_x * num_y == num_cores
    return ttnn.CoreRangeSet(
        {
            ttnn.CoreRange(
                ttnn.CoreCoord(0, 0),
                ttnn.CoreCoord(num_x - 1, num_y - 1),
            )
        }
    )


def get_out_subblock_w(per_core_n: int, out_subblock_h: int = 1) -> int:
    out_subblock_w = 4
    while out_subblock_w > 1:
        if out_subblock_w * out_subblock_h <= 4 and per_core_n % out_subblock_w == 0:
            break
        out_subblock_w -= 1
    return out_subblock_w


def align_shape_to_tile(shape, tile_size: int = TILE_SIZE):
    """Round up the last two dimensions of *shape* to multiples of *tile_size*.

    This is the recommended replacement for the deprecated ``ttnn.pad_to_tile_shape``.

    Args:
        shape: An iterable of dimension sizes (list, tuple, or ttnn.Shape).
        tile_size: Tile dimension to align to (default 32).

    Returns:
        List[int]: A new shape with the last two dims tile-aligned.
    """
    import math

    result = list(shape)
    if len(result) >= 1:
        result[-1] = math.ceil(result[-1] / tile_size) * tile_size
    if len(result) >= 2:
        result[-2] = math.ceil(result[-2] / tile_size) * tile_size
    return result


def get_rot_transformation_mat(dhead: int = TILE_SIZE) -> torch.Tensor:
    """
    Create rotation transformation matrix for RoPE.

    Constructs a permutation matrix that pairs adjacent dimensions with
    signs (+1, -1) for the RoPE rotation:
        [0, 1] → +1 at (0,1), -1 at (1,0)
        [2, 3] → +1 at (2,3), -1 at (3,2)
        ...

    Used by ttnn.experimental.rotary_embedding_llama.

    Args:
        dhead: Matrix dimension. Must equal TILE_SIZE. Use TILE_SIZE for decode.

    Returns:
        torch.Tensor of shape [1, 1, dhead, dhead].
    """
    rot_emb_matrix = torch.zeros(1, 1, dhead, dhead)
    rot_emb_matrix[..., torch.arange(0, dhead, 2), torch.arange(1, dhead, 2)] = 1
    rot_emb_matrix[..., torch.arange(1, dhead, 2), torch.arange(0, dhead, 2)] = -1
    return rot_emb_matrix


def zeros_like_kv_cache(batch_size: int, n_kv_heads: int, max_seq_len: int, head_dim: int) -> torch.Tensor:
    """Create zeros tensor for standard KV cache."""
    return torch.zeros((batch_size, n_kv_heads, max_seq_len, head_dim))


def zeros_like_paged_cache(paged_config, n_kv_heads: int, head_dim: int) -> torch.Tensor:
    """Create zeros tensor for paged KV cache."""
    return torch.zeros((paged_config.max_num_blocks, n_kv_heads, paged_config.block_size, head_dim))


# todo)) add a on-device pad_dim_to_size function?
def pad_dim_to_size(x: "torch.Tensor", dim: int, size: int) -> "torch.Tensor":
    """Pads the specified dimension of the input tensor with zeros."""
    if dim < 0:
        dim = x.dim() + dim
    current_size = x.size(dim)
    pad_size = size - current_size

    if pad_size < 0:
        raise ValueError(f"Target size {size} is smaller than current size {current_size} on dim {dim}")

    if pad_size == 0:
        return x

    pad = [0] * (2 * x.dim())
    pad_index = 2 * (x.dim() - dim - 1)
    pad[pad_index + 1] = pad_size

    return torch.nn.functional.pad(x, pad, mode="constant", value=0)


def pad_to_shape(x: "torch.Tensor", target_shape: tuple[int, ...], pad_value: float = 0.0) -> "torch.Tensor":
    """Pad tensor to target_shape in a single F.pad call (more efficient than per-dim padding)."""
    if x.shape == target_shape:
        return x

    # F.pad expects: (left_last, right_last, left_second_last, right_second_last, ...)
    pad = []
    for orig, target in zip(reversed(x.shape), reversed(target_shape)):
        if target < orig:
            raise ValueError(f"Target size {target} is smaller than current size {orig}")
        pad.extend([0, target - orig])

    return torch.nn.functional.pad(x, pad, mode="constant", value=pad_value)


def get_padded_hidden_dim(hidden_dim: int, num_devices: int, tile_size: int = 32) -> int:
    """
    Compute padded hidden_dim to satisfy ttnn.from_torch's tile alignment constraint.

    ttnn.from_torch requires physical shard shapes to be tile-aligned. When sharding
    a tensor across devices, each shard_dim = hidden_dim / num_devices must be
    divisible by tile_size.

    We pad the global tensor first, then shard evenly so only the last shard has padding.
    """
    shard_dim = hidden_dim // num_devices
    padded_shard = ((shard_dim + tile_size - 1) // tile_size) * tile_size
    return padded_shard * num_devices


def parse_shard_dims_from_mesh_mapper_config(mesh_mapper_config: ttnn.MeshMapperConfig) -> list[int]:
    """
    Parse shard dimensions from MeshMapperConfig's repr.

    MeshMapperConfig doesn't expose .placements directly, but repr shows them:
        'MeshMapperConfig(placements: [PlacementShard(-1)], mesh_shape_override=MeshShape([8]))'

    This parses out the shard dimensions (e.g., [-1]) from PlacementShard entries.
    Returns empty list if no PlacementShard found (e.g., replicated).

    Note: This is a workaround until TTNN exposes .placements directly.
    """
    config_repr = repr(mesh_mapper_config)
    matches = re.findall(r"PlacementShard\((-?\d+)\)", config_repr)
    return [int(d) for d in matches]


def memory_config_to_dict(memory_config: ttnn.MemoryConfig):
    # Convert to plain types for deterministic serialization.
    return {
        "memory_layout": str(memory_config.memory_layout),
        "buffer_type": str(memory_config.buffer_type),
        "shard_spec": str(memory_config.shard_spec),
        "is_sharded": bool(memory_config.is_sharded()),
        "interleaved": bool(memory_config.interleaved),
        "hash": int(memory_config.__hash__()),
    }


def compute_kernel_config_to_str(compute_kernel_config: ttnn.WormholeComputeKernelConfig):
    # Backward compat shim; prefer compute_kernel_config_to_dict + serialize_config.
    cfg = compute_kernel_config_to_dict(compute_kernel_config)
    return serialize_config(cfg)


def compute_kernel_config_to_dict(compute_kernel_config: ttnn.WormholeComputeKernelConfig):
    return {
        "math_fidelity": str(compute_kernel_config.math_fidelity),
        "math_approx_mode": str(compute_kernel_config.math_approx_mode),
        "fp32_dest_acc_en": bool(compute_kernel_config.fp32_dest_acc_en),
        "packer_l1_acc": bool(compute_kernel_config.packer_l1_acc),
        "dst_full_sync_en": bool(compute_kernel_config.dst_full_sync_en),
        "throttle_level": str(compute_kernel_config.throttle_level),
    }


def program_config_to_str(program_config: ttnn.MatmulMultiCoreReuseMultiCastDRAMShardedProgramConfig):
    # Backward compat shim; prefer program_config_to_dict + serialize_config.
    cfg = program_config_to_dict(program_config)
    return serialize_config(cfg)


def program_config_to_dict(program_config):
    if hasattr(program_config, "to_json"):
        d = json.loads(program_config.to_json())
        d["type"] = type(program_config).__name__
        return d
    else:
        return {"type": type(program_config).__name__, "repr": repr(program_config)}


def serialize_config(cfg_dict: dict, fmt: str = "json") -> str:
    if fmt == "json":
        return json.dumps(cfg_dict, sort_keys=True)
    if fmt == "yaml":
        try:
            import yaml
        except ImportError as exc:  # pragma: no cover - optional dependency
            raise RuntimeError("PyYAML is required for yaml serialization") from exc
        return yaml.safe_dump(cfg_dict, sort_keys=True)
    raise ValueError(f"Unsupported format: {fmt}")