stisiTT's picture
Add files using upload-large-folder tool
9aa90e0 verified
Raw History Blame Contribute Delete
19.3 kB
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import time
from abc import ABC, abstractmethod
from pathlib import Path
from typing import TYPE_CHECKING, NamedTuple, overload
from loguru import logger
from typing_extensions import deprecated
import ttnn
from ..utils import tensor
from ..utils.progress import Watchdog as _Watchdog
from ..utils.substate import pop_substate
if TYPE_CHECKING:
from collections.abc import Iterable, Iterator, Mapping, MutableSequence, Sequence
from typing import Any
import torch
class IncompatibleKeys(NamedTuple):
missing_keys: list[str]
unexpected_keys: list[str]
class LoadingError(Exception):
pass
class _LoadProgress:
"""Heartbeat for weight loading. A 46 GB load+convert is silent for minutes and
reads as a hang; this emits a timed progress line so it never looks dead."""
def __init__(self, total: int, what: str) -> None:
self._total = total
self._done = 0
self._what = what
self._t0 = time.monotonic()
self._last = 0.0
def tick(self) -> None:
self._done += 1
elapsed = time.monotonic() - self._t0
# Throttle to ~5s so fast loads stay quiet and slow ones still show life.
if elapsed - self._last >= 5.0 or self._done == self._total:
self._last = elapsed
pct = 100 * self._done / self._total if self._total else 100
logger.info(f"{self._what}: {self._done}/{self._total} tensors ({pct:.0f}%), {elapsed:.0f}s")
class Module(ABC):
def __init__(self) -> None:
self._children = {}
self._parameters = {}
self._is_loaded = False
self.coresident_exclusions = None # modules that cannot be resident in memory at the same time as this module. They should be deallocated before this module is loaded.
self._coresident_peers: list[Module] = []
def named_children(self) -> Iterator[tuple[str, Module]]:
yield from self._children.items()
def named_parameters(self) -> Iterator[tuple[str, Parameter]]:
yield from self._parameters.items()
def add_module(self, name: str, module: Module) -> None:
self._children[name] = module
def __setattr__(self, name: str, value: Any) -> None: # noqa: ANN401
super().__setattr__(name, value)
if name in ("_children", "_parameters"):
return
children = self.__dict__.get("_children")
parameters = self.__dict__.get("_parameters")
if isinstance(value, Module):
if children is None:
msg = "cannot assign child module before Module.__init__() call"
raise AttributeError(msg)
self._children[name] = value
elif isinstance(value, Parameter):
if parameters is None:
msg = "cannot assign parameter before Module.__init__() call"
raise AttributeError(msg)
self._parameters[name] = value
else:
if children is not None:
children.pop(name, None)
if parameters is not None:
parameters.pop(name, None)
def __delattr__(self, name: str) -> None:
children = self.__dict__.get("_children")
parameters = self.__dict__.get("_parameters")
if children is not None:
children.pop(name, None)
if parameters is not None:
parameters.pop(name, None)
super().__delattr__(name)
def _prepare_torch_state(self, state: dict[str, torch.Tensor]) -> None: # noqa: B027
"""Prepare a PyTorch `state_dict` in place before loading.
Override this method to adjust entries before loading them into submodules and parameters.
This method should modify `state` in place and, where possible, avoid raising exceptions for
missing keys; skip them instead. This way, missing keys can be collected and returned by
`load_torch_state_dict`.
"""
def _num_parameters(self) -> int:
return len(self._parameters) + sum(
child._num_parameters() for _, child in self.named_children()
) # noqa: SLF001
def _load_torch_state_dict_inner(
self,
state_dict: Mapping[str, torch.Tensor],
*,
module_key_prefix: str,
missing_keys: MutableSequence[str],
unexpected_keys: MutableSequence[str],
progress: _LoadProgress | None = None,
) -> None:
state_dict = dict(state_dict)
self._prepare_torch_state(state_dict)
for name, child in self.named_children():
child_state = pop_substate(state_dict, name)
try:
child._load_torch_state_dict_inner( # noqa: SLF001
child_state,
module_key_prefix=f"{module_key_prefix}{name}.",
missing_keys=missing_keys,
unexpected_keys=unexpected_keys,
progress=progress,
)
except LoadingError:
raise
except Exception as err:
msg = f"an exception occurred while loading '{module_key_prefix}{name}'"
raise LoadingError(msg) from err
for name, parameter in self.named_parameters():
if name in state_dict:
try:
parameter.load_torch_tensor(state_dict.pop(name))
except LoadingError as err:
msg = f"while loading '{module_key_prefix}{name}': {err}"
raise LoadingError(msg) from err
if progress is not None:
progress.tick()
else:
missing_keys.append(f"{module_key_prefix}{name}")
for name in state_dict:
unexpected_keys.append(f"{module_key_prefix}{name}")
def _mark_loaded(self) -> None:
"""Recursively mark this module and all descendants as loaded."""
self._is_loaded = True
for _, child in self.named_children():
child._mark_loaded() # noqa: SLF001
def load_torch_state_dict(self, state_dict: Mapping[str, torch.Tensor], *, strict: bool = True) -> IncompatibleKeys:
"""Load PyTorch state dict into module parameters.
Args:
state_dict: Mapping of parameter names to PyTorch tensors.
strict: If `True`, raises ValueError on missing or unexpected keys.
Returns:
`IncompatibleKeys` containing lists of missing and unexpected keys.
"""
missing_keys = []
unexpected_keys = []
self.evict_coresident_exclusions()
with _Watchdog(f"convert {type(self).__name__}"):
self._load_torch_state_dict_inner(
state_dict,
module_key_prefix="",
missing_keys=missing_keys,
unexpected_keys=unexpected_keys,
progress=_LoadProgress(self._num_parameters(), "converting weights to device"),
)
if strict and (missing_keys or unexpected_keys):
parts = []
if missing_keys:
parts.append("missing Torch state keys: " + ", ".join(missing_keys))
if unexpected_keys:
parts.append("unexpected Torch state keys: " + ", ".join(unexpected_keys))
raise ValueError("; ".join(parts))
self._mark_loaded()
return IncompatibleKeys(missing_keys, unexpected_keys)
@deprecated("Use load_torch_state_dict instead")
def load_state_dict(self, state_dict: Mapping[str, torch.Tensor]) -> None:
self.load_torch_state_dict(state_dict)
def save(self, directory: str | Path, /, *, prefix: str = "") -> None:
directory = Path(directory)
directory.mkdir(exist_ok=True, parents=True)
for name, child in self.named_children():
child.save(directory, prefix=f"{prefix}{name}.")
for name, parameter in self.named_parameters():
parameter.save(directory / f"{prefix}{name}.tensorbin")
def load(self, directory: str | Path, /, *, prefix: str = "") -> None:
directory = Path(directory)
# Top-level only: announce the cache load and arm a timer watchdog so a slow/stalled
# module load is never silent. Signature is left untouched — subclasses (Mochi/Wan)
# override this method.
watchdog = None
if prefix == "":
logger.info(f"loading {self._num_parameters()} cached weight tensors from '{directory}'...")
watchdog = _Watchdog(f"load-cache {type(self).__name__}").__enter__()
try:
self.evict_coresident_exclusions()
for name, child in self.named_children():
child.load(directory, prefix=f"{prefix}{name}.")
for name, parameter in self.named_parameters():
path = directory / f"{prefix}{name}.tensorbin"
try:
parameter.load(path)
except LoadingError as err:
msg = f"{err} while loading '{path}'"
raise LoadingError(msg) from err
self._is_loaded = True
finally:
if watchdog is not None:
watchdog.__exit__()
def deallocate_weights(self) -> None:
"""Deallocate all parameter weights from device memory recursively."""
for _, child in self.named_children():
child.deallocate_weights()
for _, parameter in self.named_parameters():
parameter.deallocate()
self._is_loaded = False
def is_loaded(self) -> bool:
return self._is_loaded
def register_coresident_exclusions(self, *args: Module) -> None:
"""
Register modules that cannot be resident in memory at the same time as this module.
They should be deallocated before this module is loaded. See `evict_coresident_exclusions` .
Args:
*args: Modules that cannot be co-resident in memory with this module.
"""
if self.coresident_exclusions is None:
self.coresident_exclusions = set()
self.coresident_exclusions.update(args)
def evict_coresident_exclusions(self) -> None:
"""Evict the modules that cannot be resident in memory at the same time as this module."""
if self.coresident_exclusions is not None:
for module in self.coresident_exclusions:
module.deallocate_weights()
@abstractmethod
def forward(self, *args: Any, **kwargs: Any) -> Any: # noqa: ANN401
pass
def __call__(self, *args: Any, **kwargs: Any) -> Any: # noqa: ANN401
return self.forward(*args, **kwargs)
class ModuleList(Module):
def __init__(self, modules: Iterable[Module] = ()) -> None:
super().__init__()
for i, m in enumerate(modules):
self.add_module(str(i), m)
def forward(self) -> None:
msg = "forward() should not be called on ModuleList. Iterate over the modules instead."
raise RuntimeError(msg)
def append(self, module: Module) -> None:
self.add_module(str(len(self)), module)
def __len__(self) -> int:
return len(self._children)
@overload
def __getitem__(self, key: int) -> Module:
...
@overload
def __getitem__(self, key: slice) -> ModuleList:
...
def __getitem__(self, key: int | slice) -> Module | ModuleList:
n = len(self._children)
if isinstance(key, slice):
start, stop, step = key.indices(n)
return ModuleList(self._children[str(i)] for i in range(start, stop, step))
if isinstance(key, int):
if key < 0:
key += n
if key < 0 or key >= n:
raise IndexError
return self._children[str(key)]
msg = f"expected int or slice argument, got {key}"
raise ValueError(msg)
class UnregisteredModule:
"""A wrapper for Module instances that prevents automatic registration in parent modules.
This class provides a way to hold references to Module instances without having them
automatically registered as child modules when assigned as attributes to another Module. This is
useful when you need to store a module reference but don't want it to appear in the module
hierarchy or participate in operations like parameter loading.
The UnregisteredModule acts as a transparent proxy, forwarding all attribute access and method
calls to the wrapped module.
Args:
module: The Module instance to wrap and keep unregistered.
Example:
>>> class MyModule(Module):
... def __init__(self):
... super().__init__()
... # This will be registered as a child module
... self.registered_child = SomeModule()
... # This will NOT be registered as a child module
... self.unregistered_child = UnregisteredModule(SomeModule())
...
... def forward(self, x):
... return self.registered_child(x) + self.unregistered_child(x)
...
>>> my_module = MyModule()
>>> list(my_module.named_children())
[('registered_child', <SomeModule instance>)]
"""
def __init__(self, module: Module) -> None:
self.module = module
def __getattr__(self, name: str) -> Any: # noqa: ANN401
return getattr(self.module, name)
def __call__(self, *args: Any, **kwargs: Any) -> Any: # noqa: ANN401
return self.module(*args, **kwargs)
class Parameter:
def __init__(
self,
*,
total_shape: Sequence[int],
device: ttnn.MeshDevice,
layout: ttnn.Layout = ttnn.Layout.TILE,
dtype: ttnn.DataType = ttnn.bfloat16,
memory_config: ttnn.MemoryConfig = ttnn.DRAM_MEMORY_CONFIG,
pad_value: float | None = None,
mesh_axes: Sequence[int | None] | None = None,
on_host: bool = False,
) -> None:
"""Initialize a Parameter for use in a Module.
The parameter is initially uninitialized. It is typically populated via the parent module's
`load_torch_state_dict()`. Alternatively, call `load_torch_tensor()` or assign the `data`
property directly with a correctly shaped and distributed `ttnn.Tensor`.
Args:
total_shape: The global shape of the parameter tensor across all mesh devices.
device: The mesh device on which the parameter is stored. If `on_host` is
`True`, this is used only to create the mesh mapper for distributing the tensor.
layout: See `ttnn.from_torch()`. Defaults to `ttnn.Layout.TILE`.
dtype: See `ttnn.from_torch()`. Defaults to `ttnn.bfloat16`.
memory_config: See `ttnn.from_torch()`. Defaults to `ttnn.DRAM_MEMORY_CONFIG`.
pad_value: See `ttnn.from_torch()`. Defaults to `None`.
mesh_axes: Maps tensor dimensions to mesh device axes for distribution.
For a rank-3 tensor whose second and third dimensions are sharded on mesh axes 0 and
1, respectively, use `[None, 0, 1]`.
on_host: If `True`, keep the tensor in host memory instead of device memory.
"""
total_shape = tuple(total_shape)
mesh_axes = tuple(mesh_axes) if mesh_axes is not None else (None,) * len(total_shape)
tensor.verify_tensor_mesh_axes(mesh_axes, tensor_rank=len(total_shape), mesh_rank=len(list(device.shape)))
local_shape = list(total_shape)
for tensor_dim, mesh_axis in enumerate(mesh_axes):
if mesh_axis is not None:
n = device.shape[mesh_axis]
if local_shape[tensor_dim] % n != 0:
msg = (
f"tensor with shape {total_shape} cannot be evenly distributed over mesh with shape "
f"{tuple(device.shape)} along mesh axis {mesh_axis} and tensor dimension {tensor_dim} "
)
raise ValueError(msg)
local_shape[tensor_dim] //= n
local_shape = tuple(local_shape)
self.total_shape = total_shape
self.local_shape = local_shape
self.device = device
self.layout = layout
self.dtype = dtype
self.memory_config = memory_config
self.pad_value = pad_value
self.mesh_axes = mesh_axes
self.on_host = on_host
self._data = None
def load_torch_tensor(self, torch_tensor: torch.Tensor, /) -> None:
shape = tuple(torch_tensor.shape)
if shape != self.total_shape:
msg = f"expected tensor shape {self.total_shape}, got {shape}"
raise LoadingError(msg)
self.data = tensor.from_torch(
torch_tensor,
device=self.device,
layout=self.layout,
dtype=self.dtype,
memory_config=self.memory_config,
pad_value=self.pad_value,
mesh_axes=self.mesh_axes,
on_host=self.on_host,
)
def save(self, path: str | Path, /) -> None:
ttnn.dump_tensor(path, self.data, mode=ttnn.DumpTensorMode.LOCAL)
def load(self, path: str | Path, /) -> None:
try:
tensor = ttnn.load_tensor(path, device=None if self.on_host else self.device)
except RuntimeError as err:
msg = f"TT-NN error «{err}»"
raise LoadingError(msg) from err
self.data = tensor
@property
def data(self) -> ttnn.Tensor:
if self._data is None:
msg = "parameter has no data"
raise RuntimeError(msg)
return self._data
@data.setter
def data(self, value: ttnn.Tensor) -> None:
self._check_data(value)
self._data = value
def deallocate(self) -> None:
"""Deallocate the parameter's device memory."""
if self._data is not None:
ttnn.deallocate(self._data)
self._data = None
def _check_data(self, value: ttnn.Tensor) -> None:
if self.on_host:
if value.device() is not None:
msg = "expected host tensor, got device tensor"
raise LoadingError(msg)
elif value.device() is None:
msg = "expected device tensor, got host tensor"
raise LoadingError(msg)
elif value.device() != self.device:
msg = "device mismatch"
raise LoadingError(msg)
if value.dtype != self.dtype:
msg = f"dtype mismatch: expected {self.dtype}, got {value.dtype}"
raise LoadingError(msg)
if value.layout != self.layout:
msg = f"layout mismatch: expected {self.layout}, got {value.layout}"
raise LoadingError(msg)
if value.memory_config() != self.memory_config:
msg = f"memory config mismatch: expected {self.memory_config}, got {value.memory_config()}"
raise LoadingError(msg)
if value.shape != self.local_shape:
msg = f"shape mismatch: expected {self.local_shape}, got {tuple(value.shape)}"
raise LoadingError(msg)