File size: 9,480 Bytes
9aa90e0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.

# SPDX-License-Identifier: Apache-2.0

import math
from typing import Optional

import torch

import ttnn


class PaddingConfig:
    """
    Configuration for model padding to enable tensor parallelism.

    This class handles the calculation and validation of padding requirements
    for attention heads and hidden dimensions to make them divisible by the
    tensor parallel factor.
    """

    def __init__(
        self, original_heads: int, target_heads: int, head_dim: int, tensor_parallel_factor: Optional[int] = None
    ):
        """
        Initialize padding configuration.

        Args:
            original_heads: Original number of attention heads
            target_heads: Target number of heads (must be >= original_heads)
            head_dim: Dimension per attention head (remains constant)
            tensor_parallel_factor: TP factor for validation (optional)
        """
        self.original_heads = original_heads
        self.target_heads = target_heads
        self.head_dim = head_dim

        # Calculate derived dimensions
        self.original_dim = original_heads * head_dim
        self.target_dim = target_heads * head_dim

        # Padding amounts
        self.head_padding = target_heads - original_heads
        self.dim_padding = self.target_dim - self.original_dim

        # Validation
        self._validate(tensor_parallel_factor)

    def _validate(self, tensor_parallel_factor: Optional[int]):
        """Validate padding configuration."""
        if self.target_heads < self.original_heads:
            raise ValueError(f"target_heads ({self.target_heads}) must be >= original_heads ({self.original_heads})")

        if self.head_dim <= 0:
            raise ValueError(f"head_dim must be positive, got {self.head_dim}")

        if tensor_parallel_factor is not None:
            if self.target_heads % tensor_parallel_factor != 0:
                raise ValueError(
                    f"target_heads ({self.target_heads}) must be divisible by "
                    f"tensor_parallel_factor ({tensor_parallel_factor})"
                )

    @classmethod
    def from_tensor_parallel_factor(
        cls, original_heads: int, head_dim: int, tensor_parallel_factor: int
    ) -> "PaddingConfig":
        """
        Create padding config automatically based on tensor parallel factor.

        Args:
            original_heads: Original number of attention heads
            head_dim: Dimension per attention head
            tensor_parallel_factor: Desired TP factor

        Returns:
            PaddingConfig with target_heads rounded up to be divisible by TP factor
        """
        target_heads = math.ceil(original_heads / tensor_parallel_factor) * tensor_parallel_factor
        return cls(original_heads, target_heads, head_dim, tensor_parallel_factor)

    def is_padding_needed(self) -> bool:
        """Return True if any padding is needed."""
        return self.head_padding > 0

    def __repr__(self) -> str:
        return (
            f"PaddingConfig(original_heads={self.original_heads}, "
            f"target_heads={self.target_heads}, head_dim={self.head_dim}, "
            f"dim_padding={self.dim_padding})"
        )


def pad_weight_tensor(
    weight: torch.Tensor, padding_config: PaddingConfig, pad_input_dim: bool = False, pad_output_dim: bool = False
) -> torch.Tensor:
    """
    Pad a weight tensor according to padding configuration.

    Args:
        weight: Weight tensor to pad (typically 2D: [input_dim, output_dim])
        padding_config: Padding configuration
        pad_input_dim: Whether to pad the input dimension
        pad_output_dim: Whether to pad the output dimension

    Returns:
        Padded weight tensor
    """
    if not padding_config.is_padding_needed():
        return weight

    padded_weight = weight.clone()

    # Pad input dimension (dimension 0 for transposed weights)
    if pad_input_dim and padding_config.dim_padding > 0:
        input_padding = torch.zeros(
            padding_config.dim_padding, weight.shape[1], dtype=weight.dtype, device=weight.device
        )
        padded_weight = torch.cat([padded_weight, input_padding], dim=0)

    # Pad output dimension (dimension 1 for transposed weights)
    if pad_output_dim and padding_config.dim_padding > 0:
        output_padding = torch.zeros(
            padded_weight.shape[0], padding_config.dim_padding, dtype=weight.dtype, device=weight.device
        )
        padded_weight = torch.cat([padded_weight, output_padding], dim=1)

    return padded_weight


def pad_bias_tensor(bias: torch.Tensor, padding_config: PaddingConfig) -> torch.Tensor:
    """
    Pad a bias tensor according to padding configuration.

    Args:
        bias: Bias tensor to pad
        padding_config: Padding configuration

    Returns:
        Padded bias tensor
    """
    if not padding_config.is_padding_needed():
        return bias

    if padding_config.dim_padding > 0:
        bias_padding = torch.zeros(padding_config.dim_padding, dtype=bias.dtype, device=bias.device)
        return torch.cat([bias, bias_padding], dim=-1)

    return bias


