File size: 8,536 Bytes
2415c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
# SPDX-License-Identifier: Apache-2.0

from types import SimpleNamespace

from models.common.modules.lazy_buffer import LazyBuffer
from models.common.modules.sampling import sampling_1d
from models.common.modules.sampling.sampling_1d import Sampling1D
from models.common.sampling.tt_log_probs import LogProbsCalculator


class FakeTensor:
    def __init__(self, name):
        self.name = name


class FakeLogProbsCalculator:
    instances = []

    def __init__(self, *_args):  # External construction hook; intentionally generic.
        self.buffer = FakeTensor(f"log-probs-{len(self.instances)}")
        self.instances.append(self)

    def release(self):
        if self.buffer is not None:
            sampling_1d.ttnn.deallocate(self.buffer)
            self.buffer = None


def _lazy_buffer(name):
    return LazyBuffer(
        source=name,
        dtype=object(),
        layout=object(),
        device=object(),
        mesh_mapper=object(),
    )


def _sampler_config(index_offsets, owned_specs):
    return SimpleNamespace(
        index_offsets=index_offsets,
        invalid_vocab_mask=owned_specs.get("invalid_vocab_mask"),
        invalid_vocab_tail_mask=owned_specs.get("invalid_vocab_tail_mask"),
        invalid_vocab_tail_width=32,
        seeds=owned_specs["seeds"],
        user_ids=owned_specs["user_ids"],
        mesh_device=object(),
        tt_ccl=None,
        sub_core_grids=None,
        start_core=None,
        max_batch_size=32,
        is_resolved=lambda: True,
    )


def test_sampling_release_is_idempotent_preserves_borrowed_and_allows_reload(monkeypatch):
    deallocated = []
    allocations = []

    # ttnn.from_torch is overloaded; retain backend-specific keyword absorption.
    def from_torch(source, **_kwargs):
        tensor = FakeTensor(f"{source}-{len(allocations)}")
        allocations.append(tensor)
        return tensor

    monkeypatch.setattr(sampling_1d.ttnn, "Tensor", FakeTensor)
    monkeypatch.setattr(sampling_1d.ttnn, "from_torch", from_torch)
    monkeypatch.setattr(sampling_1d.ttnn, "deallocate", deallocated.append)

    from models.common import utils as common_utils

    FakeLogProbsCalculator.instances = []
    monkeypatch.setattr(common_utils, "LogProbsCalculator", FakeLogProbsCalculator)

    borrowed_index_offsets = FakeTensor("borrowed-index-offsets")
    owned_specs = {
        "invalid_vocab_mask": _lazy_buffer("invalid-vocab-mask"),
        "invalid_vocab_tail_mask": _lazy_buffer("invalid-vocab-tail-mask"),
        "seeds": _lazy_buffer("seeds"),
        "user_ids": _lazy_buffer("user-ids"),
    }
    config = _sampler_config(borrowed_index_offsets, owned_specs)
    sampler = object.__new__(Sampling1D)
    sampler.config = config
    sampler._device_buffers_loaded = False

    sampler.load_device_buffers()
    first_owned = {name: getattr(sampler, f"_{name}") for name in owned_specs}
    first_calculator = sampler._log_probs_calculator
    assert sampler._index_offsets is borrowed_index_offsets

    sampler.release()
    deallocations_after_first_release = list(deallocated)
    sampler.release()

    assert deallocated == deallocations_after_first_release
    assert borrowed_index_offsets not in deallocated
    assert set(first_owned.values()).issubset(deallocated)
    assert first_calculator.buffer is None
    assert not sampler._device_buffers_loaded
    assert all(spec._value is None for spec in owned_specs.values())

    sampler.load_device_buffers()

    assert sampler._device_buffers_loaded
    assert sampler._index_offsets is borrowed_index_offsets
    assert sampler._log_probs_calculator is not first_calculator
    for name, old_tensor in first_owned.items():
        assert getattr(sampler, f"_{name}") is not old_tensor


