File size: 5,632 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
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0

from dataclasses import dataclass

import torch


@dataclass(frozen=True)
class InvalidVocabTailMask:
    """A compact additive mask for the tile-aligned invalid tail of the final vocab shard."""

    mask: torch.Tensor
    tail_width: int
    shard_width: int
    num_vocab_shards: int


def build_invalid_vocab_mask(
    vocab_size: int,
    padded_vocab_size: int,
    max_batch_size: int,
    *,
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor | None:
    """Build an additive logits mask for LM-head vocabulary padding.

    LM heads may pad output weights so the sharded matmul has legal tile/device
    dimensions. Those padded columns produce logits, but they are not real token
    IDs and must be masked before argmax or top-k sampling.
    """
    if vocab_size < 0:
        raise ValueError(f"vocab_size must be non-negative, got {vocab_size}")
    if padded_vocab_size < vocab_size:
        raise ValueError(f"padded_vocab_size ({padded_vocab_size}) must be >= vocab_size ({vocab_size})")
    if max_batch_size <= 0:
        raise ValueError(f"max_batch_size must be positive, got {max_batch_size}")
    if vocab_size == padded_vocab_size:
        return None

    if not torch.empty((), dtype=dtype).is_floating_point():
        raise TypeError(f"dtype must be a floating point torch dtype, got {dtype}")

    mask = torch.zeros(1, 1, max_batch_size, padded_vocab_size, dtype=dtype)
    mask[..., vocab_size:] = torch.finfo(dtype).min
    return mask


def _validate_cluster_shape(cluster_shape: tuple[int, int] | list[int]) -> tuple[int, int]:
    cluster_shape = tuple(cluster_shape)
    if len(cluster_shape) != 2:
        raise ValueError(f"cluster_shape must have two dimensions, got {cluster_shape}")

    rows, cols = int(cluster_shape[0]), int(cluster_shape[1])
    if rows <= 0 or cols <= 0:
        raise ValueError(f"cluster_shape dimensions must be positive, got {cluster_shape}")
    return rows, cols


def get_vocab_num_shards(
    cluster_shape: tuple[int, int] | list[int],
    sampling_all_gather_axis: int = 0,
) -> int:
    """Return how many mesh partitions own contiguous slices of the vocab dimension."""
    rows, cols = _validate_cluster_shape(cluster_shape)

    if rows == 1 and cols == 1:
        return 1
    if rows == 1:
        return cols
    if cols == 1:
        return rows
    if sampling_all_gather_axis == 0:
        return rows
    if sampling_all_gather_axis == 1:
        return cols
    raise ValueError(f"sampling_all_gather_axis must be 0 or 1, got {sampling_all_gather_axis}")


def build_tail_invalid_vocab_mask(
    vocab_size: int,
    padded_vocab_size: int,
    max_batch_size: int,
    cluster_shape: tuple[int, int] | list[int],
    sampling_all_gather_axis: int = 0,
    *,
    dtype: torch.dtype = torch.bfloat16,
    tile_size: int = 32,
) -> InvalidVocabTailMask | None:
    """Build a compact mask for padding that lives only at the final shard tail.

    The sampling logits are sharded into equal local vocab widths. For model
    shapes like Qwen3-32B on T3K, all invalid IDs are a small tile-aligned suffix
    of the last local shard. In that case callers can mask only the local tail
    slice instead of adding a full-vocab all-zero mask on every device.

    Returns ``None`` when the invalid range is not a tile-aligned final-shard
    suffix; callers should use ``build_invalid_vocab_mask`` as the correctness
    fallback.
    """
    if vocab_size < 0:
        raise ValueError(f"vocab_size must be non-negative, got {vocab_size}")
    if padded_vocab_size < vocab_size:
        raise ValueError(f"padded_vocab_size ({padded_vocab_size}) must be >= vocab_size ({vocab_size})")
    if max_batch_size <= 0:
        raise ValueError(f"max_batch_size must be positive, got {max_batch_size}")
    if tile_size <= 0:
        raise ValueError(f"tile_size must be positive, got {tile_size}")
    if vocab_size == padded_vocab_size:
        return None
    if not torch.empty((), dtype=dtype).is_floating_point():
        raise TypeError(f"dtype must be a floating point torch dtype, got {dtype}")

    num_vocab_shards = get_vocab_num_shards(cluster_shape, sampling_all_gather_axis)
    if padded_vocab_size % num_vocab_shards != 0:
        return None

    shard_width = padded_vocab_size // num_vocab_shards
    tail_width = padded_vocab_size - vocab_size
    if tail_width > shard_width:
        return None
    if tail_width % tile_size != 0 or (shard_width - tail_width) % tile_size != 0:
        return None

    mask = torch.zeros(1, 1, max_batch_size, tail_width * num_vocab_shards, dtype=dtype)
    final_tail_start = tail_width * (num_vocab_shards - 1)
    mask[..., final_tail_start:] = torch.finfo(dtype).min
    return InvalidVocabTailMask(
        mask=mask,
        tail_width=tail_width,
        shard_width=shard_width,
        num_vocab_shards=num_vocab_shards,
    )


def get_vocab_shard_dims(
    cluster_shape: tuple[int, int] | list[int],
    sampling_all_gather_axis: int = 0,
) -> tuple[int | None, int | None]:
    """Return the 2D mesh mapper dims for sharding vocab over the sampling TP axis."""
    rows, cols = _validate_cluster_shape(cluster_shape)

    if rows == 1 and cols == 1:
        return (None, None)
    if rows == 1:
        return (None, 3)
    if cols == 1:
        return (3, None)
    if sampling_all_gather_axis == 0:
        return (3, None)
    if sampling_all_gather_axis == 1:
        return (None, 3)
    raise ValueError(f"sampling_all_gather_axis must be 0 or 1, got {sampling_all_gather_axis}")