Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
@@ -0,0 +1,92 @@
|
||||
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# For a list of all contributors, visit:
|
||||
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
|
||||
|
||||
import sys
|
||||
|
||||
from ._compat import ( # noqa: F401
|
||||
SUPPORTS_AUTOTUNE_CACHE,
|
||||
TRITON_ABOVE_3_4_0,
|
||||
TRITON_ABOVE_3_5_1,
|
||||
TRITON_ABOVE_3_7_1,
|
||||
autotune_cache_kwargs,
|
||||
find_spec_cached,
|
||||
has_usable_nvcc,
|
||||
)
|
||||
from ._config import ( # noqa: F401
|
||||
FLA_CACHE_RESULTS,
|
||||
FLA_CI_ENV,
|
||||
FLA_DISABLE_TENSOR_CACHE,
|
||||
FLA_TENSOR_CACHE_SIZE,
|
||||
)
|
||||
from ._decorators import ( # noqa: F401
|
||||
Action,
|
||||
checkpoint,
|
||||
contiguous,
|
||||
deprecate_kwarg,
|
||||
input_guard,
|
||||
require_version,
|
||||
tensor_cache,
|
||||
)
|
||||
from ._device import ( # noqa: F401
|
||||
IS_AMD,
|
||||
IS_ARM,
|
||||
IS_GATHER_SUPPORTED,
|
||||
IS_INTEL,
|
||||
IS_INTEL_ALCHEMIST,
|
||||
IS_NPU,
|
||||
IS_NVIDIA,
|
||||
IS_NVIDIA_BLACKWELL,
|
||||
IS_NVIDIA_HOPPER,
|
||||
IS_NVIDIA_SM100,
|
||||
IS_NVIDIA_SM120,
|
||||
IS_TF32_SUPPORTED,
|
||||
IS_TMA_SUPPORTED,
|
||||
Backend,
|
||||
autocast_custom_bwd,
|
||||
autocast_custom_fwd,
|
||||
check_environments,
|
||||
check_pytorch_version,
|
||||
check_shared_mem,
|
||||
custom_device_ctx,
|
||||
device,
|
||||
device_name,
|
||||
device_platform,
|
||||
device_torch_lib,
|
||||
get_all_max_shared_mem,
|
||||
get_available_device,
|
||||
get_device_capability,
|
||||
get_device_smem_optin,
|
||||
get_multiprocessor_count,
|
||||
map_triton_backend_to_torch_device,
|
||||
)
|
||||
from ._testing import assert_close, get_abs_err, get_err_ratio # noqa: F401
|
||||
|
||||
|
||||
def _register_aliases():
|
||||
current_module = sys.modules[__name__]
|
||||
for key in (
|
||||
'IS_AMD',
|
||||
'IS_ARM',
|
||||
'IS_INTEL',
|
||||
'IS_INTEL_ALCHEMIST',
|
||||
'IS_NVIDIA',
|
||||
'IS_NPU',
|
||||
'IS_NVIDIA_BLACKWELL',
|
||||
'IS_NVIDIA_HOPPER',
|
||||
'IS_NVIDIA_SM100',
|
||||
'IS_NVIDIA_SM120',
|
||||
'IS_TF32_SUPPORTED',
|
||||
'IS_GATHER_SUPPORTED',
|
||||
'IS_TMA_SUPPORTED',
|
||||
):
|
||||
if hasattr(current_module, key):
|
||||
setattr(current_module, key.lower(), getattr(current_module, key))
|
||||
|
||||
|
||||
_register_aliases()
|
||||
|
||||
del _register_aliases
|
||||
@@ -0,0 +1,65 @@
|
||||
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# For a list of all contributors, visit:
|
||||
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
|
||||
|
||||
import functools
|
||||
import importlib.metadata
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
from importlib.util import find_spec
|
||||
from pathlib import Path
|
||||
|
||||
import triton
|
||||
from packaging import version as package_version
|
||||
|
||||
from ._config import FLA_CACHE_RESULTS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TRITON_ABOVE_3_4_0 = package_version.parse(triton.__version__) >= package_version.parse("3.4.0")
|
||||
TRITON_ABOVE_3_5_1 = package_version.parse(triton.__version__) >= package_version.parse("3.5.1")
|
||||
TRITON_ABOVE_3_7_1 = package_version.parse(triton.__version__) >= package_version.parse("3.7.1")
|
||||
|
||||
SUPPORTS_AUTOTUNE_CACHE = "cache_results" in inspect.signature(triton.autotune).parameters
|
||||
autotune_cache_kwargs = {"cache_results": FLA_CACHE_RESULTS} if SUPPORTS_AUTOTUNE_CACHE else {}
|
||||
|
||||
|
||||
@functools.cache
|
||||
def find_spec_cached(name):
|
||||
return find_spec(name)
|
||||
|
||||
|
||||
@functools.cache
|
||||
def has_usable_nvcc() -> bool:
|
||||
"""Whether a usable nvcc compiler is available for TileLang's JIT.
|
||||
|
||||
Mirrors the guesses in ``tilelang.env._find_cuda_home`` (env
|
||||
CUDA_HOME/CUDA_PATH, nvcc on PATH, the ``nvidia-cuda-nvcc`` wheel,
|
||||
/usr/local/cuda), but verifies the nvcc binary actually exists —
|
||||
only ``nvidia-cuda-nvcc`` >= 13.0 ships it, the ``-cu12`` variant
|
||||
installs just ptxas.
|
||||
"""
|
||||
cuda_home = os.environ.get("CUDA_HOME") or os.environ.get("CUDA_PATH")
|
||||
if cuda_home is not None and (Path(cuda_home) / "bin" / "nvcc").exists():
|
||||
return True
|
||||
if shutil.which("nvcc") is not None:
|
||||
return True
|
||||
try:
|
||||
files = importlib.metadata.files("nvidia-cuda-nvcc") or []
|
||||
except importlib.metadata.PackageNotFoundError:
|
||||
files = []
|
||||
if any(f.name in ("nvcc", "nvcc.exe") for f in files):
|
||||
return True
|
||||
if (Path("/usr/local/cuda") / "bin" / "nvcc").exists():
|
||||
return True
|
||||
|
||||
logger.info(
|
||||
"[FLA Backend] TileLang is installed but no usable nvcc compiler was found; falling back to Triton. "
|
||||
"Install a CUDA toolkit or nvidia-cuda-nvcc, or set FLA_TILELANG=0 to disable TileLang explicitly."
|
||||
)
|
||||
return False
|
||||
@@ -0,0 +1,17 @@
|
||||
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# For a list of all contributors, visit:
|
||||
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
|
||||
|
||||
import os
|
||||
|
||||
FLA_CI_ENV = os.getenv("FLA_CI_ENV") == "1"
|
||||
FLA_CACHE_RESULTS = os.getenv('FLA_CACHE_RESULTS', '1') == '1'
|
||||
|
||||
FLA_DISABLE_TENSOR_CACHE = os.getenv('FLA_DISABLE_TENSOR_CACHE', '0') == '1'
|
||||
try:
|
||||
FLA_TENSOR_CACHE_SIZE = int(os.getenv('FLA_TENSOR_CACHE_SIZE', "4"))
|
||||
except ValueError:
|
||||
FLA_TENSOR_CACHE_SIZE = 4
|
||||
@@ -0,0 +1,336 @@
|
||||
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# For a list of all contributors, visit:
|
||||
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
|
||||
|
||||
import contextlib
|
||||
import functools
|
||||
import inspect
|
||||
import sys
|
||||
import warnings
|
||||
from collections import deque
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from packaging import version as package_version
|
||||
|
||||
from .. import __version__
|
||||
from ._config import FLA_DISABLE_TENSOR_CACHE, FLA_TENSOR_CACHE_SIZE
|
||||
from ._device import custom_device_ctx
|
||||
|
||||
|
||||
class Action(Enum):
|
||||
NONE = "none"
|
||||
NOTIFY = "notify"
|
||||
NOTIFY_ALWAYS = "notify_always"
|
||||
RAISE = "raise"
|
||||
|
||||
|
||||
def tensor_cache(
|
||||
fn: Callable[..., torch.Tensor],
|
||||
) -> Callable[..., torch.Tensor]:
|
||||
"""
|
||||
A decorator that memoizes the most recent results of a function call by argument identity.
|
||||
|
||||
The decorator keeps a bounded queue of up to ``FLA_TENSOR_CACHE_SIZE`` (default 4)
|
||||
recent ``(args, kwargs, result)`` triples. On each call, every cached entry is checked
|
||||
in order; an entry is considered a hit when the positional arg count and kwarg key set
|
||||
match and every argument is the *same object* (``is`` identity) as the cached one. On a
|
||||
hit the cached result is returned and ``fn`` is skipped; on a miss ``fn`` is invoked and
|
||||
the new triple is appended (evicting the oldest when the queue is full).
|
||||
|
||||
Caching is fully bypassed when the ``FLA_DISABLE_TENSOR_CACHE`` environment variable is
|
||||
set to ``'1'``.
|
||||
|
||||
Args:
|
||||
fn (Callable[..., torch.Tensor]):
|
||||
The function to be decorated. Intended for functions whose inputs are tensors
|
||||
(or other objects compared by identity) and whose output is a tensor.
|
||||
|
||||
Returns:
|
||||
Callable[..., torch.Tensor]:
|
||||
A wrapped version of ``fn`` backed by an identity-based bounded cache.
|
||||
"""
|
||||
cached: deque = deque(maxlen=FLA_TENSOR_CACHE_SIZE)
|
||||
|
||||
def cache_disabled() -> bool:
|
||||
utils_module = sys.modules.get('kda._fla.utils')
|
||||
return getattr(utils_module, 'FLA_DISABLE_TENSOR_CACHE', FLA_DISABLE_TENSOR_CACHE)
|
||||
|
||||
@functools.wraps(fn)
|
||||
def wrapper(*args: Any, **kwargs: Any) -> Any:
|
||||
if cache_disabled():
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
for cached_args, cached_kwargs, cached_result in cached:
|
||||
if len(args) != len(cached_args) or len(kwargs) != len(cached_kwargs):
|
||||
continue
|
||||
if all(a is b for a, b in zip(args, cached_args, strict=False)) and \
|
||||
all(k in cached_kwargs and v is cached_kwargs[k] for k, v in kwargs.items()):
|
||||
return cached_result
|
||||
|
||||
result = fn(*args, **kwargs)
|
||||
cached.append((args, kwargs, result))
|
||||
return result
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def _skip_contiguous(
|
||||
no_guard_contiguous: bool | list[str] | tuple[str, ...] | set[str],
|
||||
param_name: str,
|
||||
skip_params: set[str],
|
||||
) -> bool:
|
||||
return no_guard_contiguous is True or param_name in skip_params
|
||||
|
||||
|
||||
def _contiguous_if_needed(arg: Any, skip: bool) -> Any:
|
||||
if isinstance(arg, torch.Tensor) and not skip:
|
||||
return arg.contiguous()
|
||||
return arg
|
||||
|
||||
|
||||
def input_guard(
|
||||
fn: Callable[..., torch.Tensor] | None = None,
|
||||
*,
|
||||
no_guard_contiguous: bool | list[str] | tuple[str, ...] | set[str] = False,
|
||||
) -> Callable[[Callable[..., torch.Tensor]], Callable[..., torch.Tensor]] | Callable[..., torch.Tensor]:
|
||||
"""
|
||||
A decorator to make sure all input tensors are contiguous and set the device based on input tensors.
|
||||
|
||||
Args:
|
||||
no_guard_contiguous (bool | list[str] | tuple[str, ...] | set[str]):
|
||||
If True, skip all contiguous checks. If a list/tuple/set of parameter names, skip contiguous check for those parameters.
|
||||
"""
|
||||
|
||||
def decorator(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
|
||||
# Get function signature for parameter name mapping
|
||||
sig = inspect.signature(fn)
|
||||
param_names = list(sig.parameters.keys())
|
||||
skip_params = set(no_guard_contiguous) if isinstance(no_guard_contiguous, (list, tuple, set)) else set()
|
||||
|
||||
@functools.wraps(fn)
|
||||
def wrapper(*args, **kwargs):
|
||||
# Process args with parameter name mapping
|
||||
processed_args = []
|
||||
for i, arg in enumerate(args):
|
||||
if i < len(param_names):
|
||||
param_name = param_names[i]
|
||||
else:
|
||||
# For *args beyond signature, use position as name
|
||||
param_name = f"__arg_{i}"
|
||||
|
||||
processed_args.append(_contiguous_if_needed(
|
||||
arg, _skip_contiguous(no_guard_contiguous, param_name, skip_params)))
|
||||
|
||||
# Process kwargs
|
||||
processed_kwargs = {}
|
||||
for k, v in kwargs.items():
|
||||
processed_kwargs[k] = _contiguous_if_needed(v, _skip_contiguous(no_guard_contiguous, k, skip_params))
|
||||
|
||||
tensor = None
|
||||
for arg in args:
|
||||
if isinstance(arg, torch.Tensor):
|
||||
tensor = arg
|
||||
break
|
||||
if tensor is None:
|
||||
for value in kwargs.values():
|
||||
if isinstance(value, torch.Tensor):
|
||||
tensor = value
|
||||
break
|
||||
|
||||
if tensor is not None:
|
||||
ctx = custom_device_ctx(tensor.device.index)
|
||||
else:
|
||||
ctx = contextlib.nullcontext()
|
||||
|
||||
with ctx:
|
||||
return fn(*processed_args, **processed_kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
# Handle direct usage without parentheses: @input_guard
|
||||
if fn is not None:
|
||||
return decorator(fn)
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def contiguous(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
|
||||
"""Alias for input_guard() without parameters."""
|
||||
return input_guard(fn)
|
||||
|
||||
|
||||
def require_version(version, hint):
|
||||
"""
|
||||
Perform a runtime check of the dependency versions, using the exact same syntax used by pip.
|
||||
"""
|
||||
def decorator(fn):
|
||||
@functools.wraps(fn)
|
||||
def wrapper(ctx, *args, **kwargs):
|
||||
from transformers.utils.versions import require_version
|
||||
require_version(version, hint)
|
||||
return fn(
|
||||
ctx,
|
||||
*(i if not isinstance(i, torch.Tensor) else i.contiguous() for i in args),
|
||||
**{k: (v if not isinstance(v, torch.Tensor) else v.contiguous()) for k, v in kwargs.items()},
|
||||
)
|
||||
return wrapper
|
||||
return decorator
|
||||
|
||||
|
||||
def deprecate_kwarg(
|
||||
old_name: str,
|
||||
version: str,
|
||||
new_name: str | None = None,
|
||||
warn_if_greater_or_equal_version: bool = False,
|
||||
raise_if_greater_or_equal_version: bool = False,
|
||||
raise_if_both_names: bool = False,
|
||||
additional_message: str | None = None,
|
||||
):
|
||||
"""
|
||||
Decorator to notify users about deprecated keyword arguments, replacing them with a new name if specified.
|
||||
|
||||
This decorator allows you to:
|
||||
- Notify users when a keyword argument is deprecated.
|
||||
- Automatically replace deprecated keyword arguments with new ones.
|
||||
- Raise an error if deprecated arguments are used, depending on the specified conditions.
|
||||
|
||||
By default, the decorator notifies the user about the deprecated argument while the `fla.__version__` < specified `version`
|
||||
in the decorator. To keep notifications with any version `warn_if_greater_or_equal_version=True` can be set.
|
||||
|
||||
Args:
|
||||
old_name (`str`):
|
||||
Name of the deprecated keyword argument.
|
||||
version (`str`):
|
||||
The version in which the keyword argument was (or will be) deprecated.
|
||||
new_name (`Optional[str]`, *optional*):
|
||||
The new name for the deprecated keyword argument.
|
||||
If specified, the deprecated keyword argument will be replaced with this new name.
|
||||
warn_if_greater_or_equal_version (`bool`, *optional*, defaults to `False`):
|
||||
Whether to show warning if current `fla` version is greater or equal to the deprecated version.
|
||||
raise_if_greater_or_equal_version (`bool`, *optional*, defaults to `False`):
|
||||
Whether to raise `ValueError` if current `fla` version is greater or equal to the deprecated version.
|
||||
raise_if_both_names (`bool`, *optional*, defaults to `False`):
|
||||
Whether to raise `ValueError` if both deprecated and new keyword arguments are set.
|
||||
additional_message (`Optional[str]`, *optional*):
|
||||
An additional message to append to the default deprecation message.
|
||||
|
||||
Raises:
|
||||
ValueError:
|
||||
If `raise_if_greater_or_equal_version` is `True` and the current version >= the deprecated one,
|
||||
or if `raise_if_both_names` is `True` and both old and new keyword arguments are provided.
|
||||
|
||||
Returns:
|
||||
Callable:
|
||||
A wrapped function that handles the deprecated keyword arguments according to the specified parameters.
|
||||
|
||||
Example usage with renaming argument:
|
||||
|
||||
```python
|
||||
@deprecate_kwarg("reduce_labels", new_name="do_reduce_labels", version="6.0.0")
|
||||
def my_function(do_reduce_labels):
|
||||
print(do_reduce_labels)
|
||||
|
||||
my_function(reduce_labels=True) # Will show a deprecation warning and use do_reduce_labels=True
|
||||
```
|
||||
|
||||
Example usage without renaming argument:
|
||||
|
||||
```python
|
||||
@deprecate_kwarg("max_size", version="6.0.0")
|
||||
def my_function(max_size):
|
||||
print(max_size)
|
||||
|
||||
my_function(max_size=1333) # Will show a deprecation warning
|
||||
```
|
||||
|
||||
"""
|
||||
deprecated_version = package_version.parse(version)
|
||||
current_version = package_version.parse(__version__)
|
||||
is_greater_or_equal_version = current_version >= deprecated_version
|
||||
|
||||
if is_greater_or_equal_version:
|
||||
version_message = f"and removed starting from version {version}"
|
||||
else:
|
||||
version_message = f"and will be removed in version {version}"
|
||||
|
||||
def wrapper(func):
|
||||
# Required for better warning message
|
||||
sig = inspect.signature(func)
|
||||
function_named_args = set(sig.parameters.keys())
|
||||
is_instance_method = "self" in function_named_args
|
||||
is_class_method = "cls" in function_named_args
|
||||
|
||||
@functools.wraps(func)
|
||||
def wrapped_func(*args, **kwargs):
|
||||
# Get class + function name (just for better warning message)
|
||||
func_name = func.__name__
|
||||
if is_instance_method:
|
||||
func_name = f"{args[0].__class__.__name__}.{func_name}"
|
||||
elif is_class_method:
|
||||
func_name = f"{args[0].__name__}.{func_name}"
|
||||
|
||||
minimum_action = Action.NONE
|
||||
message = None
|
||||
|
||||
# deprecated kwarg and its new version are set for function call -> replace it with new name
|
||||
if old_name in kwargs and new_name in kwargs:
|
||||
minimum_action = Action.RAISE if raise_if_both_names else Action.NOTIFY_ALWAYS
|
||||
message = (
|
||||
f"Both `{old_name}` and `{new_name}` are set for `{func_name}`. "
|
||||
f"Using `{new_name}={kwargs[new_name]}` and ignoring deprecated `{old_name}={kwargs[old_name]}`."
|
||||
)
|
||||
kwargs.pop(old_name)
|
||||
|
||||
# only deprecated kwarg is set for function call -> replace it with new name
|
||||
elif old_name in kwargs and new_name is not None and new_name not in kwargs:
|
||||
minimum_action = Action.NOTIFY
|
||||
message = (
|
||||
f"`{old_name}` is deprecated {version_message} for `{func_name}`. "
|
||||
f"Use `{new_name}` instead."
|
||||
)
|
||||
kwargs[new_name] = kwargs.pop(old_name)
|
||||
|
||||
# deprecated kwarg is not set for function call and new name is not specified -> just notify
|
||||
elif old_name in kwargs:
|
||||
minimum_action = Action.NOTIFY
|
||||
message = f"`{old_name}` is deprecated {version_message} for `{func_name}`."
|
||||
|
||||
if message is not None and additional_message is not None:
|
||||
message = f"{message} {additional_message}"
|
||||
|
||||
# update minimum_action if argument is ALREADY deprecated (current version >= deprecated version)
|
||||
if is_greater_or_equal_version:
|
||||
# change to (NOTIFY, NOTIFY_ALWAYS) -> RAISE if specified
|
||||
# in case we want to raise error for already deprecated arguments
|
||||
if raise_if_greater_or_equal_version and minimum_action != Action.NONE:
|
||||
minimum_action = Action.RAISE
|
||||
|
||||
# change to NOTIFY -> NONE if specified (NOTIFY_ALWAYS can't be changed to NONE)
|
||||
elif not warn_if_greater_or_equal_version and minimum_action == Action.NOTIFY:
|
||||
minimum_action = Action.NONE
|
||||
|
||||
# raise error or notify user
|
||||
if minimum_action == Action.RAISE:
|
||||
raise ValueError(message)
|
||||
elif minimum_action in (Action.NOTIFY, Action.NOTIFY_ALWAYS):
|
||||
# DeprecationWarning is ignored by default, so we use FutureWarning instead
|
||||
warnings.warn(message, FutureWarning, stacklevel=2)
|
||||
|
||||
return func(*args, **kwargs)
|
||||
|
||||
return wrapped_func
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def checkpoint(fn):
|
||||
@functools.wraps(fn)
|
||||
def wrapper(*args, **kwargs):
|
||||
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs)
|
||||
return wrapper
|
||||
@@ -0,0 +1,245 @@
|
||||
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# For a list of all contributors, visit:
|
||||
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
|
||||
|
||||
import contextlib
|
||||
import functools
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import sys
|
||||
import warnings
|
||||
from enum import Enum
|
||||
from functools import cache, lru_cache
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from packaging import version as package_version
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def check_environments():
|
||||
"""
|
||||
Checks the current operating system, Triton version, and Python version,
|
||||
issuing warnings if they don't meet recommendations.
|
||||
This function's body only runs once due to lru_cache.
|
||||
"""
|
||||
# Check Operating System
|
||||
if sys.platform == 'win32':
|
||||
# Check if triton-windows is installed
|
||||
try:
|
||||
from importlib.metadata import PackageNotFoundError, metadata
|
||||
metadata('triton-windows')
|
||||
# triton-windows is installed, no warning needed
|
||||
except PackageNotFoundError:
|
||||
logger.warning(
|
||||
"Detected Windows operating system. Consider installing triton-windows "
|
||||
"(https://github.com/triton-lang/triton-windows) for better compatibility. "
|
||||
"Without it, some features may not work correctly.",
|
||||
)
|
||||
|
||||
triton_version = package_version.parse(triton.__version__)
|
||||
required_triton_version = package_version.parse("3.3.0")
|
||||
|
||||
if triton_version < required_triton_version:
|
||||
logger.warning(
|
||||
f"Current Triton version {triton_version} is below the recommended 3.3.0 version. "
|
||||
"Errors may occur and these issues will not be fixed. "
|
||||
"Please consider upgrading Triton.",
|
||||
)
|
||||
|
||||
# Check Python version
|
||||
py_version = package_version.parse(f"{sys.version_info.major}.{sys.version_info.minor}")
|
||||
required_py_version = package_version.parse("3.11")
|
||||
|
||||
if py_version < required_py_version:
|
||||
logger.warning(
|
||||
f"Current Python version {py_version} is below the recommended 3.11 version. "
|
||||
"It is recommended to upgrade to Python 3.11 or higher for the best experience.",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
check_environments()
|
||||
|
||||
|
||||
def _cpu_device_warning():
|
||||
warnings.warn(('Triton is not supported on current platform, roll back to CPU.'), stacklevel=2)
|
||||
|
||||
|
||||
@cache
|
||||
def check_pytorch_version(version_s: str = '2.4') -> bool:
|
||||
return package_version.parse(torch.__version__) >= package_version.parse(version_s)
|
||||
|
||||
|
||||
@cache
|
||||
def get_multiprocessor_count(tensor_idx: int = 0, *, use_aicore: bool = False) -> int:
|
||||
try:
|
||||
return triton.runtime.driver.active.utils.get_device_properties(tensor_idx)['multiprocessor_count']
|
||||
except Exception:
|
||||
# Maybe we use a NPU device.
|
||||
try:
|
||||
if triton.runtime.driver.active.get_current_target().backend == 'npu':
|
||||
props = triton.runtime.driver.active.utils.get_device_properties(tensor_idx)
|
||||
return props['num_aicore'] if use_aicore else props['num_vectorcore']
|
||||
except Exception:
|
||||
logger.debug('Failed to get NPU multiprocessor count, falling back to 1.', exc_info=True)
|
||||
return 1
|
||||
|
||||
|
||||
@cache
|
||||
def get_device_capability(device_index: int = 0) -> tuple[int, int]:
|
||||
major, minor = torch.cuda.get_device_capability(device_index)
|
||||
return int(major), int(minor)
|
||||
|
||||
|
||||
@cache
|
||||
def get_device_smem_optin(device_index: int = 0) -> int:
|
||||
props = torch.cuda.get_device_properties(device_index)
|
||||
return int(getattr(props, 'shared_memory_per_block_optin', props.shared_memory_per_block))
|
||||
|
||||
|
||||
@cache
|
||||
def get_available_device() -> str:
|
||||
try:
|
||||
return triton.runtime.driver.active.get_current_target().backend
|
||||
except Exception:
|
||||
_cpu_device_warning()
|
||||
return 'cpu'
|
||||
|
||||
|
||||
def map_triton_backend_to_torch_device() -> str:
|
||||
backend = get_available_device() # 'cuda' | 'hip' | 'xpu' | 'cpu' | ...
|
||||
return {'cuda': 'cuda', 'hip': 'cuda', 'xpu': 'xpu'}.get(backend, backend)
|
||||
|
||||
|
||||
# For AMD GPUs, the triton backend is 'hip', while for Nvidia GPUs, the triton backend is 'cuda'.
|
||||
# However, the torch backend is 'cuda' for both Nvidia and AMD GPUs.
|
||||
# Therefore, we need to check the triton backend to determine the actual GPU vendor.
|
||||
device = get_available_device() if get_available_device() != 'hip' else 'cuda'
|
||||
device_torch_lib = getattr(torch, device)
|
||||
device_platform = get_available_device()
|
||||
device_name = map_triton_backend_to_torch_device()
|
||||
|
||||
IS_AMD = (device_platform == 'hip')
|
||||
|
||||
IS_ARM = platform.machine().lower() in ('aarch64', 'arm64')
|
||||
|
||||
IS_INTEL = (device_platform == 'xpu')
|
||||
IS_INTEL_ALCHEMIST = (IS_INTEL and 'Intel(R) Arc(TM) A' in torch.xpu.get_device_name(0))
|
||||
|
||||
IS_NPU = (device_platform == 'npu')
|
||||
|
||||
IS_NVIDIA = (device_platform == 'cuda')
|
||||
IS_NVIDIA_HOPPER = (
|
||||
IS_NVIDIA and (
|
||||
'NVIDIA H' in torch.cuda.get_device_name(0)
|
||||
or torch.cuda.get_device_capability()[0] == 9
|
||||
)
|
||||
)
|
||||
IS_NVIDIA_SM100 = (IS_NVIDIA and torch.cuda.get_device_capability()[0] == 10)
|
||||
# NOTE: exactly 12.0 — 12.1 (GB10) is a different target that FlashQLA rejects at import time.
|
||||
IS_NVIDIA_SM120 = (IS_NVIDIA and torch.cuda.get_device_capability() == (12, 0))
|
||||
IS_NVIDIA_BLACKWELL = (IS_NVIDIA and torch.cuda.get_device_capability()[0] in (10, 12))
|
||||
|
||||
# Nvidia Ampere or newer, haven't check AMD and intel yet.
|
||||
IS_TF32_SUPPORTED = (IS_NVIDIA and torch.cuda.get_device_capability(0)[0] >= 8)
|
||||
IS_GATHER_SUPPORTED = hasattr(triton.language, 'gather')
|
||||
IS_TMA_SUPPORTED = (
|
||||
IS_NVIDIA
|
||||
and torch.cuda.get_device_capability(0)[0] >= 9
|
||||
and os.environ.get('FLA_USE_TMA', '0') == '1'
|
||||
and (
|
||||
hasattr(triton.language, '_experimental_make_tensor_descriptor')
|
||||
or hasattr(triton.language, 'make_tensor_descriptor')
|
||||
)
|
||||
)
|
||||
|
||||
if IS_NVIDIA and not IS_TF32_SUPPORTED:
|
||||
# Make old card happy, since triton will use tf32 by default.
|
||||
# This is a workaround for old nvidia card.
|
||||
os.environ['TRITON_F32_DEFAULT'] = 'ieee'
|
||||
|
||||
|
||||
def _default_alloc_fn(size: int, alignment: int, stream: int | None):
|
||||
return torch.empty(size, device=torch.device(device_name, device_torch_lib.current_device()), dtype=torch.int8)
|
||||
|
||||
|
||||
if IS_TMA_SUPPORTED:
|
||||
logger.info('TMA is supported, using TMA by default.')
|
||||
triton.set_allocator(_default_alloc_fn)
|
||||
elif IS_NVIDIA_BLACKWELL:
|
||||
# Blackwell (SM100 datacenter / SM120 consumer): Triton compiler may emit global_scratch for
|
||||
# autotuned kernels even without TMA. Register a default allocator to
|
||||
# prevent NullAllocator crashes. See triton-lang/triton#10002.
|
||||
logger.info('Blackwell detected: registering default global_scratch allocator.')
|
||||
triton.set_allocator(_default_alloc_fn)
|
||||
|
||||
|
||||
def get_all_max_shared_mem():
|
||||
try:
|
||||
return [
|
||||
triton.runtime.driver.active.utils.get_device_properties(i)['max_shared_mem']
|
||||
for i in range(device_torch_lib.device_count())
|
||||
]
|
||||
except Exception:
|
||||
_cpu_device_warning()
|
||||
return [-1]
|
||||
|
||||
|
||||
class Backend(Enum):
|
||||
ADA = 101376 # RTX 4090
|
||||
AMPERE = 166912 # A100
|
||||
HOPPER = 232448 # H100
|
||||
DEFAULT = 102400 # Default
|
||||
|
||||
@classmethod
|
||||
def get_shared_memory(cls, arch: str) -> int:
|
||||
try:
|
||||
return cls[arch.upper()].value
|
||||
except KeyError:
|
||||
return cls.DEFAULT.value
|
||||
|
||||
|
||||
@cache
|
||||
def check_shared_mem(arch: str = "none", tensor_idx: int = 0) -> bool:
|
||||
try:
|
||||
device_shared_mem_list = get_all_max_shared_mem()
|
||||
max_shared_memory = device_shared_mem_list[tensor_idx]
|
||||
return max_shared_memory >= Backend.get_shared_memory(arch)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
if check_pytorch_version('2.4'):
|
||||
if device == 'cpu':
|
||||
device = 'cuda'
|
||||
device_torch_lib = getattr(torch, device)
|
||||
autocast_custom_fwd = functools.partial(torch.amp.custom_fwd, device_type=device)
|
||||
autocast_custom_bwd = functools.partial(torch.amp.custom_bwd, device_type=device)
|
||||
|
||||
def custom_device_ctx(index: int):
|
||||
if index is None:
|
||||
return contextlib.nullcontext()
|
||||
try:
|
||||
return device_torch_lib.device(index)
|
||||
except (AttributeError, AssertionError, RuntimeError):
|
||||
return contextlib.nullcontext()
|
||||
else:
|
||||
assert device == 'cuda', 'Only cuda device is supported for PyTorch version < 2.4.0.'
|
||||
autocast_custom_fwd = device_torch_lib.amp.custom_fwd
|
||||
autocast_custom_bwd = device_torch_lib.amp.custom_bwd
|
||||
|
||||
def custom_device_ctx(index: int):
|
||||
if index is None:
|
||||
return contextlib.nullcontext()
|
||||
try:
|
||||
return torch.cuda.device(index)
|
||||
except (AttributeError, AssertionError, RuntimeError):
|
||||
return contextlib.nullcontext()
|
||||
@@ -0,0 +1,41 @@
|
||||
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# For a list of all contributors, visit:
|
||||
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
|
||||
|
||||
import logging
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
|
||||
from ._config import FLA_CI_ENV
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def get_abs_err(x, y):
|
||||
return (x.detach() - y.detach()).flatten().abs().max().item()
|
||||
|
||||
|
||||
def get_err_ratio(x, y):
|
||||
err = (x.detach() - y.detach()).flatten().square().mean().sqrt().item()
|
||||
base = (x.detach()).flatten().square().mean().sqrt().item()
|
||||
return err / (base + 1e-8)
|
||||
|
||||
|
||||
def assert_close(prefix, ref, tri, ratio, warning=False, err_atol=1e-6):
|
||||
abs_atol = get_abs_err(ref, tri)
|
||||
error_rate = get_err_ratio(ref, tri)
|
||||
msg = f"{prefix:>16} diff: {abs_atol:.6f} ratio: {error_rate:.6f}"
|
||||
logger.info(msg)
|
||||
if abs_atol <= err_atol:
|
||||
return
|
||||
assert not torch.isnan(ref).any(), f"{prefix}: NaN detected in ref"
|
||||
assert not torch.isnan(tri).any(), f"{prefix}: NaN detected in tri"
|
||||
if warning or (FLA_CI_ENV and (error_rate < 0.01 or abs_atol <= 0.3)):
|
||||
if error_rate > ratio:
|
||||
warnings.warn(msg)
|
||||
else:
|
||||
assert error_rate < ratio, msg
|
||||
Reference in New Issue
Block a user