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
|