File size: 3,518 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
# SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
# SPDX-License-Identifier: Apache-2.0

"""Low-level helpers for releasing runtime-owned TT tensor trees."""

from __future__ import annotations

from dataclasses import dataclass, field
from typing import Any, Iterable

import ttnn


@dataclass
class TensorResourceOrphan:
    """An owned tensor tree whose failed releases remain retryable."""

    values: Any
    deallocated_tensor_ids: set[int] = field(default_factory=set)


def best_effort_deallocate_owned_tensors(
    value: Any,
    completed: set[int] | None = None,
) -> list[BaseException]:
    """Release all reachable TT tensors while retaining successful progress."""

    if completed is None:
        completed = set()
    failures: list[BaseException] = []
    visiting: set[int] = set()

    def visit(item: Any) -> None:
        if item is None:
            return
        item_id = id(item)
        if isinstance(item, ttnn.Tensor):
            if item_id in completed:
                return
            try:
                ttnn.deallocate(item)
            except BaseException as error:
                failures.append(error)
            else:
                completed.add(item_id)
            return
        if item_id in visiting:
            return
        if isinstance(item, dict):
            visiting.add(item_id)
            for nested in item.values():
                visit(nested)
            visiting.remove(item_id)
        elif isinstance(item, (list, tuple, set)):
            visiting.add(item_id)
            for nested in item:
                visit(nested)
            visiting.remove(item_id)
        else:
            owned_values = getattr(item, "owned_tensor_values", None)
            if callable(owned_values):
                visiting.add(item_id)
                visit(owned_values())
                visiting.remove(item_id)

    visit(value)
    return failures


def release_orphans(orphans: list[TensorResourceOrphan]) -> list[BaseException]:
    """Retry orphan releases in place, retaining only incomplete entries."""

    failures: list[BaseException] = []
    remaining: list[TensorResourceOrphan] = []
    for orphan in orphans:
        orphan_failures = best_effort_deallocate_owned_tensors(
            orphan.values,
            orphan.deallocated_tensor_ids,
        )
        failures.extend(orphan_failures)
        if orphan_failures:
            remaining.append(orphan)
    orphans[:] = remaining
    return failures


def attach_cleanup_failures(
    primary: BaseException,
    failures: Iterable[BaseException],
    *,
    note: str = "Cleanup also encountered {count} failure(s)",
) -> None:
    """Attach cleanup failures without replacing the primary exception."""

    failures = tuple(failures)
    if not failures:
        return
    previous = tuple(getattr(primary, "cleanup_failures", ()))
    primary.cleanup_failures = previous + failures
    add_note = getattr(primary, "add_note", None)
    if callable(add_note):
        add_note(note.format(count=len(failures)))


def raise_cleanup_failures(failures: Iterable[BaseException]) -> None:
    """Raise the first cleanup failure while retaining every additional one."""

    failures = tuple(failures)
    if not failures:
        raise ValueError("at least one cleanup failure is required")
    primary = failures[0]
    attach_cleanup_failures(
        primary,
        failures[1:],
        note="Cleanup also encountered {count} additional failure(s)",
    )
    raise primary