Download code/models/tt_dit/layers/module.py from stisiTT/flux2-dev-qb2: direct link, hf CLI and curl.
- Browser
- Download file 19.3 kB
-
https://huggingface.co/stisiTT/flux2-dev-qb2/resolve/main/code/models/tt_dit/layers/module.py
- Command line
-
hf download hf://stisiTT/flux2-dev-qb2/code/models/tt_dit/layers/module.py
-
curl -L -o module.py https://huggingface.co/stisiTT/flux2-dev-qb2/resolve/main/code/models/tt_dit/layers/module.py
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) | |
| 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() | |
| 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) | |
| def __getitem__(self, key: int) -> Module: | |
| ... | |
| 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 | |
| def data(self) -> ttnn.Tensor: | |
| if self._data is None: | |
| msg = "parameter has no data" | |
| raise RuntimeError(msg) | |
| return self._data | |
| 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) | |