def test_partial_load_failure_preserves_primary_and_releases_before_reload(monkeypatch, expect_error):
    allocations = []
    deallocated = []
    allocation_error = RuntimeError("seed allocation failed")
    cleanup_error = RuntimeError("mask cleanup failed once")
    fail_allocation = True
    fail_cleanup = True

    def from_torch(source, **_kwargs):
        if source == "seeds" and fail_allocation:
            raise allocation_error
        tensor = FakeTensor(f"{source}-{len(allocations)}")
        allocations.append(tensor)
        return tensor

    def deallocate(value):
        nonlocal fail_cleanup
        deallocated.append(value)
        if value.name.startswith("invalid-vocab-mask") and fail_cleanup:
            fail_cleanup = False
            raise cleanup_error

    monkeypatch.setattr(sampling_1d.ttnn, "Tensor", FakeTensor)
    monkeypatch.setattr(sampling_1d.ttnn, "from_torch", from_torch)
    monkeypatch.setattr(sampling_1d.ttnn, "deallocate", deallocate)

    from models.common import utils as common_utils

    FakeLogProbsCalculator.instances = []
    monkeypatch.setattr(common_utils, "LogProbsCalculator", FakeLogProbsCalculator)

    borrowed_index_offsets = FakeTensor("borrowed-index-offsets")
    owned_specs = {
        "invalid_vocab_mask": _lazy_buffer("invalid-vocab-mask"),
        "seeds": _lazy_buffer("seeds"),
        "user_ids": _lazy_buffer("user-ids"),
    }
    sampler = object.__new__(Sampling1D)
    sampler.config = _sampler_config(borrowed_index_offsets, owned_specs)
    sampler._device_buffers_loaded = False

    with expect_error(RuntimeError, "seed allocation failed") as caught:
        sampler.load_device_buffers()

    assert caught.value is allocation_error
    assert cleanup_error in caught.value.cleanup_failures
    assert owned_specs["invalid_vocab_mask"]._value is not None
    assert not sampler._device_buffers_loaded
    assert borrowed_index_offsets not in deallocated

    sampler.release()
    assert owned_specs["invalid_vocab_mask"]._value is None

    fail_allocation = False
    sampler.load_device_buffers()
    assert sampler._device_buffers_loaded
    sampler.release()


def test_sampling_release_is_best_effort_and_retries_only_failed_buffer(monkeypatch, expect_error):
    attempts = []
    cleanup_error = RuntimeError("seed cleanup failed once")

    monkeypatch.setattr(sampling_1d.ttnn, "Tensor", FakeTensor)
    monkeypatch.setattr(
        sampling_1d.ttnn,
        "from_torch",
        lambda source, **_kwargs: FakeTensor(source),
    )

    from models.common import utils as common_utils

    FakeLogProbsCalculator.instances = []
    monkeypatch.setattr(common_utils, "LogProbsCalculator", FakeLogProbsCalculator)

    borrowed_index_offsets = FakeTensor("borrowed-index-offsets")
    owned_specs = {
        "seeds": _lazy_buffer("seeds"),
        "user_ids": _lazy_buffer("user-ids"),
    }
    sampler = object.__new__(Sampling1D)
    sampler.config = _sampler_config(borrowed_index_offsets, owned_specs)
    sampler._device_buffers_loaded = False
    sampler.load_device_buffers()
    failed = sampler._seeds

    def deallocate(value):
        attempts.append(value)
        if value is failed and attempts.count(value) == 1:
            raise cleanup_error

    monkeypatch.setattr(sampling_1d.ttnn, "deallocate", deallocate)

    with expect_error(RuntimeError, "seed cleanup failed once") as caught:
        sampler.release()

    assert caught.value is cleanup_error
    assert sampler._seeds is failed
    assert owned_specs["seeds"]._value is failed
    assert owned_specs["user_ids"]._value is None

    sampler.release()

    assert attempts.count(failed) == 2
    assert attempts.count(borrowed_index_offsets) == 0
    assert owned_specs["seeds"]._value is None


def test_log_probs_calculator_release_deallocates_unique_owned_tensors_once(monkeypatch):
    deallocated = []
    monkeypatch.setattr(sampling_1d.ttnn, "deallocate", deallocated.append)

    shared = FakeTensor("shared")
    tensors = {
        "global_max": shared,
        "global_exp_sum": FakeTensor("global-exp-sum"),
        "mask": FakeTensor("mask"),
        "output_tensor": shared,
        "topk_logprobs_output": FakeTensor("topk-logprobs"),
        "topk_indices_output": FakeTensor("topk-indices"),
    }
    calculator = object.__new__(LogProbsCalculator)
    for name, tensor in tensors.items():
        setattr(calculator, name, tensor)

    calculator.release()
    calculator.release()

    assert len(deallocated) == len({id(tensor) for tensor in tensors.values()})
    assert {id(tensor) for tensor in deallocated} == {id(tensor) for tensor in tensors.values()}
    assert all(getattr(calculator, name) is None for name in tensors)