def pad_qkv_weights(
    q_weight: torch.Tensor, k_weight: torch.Tensor, v_weight: torch.Tensor, padding_config: PaddingConfig
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    """
    Pad QKV weight tensors for attention layers using structured padding.

    Args:
        q_weight: Query projection weight (in_dim, out_dim)
        k_weight: Key projection weight (in_dim, out_dim)
        v_weight: Value projection weight (in_dim, out_dim)
        padding_config: Padding configuration

    Returns:
        Tuple of padded (q_weight, k_weight, v_weight)
    """
    if not padding_config.is_padding_needed():
        return q_weight, k_weight, v_weight

    original_dim = padding_config.original_dim
    target_dim = padding_config.target_dim

    def pad_qkv_weight(weight):
        in_dim, out_dim = weight.shape
        mult = out_dim // original_dim
        assert mult == 3, f"Only 3-way fused QKV weight matrices are supported, given weight shape {weight.shape}"

        # Reshape: (in_dim, mult_factor * original_dim) -> (in_dim, mult_factor, original_dim)
        weight = weight.reshape(weight.shape[0], mult_factor, original_dim)

        # Pad output dimension: (in_dim, mult_factor, original_dim) -> (in_dim, mult_factor, target_dim)
        output_padding = torch.zeros(
            weight.shape[0], mult_factor, target_dim - original_dim, dtype=weight.dtype, device=weight.device
        )
        weight = torch.cat([weight, output_padding], dim=2)

        # Reshape back: (in_dim, mult_factor, target_dim) -> (in_dim, mult_factor * target_dim)
        weight = weight.reshape(weight.shape[0], -1)

        return weight

    padded_q = pad_qkv_weight(q_weight)
    padded_k = pad_qkv_weight(k_weight)
    padded_v = pad_qkv_weight(v_weight)

    return padded_q, padded_k, padded_v


def pad_qkv_biases(
    q_bias: Optional[torch.Tensor],
    k_bias: Optional[torch.Tensor],
    v_bias: Optional[torch.Tensor],
    padding_config: PaddingConfig,
) -> tuple[Optional[torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor]]:
    """
    Pad QKV bias tensors for attention layers using structured padding.

    Args:
        q_bias: Query projection bias (can be None)
        k_bias: Key projection bias (can be None)
        v_bias: Value projection bias (can be None)
        padding_config: Padding configuration

    Returns:
        Tuple of padded (q_bias, k_bias, v_bias)
    """
    if not padding_config.is_padding_needed():
        return q_bias, k_bias, v_bias

    original_dim = padding_config.original_dim
    target_dim = padding_config.target_dim

    def pad_qkv_bias(bias):
        if bias is None:
            return None

        orig_shape = bias.shape
        mult_factor = orig_shape[0] // original_dim
        assert mult_factor == 3, "Only 3-way fused QKV bias matrices are supported"

        bias = bias.reshape(mult_factor, original_dim)

        # Pad: (mult_factor, original_dim) -> (mult_factor, target_dim)
        bias_padding = torch.zeros(mult_factor, target_dim - original_dim, dtype=bias.dtype, device=bias.device)
        bias = torch.cat([bias, bias_padding], dim=1)

        # Reshape back: (mult_factor, target_dim) -> (mult_factor * target_dim,)
        bias = bias.reshape(orig_shape)

        return bias

    padded_q_bias = pad_qkv_bias(q_bias)
    padded_k_bias = pad_qkv_bias(k_bias)
    padded_v_bias = pad_qkv_bias(v_bias)

    return padded_q_bias, padded_k_bias, padded_v_bias


def get_padded_vision_seq_len(N, num_devices):
    divisor = ttnn.TILE_SIZE * num_devices

    # Calculate padding needed to make seq_len divisible by both tile size and num_devices
    padded_seq_len = math.ceil(N / divisor) * divisor
    padding = padded_seq_len - N
    shard_size = padded_seq_len // num_devices
    return padded_seq_len


def pad_vision_seq_parallel(tensor, num_devices):
    """
    Sequence parallelism shards the vision tensor in dim2.
    dim2 must be divisible by tile size and num_devices.
    """
    seq_len = tensor.shape[2]
    padded_seq_len = get_padded_vision_seq_len(seq_len, num_devices)
    pad_len = padded_seq_len - seq_len

    # Pad the sequence length dimension (dim2) on the right
    if pad_len > 0:
        tensor = torch.nn.functional.pad(tensor, (0, 0, 0, pad_len))

    return tensor