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,14 @@
|
||||
from .cumsum import (
|
||||
chunk_local_cumsum,
|
||||
chunk_local_cumsum_scalar,
|
||||
chunk_local_cumsum_vector,
|
||||
)
|
||||
from .index import prepare_chunk_indices, prepare_chunk_offsets
|
||||
|
||||
__all__ = [
|
||||
"chunk_local_cumsum",
|
||||
"chunk_local_cumsum_scalar",
|
||||
"chunk_local_cumsum_vector",
|
||||
"prepare_chunk_indices",
|
||||
"prepare_chunk_offsets",
|
||||
]
|
||||
@@ -0,0 +1,449 @@
|
||||
# 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 dataclasses
|
||||
import enum
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from functools import cache, lru_cache
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from packaging import version
|
||||
from triton.runtime.autotuner import Autotuner
|
||||
|
||||
TRITON_ABOVE_3_5_1 = version.parse(triton.__version__) >= version.parse("3.5.1")
|
||||
TRITON_ABOVE_3_4_0 = version.parse(triton.__version__) >= version.parse("3.4.0")
|
||||
|
||||
|
||||
class FlaCacheMode(enum.Enum):
|
||||
"""Controls how FLA loads kernel configs from its config cache (FLA_CACHE_MODE env var).
|
||||
|
||||
DISABLED — skip all cache lookups, always fall back to Triton autotune (default when FLA_CACHE_MODE is unset)
|
||||
STRICT — exact key match only; falls back to Triton autotune if no match
|
||||
FUZZY — exact key match → fuzzy key match; falls back to Triton autotune if no match
|
||||
FULL — exact key match → fuzzy key match → default_config fallback
|
||||
DEFAULT — use only the top-level default_config field, skip key-based lookup
|
||||
ALWAYS — like DEFAULT, but re-reads config files on every kernel call;
|
||||
useful for debugging: edit default_config in a JSON file and the next
|
||||
kernel call picks it up without restarting the process
|
||||
"""
|
||||
DISABLED = "disabled"
|
||||
STRICT = "strict"
|
||||
FUZZY = "fuzzy"
|
||||
FULL = "full"
|
||||
DEFAULT = "default"
|
||||
ALWAYS = "always"
|
||||
|
||||
def uses_default_config(self) -> bool:
|
||||
"""Return True for modes that may fall back to default_config (FULL, DEFAULT, ALWAYS)."""
|
||||
return self in (FlaCacheMode.FULL, FlaCacheMode.DEFAULT, FlaCacheMode.ALWAYS)
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> "FlaCacheMode":
|
||||
mode_str = os.environ.get("FLA_CACHE_MODE", cls.DISABLED.value)
|
||||
try:
|
||||
return cls(mode_str)
|
||||
except ValueError:
|
||||
valid = [m.value for m in cls]
|
||||
raise ValueError(
|
||||
f"Invalid FLA_CACHE_MODE={mode_str!r}. Valid values: {valid}"
|
||||
) from None
|
||||
|
||||
|
||||
FLA_CACHE_MODE: FlaCacheMode = FlaCacheMode.from_env()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def sanitize_gpu_name(gpu_name: str) -> str:
|
||||
sanitized = re.sub(r"[^0-9A-Za-z]+", "_", gpu_name)
|
||||
sanitized = sanitized.strip("_")
|
||||
return sanitized or "unknown_gpu"
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_gpu_info():
|
||||
"""Get GPU model information.
|
||||
|
||||
This function detects the GPU model and returns a sanitized string identifier.
|
||||
It prioritizes FLA_GPU_NAME environment variable if set, then detects from
|
||||
available hardware (CUDA, ROCm, Intel GPU, or CPU).
|
||||
"""
|
||||
# Check if GPU name is overridden via environment variable
|
||||
gpu_name = None
|
||||
# Check if GPU name is overridden via environment variable
|
||||
if "FLA_GPU_NAME" in os.environ:
|
||||
gpu_name = os.environ["FLA_GPU_NAME"]
|
||||
# Try to get device name based on availability
|
||||
elif torch.cuda.is_available():
|
||||
# Works for both NVIDIA and AMD GPUs (ROCm)
|
||||
gpu_name = torch.cuda.get_device_name(0)
|
||||
elif hasattr(torch, 'xpu') and torch.xpu.is_available():
|
||||
gpu_name = torch.xpu.get_device_name(0)
|
||||
|
||||
if gpu_name:
|
||||
return sanitize_gpu_name(gpu_name)
|
||||
|
||||
# Default to CPU if no GPU available
|
||||
return "cpu"
|
||||
|
||||
|
||||
def get_fla_config_dir() -> Path:
|
||||
"""Get FLA's configs directory.
|
||||
|
||||
The directory can be overridden by setting the FLA_CONFIG_DIR environment variable.
|
||||
If set, configs will be loaded directly from $FLA_CONFIG_DIR/. Otherwise FLA
|
||||
falls back to the default fla/configs/{GPU}/ directory in the project.
|
||||
"""
|
||||
# Check if custom config dir is set via environment variable
|
||||
if "FLA_CONFIG_DIR" in os.environ:
|
||||
return Path(os.environ["FLA_CONFIG_DIR"])
|
||||
|
||||
# Default: project_dir/fla/configs/{GPU}/
|
||||
project_dir = Path(__file__).parent.parent.parent
|
||||
return project_dir / "configs" / get_gpu_info()
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class AutotuneKey:
|
||||
"""Autotune key with exact/fuzzy matching, serialization, and construction helpers."""
|
||||
autotune_key: tuple[Any, ...]
|
||||
|
||||
@staticmethod
|
||||
def normalize_autotune_key(value: Any) -> Any:
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [AutotuneKey.normalize_autotune_key(v) for v in value]
|
||||
if isinstance(value, dict):
|
||||
return {k: AutotuneKey.normalize_autotune_key(v) for k, v in value.items()}
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def serialize(key: Any) -> str:
|
||||
return json.dumps(AutotuneKey.normalize_autotune_key(key), separators=(",", ":"), sort_keys=True)
|
||||
|
||||
@staticmethod
|
||||
def key_hash(key: Any) -> str:
|
||||
import hashlib
|
||||
return hashlib.md5(AutotuneKey.serialize(key).encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def is_numeric(value: Any) -> bool:
|
||||
return isinstance(value, (int, float)) and not isinstance(value, bool)
|
||||
|
||||
@staticmethod
|
||||
def keys_fuzzy_match(cached_key: Any, requested_key: Any) -> bool:
|
||||
# Fuzzy match: numeric leaves are compatible regardless of their actual numeric values
|
||||
# (e.g. a config tuned for seq_len=1024 can apply to seq_len=2048).
|
||||
# Structure (type, length, dict keys) must still match exactly.
|
||||
if AutotuneKey.is_numeric(cached_key) and AutotuneKey.is_numeric(requested_key):
|
||||
return True
|
||||
if isinstance(cached_key, (list, tuple)) and isinstance(requested_key, (list, tuple)):
|
||||
return len(cached_key) == len(requested_key) and all(
|
||||
AutotuneKey.keys_fuzzy_match(c, r) for c, r in zip(cached_key, requested_key)
|
||||
)
|
||||
if isinstance(cached_key, dict) and isinstance(requested_key, dict):
|
||||
return cached_key.keys() == requested_key.keys() and all(
|
||||
AutotuneKey.keys_fuzzy_match(cached_key[k], requested_key[k]) for k in cached_key
|
||||
)
|
||||
return cached_key == requested_key
|
||||
|
||||
@classmethod
|
||||
def build(
|
||||
cls,
|
||||
arg_names: list[str],
|
||||
key_names: list[str],
|
||||
positional_args: tuple[Any, ...],
|
||||
runtime_kwargs: dict[str, Any],
|
||||
) -> "AutotuneKey":
|
||||
named_args = dict(zip(arg_names, positional_args))
|
||||
all_args = {**named_args, **runtime_kwargs}
|
||||
tracked_args = {k: v for (k, v) in all_args.items() if k in arg_names}
|
||||
tuning_key = [tracked_args[name] for name in key_names if name in tracked_args]
|
||||
for arg in tracked_args.values():
|
||||
if hasattr(arg, "dtype"):
|
||||
tuning_key.append(str(arg.dtype))
|
||||
return cls(autotune_key=tuple(tuning_key))
|
||||
|
||||
def exact_matches(self, entry_key: Any) -> bool:
|
||||
return self.serialize(self.autotune_key) == self.serialize(entry_key)
|
||||
|
||||
def fuzzy_matches(self, entry_key: Any) -> bool:
|
||||
self_normalized = self.normalize_autotune_key(self.autotune_key)
|
||||
entry_normalized = self.normalize_autotune_key(entry_key)
|
||||
return (
|
||||
isinstance(self_normalized, list)
|
||||
and isinstance(entry_normalized, list)
|
||||
and len(self_normalized) == len(entry_normalized)
|
||||
and AutotuneKey.keys_fuzzy_match(self_normalized, entry_normalized)
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class KernelConfigFile:
|
||||
"""Validated in-memory representation of a {kernel_name}.json config file."""
|
||||
kernel_name: str | None
|
||||
triton_version: str | None
|
||||
autotune_entries: dict[str, dict[str, Any]] | None
|
||||
default_config: dict[str, Any] | None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config_file: Path, data: Any) -> "KernelConfigFile | None":
|
||||
"""Parse and validate a raw JSON dict. Returns None (with a warning) if malformed."""
|
||||
def fail(msg, *args):
|
||||
logger.warning(msg, *args)
|
||||
raise ValueError
|
||||
|
||||
try:
|
||||
if not isinstance(data, dict):
|
||||
fail("Malformed config %s: root is %s, expected dict", config_file, type(data).__name__)
|
||||
raw_entries = data.get("autotune_entries")
|
||||
entries: dict[str, dict[str, Any]] | None = None
|
||||
if raw_entries is not None:
|
||||
if not isinstance(raw_entries, dict):
|
||||
fail("Malformed config %s: 'autotune_entries' is %s, expected dict",
|
||||
config_file, type(raw_entries).__name__)
|
||||
for h, entry in raw_entries.items():
|
||||
if not isinstance(entry, dict):
|
||||
fail("Malformed config %s: autotune_entries[%r] is %s, expected dict",
|
||||
config_file, h, type(entry).__name__)
|
||||
if not isinstance(entry.get("config"), dict):
|
||||
fail("Malformed config %s: autotune_entries[%r] missing valid 'config' field", config_file, h)
|
||||
entries = raw_entries
|
||||
default_config = data.get("default_config")
|
||||
if default_config is not None and not isinstance(default_config, dict):
|
||||
fail("Malformed config %s: 'default_config' is %s, expected dict", config_file, type(default_config).__name__)
|
||||
return cls(
|
||||
kernel_name=data.get("kernel_name"),
|
||||
triton_version=data.get("triton_version"),
|
||||
autotune_entries=entries,
|
||||
default_config=default_config,
|
||||
)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def from_file(cls, config_file: Path) -> "KernelConfigFile | None":
|
||||
"""Read and validate a config file. Returns None if the file is missing or malformed."""
|
||||
config_data = read_config_file(config_file)
|
||||
if config_data is None:
|
||||
return None
|
||||
return cls.from_dict(config_file, config_data)
|
||||
|
||||
def lookup_exact(self, key: AutotuneKey) -> dict[str, Any] | None:
|
||||
if self.autotune_entries is None:
|
||||
return None
|
||||
return self.autotune_entries.get(AutotuneKey.key_hash(key.autotune_key))
|
||||
|
||||
def lookup_fuzzy(self, key: AutotuneKey) -> dict[str, Any] | None:
|
||||
if self.autotune_entries is None:
|
||||
return None
|
||||
for entry in self.autotune_entries.values():
|
||||
if key.fuzzy_matches(entry.get("autotune_key")):
|
||||
return entry
|
||||
return None
|
||||
|
||||
|
||||
@cache
|
||||
def load_config_file(config_file: Path) -> dict[str, Any] | None:
|
||||
try:
|
||||
with open(config_file) as f:
|
||||
return json.load(f)
|
||||
except Exception as e:
|
||||
logger.warning("Error reading config file %s: %s", config_file, e)
|
||||
return None
|
||||
|
||||
|
||||
def read_config_file(config_file: Path) -> dict[str, Any] | None:
|
||||
"""Read a config file, bypassing the in-process cache in ALWAYS mode."""
|
||||
if FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
|
||||
return load_config_file.__wrapped__(config_file)
|
||||
return load_config_file(config_file)
|
||||
|
||||
|
||||
def load_cached_config(kernel_name: str, autotune_key: AutotuneKey | None = None) -> dict[str, Any] | None:
|
||||
"""
|
||||
Load cached best config for a kernel from FLA configs directory.
|
||||
|
||||
This function loads the cached best configuration for a given kernel name
|
||||
from get_fla_config_dir()/{kernel_name}.json.
|
||||
|
||||
Cache files may contain multiple autotune entries keyed by Triton's
|
||||
runtime tuning key plus a top-level default config.
|
||||
|
||||
If the config file is not found or cannot be loaded, a warning is printed
|
||||
and None is returned, allowing fallback to Triton's autotune.
|
||||
|
||||
The lookup mode is controlled by the FLA_CACHE_MODE environment variable (see FlaCacheMode).
|
||||
|
||||
Args:
|
||||
kernel_name: Name of the kernel (e.g., "causal_conv1d_fwd_kernel")
|
||||
autotune_key: Triton autotune key for the current invocation
|
||||
|
||||
Returns:
|
||||
Best config dictionary or None if not found or disabled
|
||||
"""
|
||||
if FLA_CACHE_MODE is FlaCacheMode.DISABLED:
|
||||
return None
|
||||
|
||||
config_dir = get_fla_config_dir()
|
||||
config_file = config_dir / f"{kernel_name}.json"
|
||||
|
||||
if not config_file.exists():
|
||||
return None
|
||||
|
||||
config_data = read_config_file(config_file)
|
||||
if config_data is None:
|
||||
return None
|
||||
config = KernelConfigFile.from_dict(config_file, config_data)
|
||||
if config is None:
|
||||
return None
|
||||
|
||||
if FLA_CACHE_MODE is FlaCacheMode.DEFAULT or FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
|
||||
return config.default_config
|
||||
|
||||
# STRICT mode: exact match only, no fuzzy fallback
|
||||
if FLA_CACHE_MODE is FlaCacheMode.STRICT:
|
||||
if autotune_key is not None:
|
||||
entry = config.lookup_exact(autotune_key)
|
||||
if entry is not None:
|
||||
return entry["config"]
|
||||
return None
|
||||
|
||||
# FULL and FUZZY modes: try exact key match first, then fuzzy match
|
||||
if autotune_key is not None:
|
||||
entry = config.lookup_exact(autotune_key) or config.lookup_fuzzy(autotune_key)
|
||||
if entry is not None:
|
||||
return entry["config"]
|
||||
|
||||
if FLA_CACHE_MODE is FlaCacheMode.FUZZY:
|
||||
return None
|
||||
|
||||
# FULL mode: fall back to default_config, then legacy raw config (no autotune_entries)
|
||||
if config.default_config is not None:
|
||||
return config.default_config
|
||||
if config.autotune_entries is not None:
|
||||
return None
|
||||
return config_data
|
||||
|
||||
|
||||
class CachedAutotuner(Autotuner):
|
||||
"""
|
||||
A modified autotuner that loads best config from FLA's config directory.
|
||||
|
||||
This class extends Triton's Autotuner but overrides the run method to
|
||||
try loading cached configuration first before falling back to autotune.
|
||||
"""
|
||||
|
||||
def __init__(self, fn, arg_names, configs, key, reset_to_zero, restore_value, **kwargs):
|
||||
super().__init__(fn, arg_names, configs, key, reset_to_zero, restore_value, **kwargs)
|
||||
self.kernel_name = fn.fn.__name__ if hasattr(fn, 'fn') else fn.__name__
|
||||
|
||||
# None-safe pre/post hooks: Triton's defaults crash when a restore_value / reset_to_zero arg
|
||||
# is None (idiomatic for optional pointers gated by a tl.constexpr flag).
|
||||
# Fixed upstream in triton-lang/triton#10295 — remove this override once FLA's minimum Triton version has it.
|
||||
if not self.user_defined_pre_hook and (self.reset_to_zero or self.restore_value):
|
||||
def _pre_hook(kw, reset_only=False):
|
||||
for n in self.reset_to_zero:
|
||||
if kw[n] is not None:
|
||||
kw[n].zero_()
|
||||
if not reset_only:
|
||||
self.restore_copies = {n: kw[n].clone() for n in self.restore_value if kw[n] is not None}
|
||||
self.pre_hook = _pre_hook
|
||||
if not self.user_defined_post_hook and self.restore_value:
|
||||
def _post_hook(kw, exception):
|
||||
for n, copy in self.restore_copies.items():
|
||||
kw[n].copy_(copy)
|
||||
self.restore_copies = {}
|
||||
self.post_hook = _post_hook
|
||||
|
||||
def should_check_fla_cache(self, key: AutotuneKey) -> bool:
|
||||
if FLA_CACHE_MODE is FlaCacheMode.DISABLED:
|
||||
return False
|
||||
if FLA_CACHE_MODE is FlaCacheMode.ALWAYS:
|
||||
return True
|
||||
return key.autotune_key not in self.cache
|
||||
|
||||
def run(self, *args, **kwargs):
|
||||
key = AutotuneKey.build(self.arg_names, self.keys, args, kwargs)
|
||||
if self.should_check_fla_cache(key):
|
||||
self.maybe_load_cached_config(key)
|
||||
return super().run(*args, **kwargs)
|
||||
|
||||
def maybe_load_cached_config(self, key: AutotuneKey):
|
||||
best_config = load_cached_config(self.kernel_name, key)
|
||||
|
||||
if best_config is not None:
|
||||
kw = best_config["kwargs"]
|
||||
num_warps = best_config["num_warps"]
|
||||
num_stages = best_config["num_stages"]
|
||||
|
||||
extra = {
|
||||
"num_ctas": best_config["num_ctas"],
|
||||
"maxnreg": best_config.get("maxnreg"),
|
||||
"pre_hook": None,
|
||||
"ir_override": best_config.get("ir_override"),
|
||||
} if TRITON_ABOVE_3_5_1 else {}
|
||||
cfg = triton.Config(kw, num_warps=num_warps, num_stages=num_stages, **extra)
|
||||
|
||||
self.cache[key.autotune_key] = cfg
|
||||
else:
|
||||
logger.debug(
|
||||
"No cached config found for kernel %s and key %s; falling back to Triton autotune",
|
||||
self.kernel_name,
|
||||
list(key.autotune_key),
|
||||
)
|
||||
|
||||
|
||||
def fla_cache_autotune(configs, key=None, prune_configs_by=None, reset_to_zero=None, restore_value=None,
|
||||
pre_hook=None, post_hook=None, warmup=None, rep=None, use_cuda_graph=False,
|
||||
do_bench=None, cache_results=False):
|
||||
"""
|
||||
Decorator for auto-tuning a :code:`triton.jit`'d function with FLA config support.
|
||||
|
||||
Extends Triton's autotune to load best configurations from FLA's config directory
|
||||
(default: fla/configs/{GPU}/, or FLA_CONFIG_DIR/ when overridden), keyed by kernel
|
||||
name from {kernel_name}.json. Lookup behaviour is controlled by FLA_CACHE_MODE.
|
||||
Falls back to normal Triton autotuning when no cached config is found.
|
||||
"""
|
||||
# key can be None when we want to use cache only (no fallback autotune)
|
||||
if key is None:
|
||||
key = []
|
||||
|
||||
def decorator(fn):
|
||||
kwargs = {}
|
||||
if TRITON_ABOVE_3_4_0:
|
||||
kwargs = {"cache_results": cache_results}
|
||||
|
||||
return CachedAutotuner(fn, fn.arg_names, configs, key, reset_to_zero, restore_value,
|
||||
pre_hook=pre_hook, post_hook=post_hook,
|
||||
prune_configs_by=prune_configs_by, warmup=warmup, rep=rep,
|
||||
use_cuda_graph=use_cuda_graph, do_bench=do_bench,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def configure_fla_cache_autotune():
|
||||
triton.autotune = fla_cache_autotune
|
||||
logger.info(
|
||||
"configure_fla_cache_autotune() is enabling FLA fla_cache_autotune; "
|
||||
"triton.autotune will be replaced with fla_cache_autotune."
|
||||
)
|
||||
|
||||
|
||||
def restore_autotune_backend():
|
||||
from triton.runtime.autotuner import autotune as original_autotune
|
||||
triton.autotune = original_autotune
|
||||
logger.info(
|
||||
"restore_autotune_backend() is restoring Triton's original autotune; "
|
||||
"triton.autotune will be replaced with triton.runtime.autotuner.autotune."
|
||||
)
|
||||
@@ -0,0 +1,10 @@
|
||||
# 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
|
||||
|
||||
# Approximate value of 1/ln(2), used for log/exp base conversion
|
||||
# Best FP32 approximation: 1.4426950216 (hex 0x3FB8AA3B)
|
||||
RCP_LN2 = 1.4426950216
|
||||
@@ -0,0 +1,468 @@
|
||||
# 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 torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from kda._fla.ops.backends import dispatch
|
||||
from kda._fla.ops.utils.cache import fla_cache_autotune
|
||||
from kda._fla.ops.utils.index import prepare_chunk_indices
|
||||
from kda._fla.utils import autotune_cache_kwargs, check_shared_mem, input_guard
|
||||
|
||||
BS_LIST = [32, 64] if check_shared_mem() else [16, 32]
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
'HAS_SCALE': lambda args: args['scale'] is not None,
|
||||
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
|
||||
})
|
||||
@fla_cache_autotune(
|
||||
configs=[
|
||||
triton.Config({}, num_warps=num_warps)
|
||||
for num_warps in [1, 2, 4, 8]
|
||||
],
|
||||
key=['B', 'H', 'BT', 'IS_VARLEN', 'REVERSE'],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=['T'])
|
||||
def chunk_local_cumsum_scalar_kernel(
|
||||
s,
|
||||
o,
|
||||
scale,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
B: tl.constexpr,
|
||||
H: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
REVERSE: tl.constexpr,
|
||||
HAS_SCALE: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
i_t, i_bh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64)
|
||||
i_b, i_h = i_bh // H, i_bh % H
|
||||
if IS_VARLEN:
|
||||
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
|
||||
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
|
||||
T = eos - bos
|
||||
else:
|
||||
bos, eos = i_b * T, i_b * T + T
|
||||
|
||||
o_t = i_t * BT + tl.arange(0, BT)
|
||||
m_t = o_t < T
|
||||
p_s = s + bos*H + i_h + o_t * H
|
||||
p_o = o + bos*H + i_h + o_t * H
|
||||
# [BT]
|
||||
b_s = tl.load(p_s, mask=m_t, other=0.0).to(tl.float32)
|
||||
if REVERSE:
|
||||
b_o = tl.cumsum(b_s, axis=0, reverse=True)
|
||||
else:
|
||||
b_o = tl.cumsum(b_s, axis=0)
|
||||
if HAS_SCALE:
|
||||
b_o *= scale
|
||||
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
'HAS_SCALE': lambda args: args['scale'] is not None,
|
||||
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
|
||||
})
|
||||
@fla_cache_autotune(
|
||||
configs=[
|
||||
triton.Config({'BS': BS}, num_warps=num_warps)
|
||||
for BS in BS_LIST
|
||||
for num_warps in [2, 4, 8]
|
||||
],
|
||||
key=['B', 'H', 'S', 'BT', 'IS_VARLEN', 'REVERSE'],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=['T'])
|
||||
def chunk_local_cumsum_vector_kernel(
|
||||
s,
|
||||
o,
|
||||
scale,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
B: tl.constexpr,
|
||||
H: tl.constexpr,
|
||||
S: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BS: tl.constexpr,
|
||||
REVERSE: tl.constexpr,
|
||||
HAS_SCALE: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
i_s, i_t, i_bh = tl.program_id(0), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
|
||||
i_b, i_h = i_bh // H, i_bh % H
|
||||
if IS_VARLEN:
|
||||
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int64)
|
||||
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
|
||||
T = eos - bos
|
||||
else:
|
||||
bos, eos = i_b * T, i_b * T + T
|
||||
|
||||
o_t = i_t * BT + tl.arange(0, BT)
|
||||
o_s = i_s * BS + tl.arange(0, BS)
|
||||
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
|
||||
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
|
||||
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
|
||||
# [BT, BS]
|
||||
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
|
||||
if REVERSE:
|
||||
b_o = tl.cumsum(b_s, axis=0, reverse=True)
|
||||
else:
|
||||
b_o = tl.cumsum(b_s, axis=0)
|
||||
if HAS_SCALE:
|
||||
b_o *= scale
|
||||
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_s)
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
'HAS_SCALE': lambda args: args['scale'] is not None,
|
||||
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
|
||||
})
|
||||
@triton.autotune(
|
||||
configs=[
|
||||
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
|
||||
for BT in [32, 64, 128, 256]
|
||||
for num_warps in [2, 4, 8]
|
||||
for num_stages in [1, 2, 3, 4]
|
||||
],
|
||||
key=['B', 'H', 'IS_VARLEN', 'REVERSE'],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=['T'])
|
||||
def chunk_global_cumsum_scalar_kernel(
|
||||
s,
|
||||
o,
|
||||
scale,
|
||||
cu_seqlens,
|
||||
T,
|
||||
B: tl.constexpr,
|
||||
H: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
REVERSE: tl.constexpr,
|
||||
HAS_SCALE: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
i_nh = tl.program_id(0).to(tl.int64)
|
||||
i_n, i_h = i_nh // H, i_nh % H
|
||||
if IS_VARLEN:
|
||||
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
|
||||
else:
|
||||
bos, eos = i_n * T, i_n * T + T
|
||||
T = eos - bos
|
||||
|
||||
b_z = tl.zeros([], dtype=tl.float32)
|
||||
NT = tl.cdiv(T, BT)
|
||||
for i_c in range(NT):
|
||||
i_t = NT - 1 - i_c if REVERSE else i_c
|
||||
o_t = i_t * BT + tl.arange(0, BT)
|
||||
m_t = o_t < T
|
||||
p_s = s + bos*H + i_h + o_t * H
|
||||
p_o = o + bos*H + i_h + o_t * H
|
||||
b_s = tl.load(p_s, mask=m_t, other=0.0).to(tl.float32)
|
||||
if REVERSE:
|
||||
b_o = tl.cumsum(b_s, axis=0, reverse=True)
|
||||
else:
|
||||
b_o = tl.cumsum(b_s, axis=0)
|
||||
b_ss = tl.sum(b_s, 0)
|
||||
b_o += b_z
|
||||
if i_c >= 0:
|
||||
b_z += b_ss
|
||||
if HAS_SCALE:
|
||||
b_o *= scale
|
||||
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_t)
|
||||
|
||||
|
||||
@triton.heuristics({
|
||||
'HAS_SCALE': lambda args: args['scale'] is not None,
|
||||
'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
|
||||
})
|
||||
@triton.autotune(
|
||||
configs=[
|
||||
triton.Config({'BT': BT}, num_warps=num_warps, num_stages=num_stages)
|
||||
for BT in [16, 32, 64, 128]
|
||||
for num_warps in [2, 4, 8]
|
||||
for num_stages in [1, 2, 3, 4]
|
||||
],
|
||||
key=['B', 'H', 'S', 'IS_VARLEN', 'REVERSE'],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=['T'])
|
||||
def chunk_global_cumsum_vector_kernel(
|
||||
s,
|
||||
o,
|
||||
scale,
|
||||
cu_seqlens,
|
||||
T,
|
||||
B: tl.constexpr,
|
||||
H: tl.constexpr,
|
||||
S: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BS: tl.constexpr,
|
||||
REVERSE: tl.constexpr,
|
||||
HAS_SCALE: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
i_s, i_nh = tl.program_id(0), tl.program_id(1).to(tl.int64)
|
||||
i_n, i_h = i_nh // H, i_nh % H
|
||||
if IS_VARLEN:
|
||||
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
|
||||
else:
|
||||
bos, eos = i_n * T, i_n * T + T
|
||||
T = eos - bos
|
||||
|
||||
b_z = tl.zeros([BS], dtype=tl.float32)
|
||||
NT = tl.cdiv(T, BT)
|
||||
for i_c in range(NT):
|
||||
i_t = NT - 1 - i_c if REVERSE else i_c
|
||||
o_t = i_t * BT + tl.arange(0, BT)
|
||||
o_s = i_s * BS + tl.arange(0, BS)
|
||||
m_s = (o_t[:, None] < T) & (o_s[None, :] < S)
|
||||
p_s = s + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
|
||||
p_o = o + (bos * H + i_h) * S + o_t[:, None] * (H*S) + o_s[None, :]
|
||||
# [BT, BS]
|
||||
b_s = tl.load(p_s, mask=m_s, other=0.0).to(tl.float32)
|
||||
if REVERSE:
|
||||
b_c = b_z[None, :] + tl.cumsum(b_s, axis=0, reverse=True)
|
||||
else:
|
||||
b_c = b_z[None, :] + tl.cumsum(b_s, axis=0)
|
||||
if HAS_SCALE:
|
||||
b_c *= scale
|
||||
tl.store(p_o, b_c.to(p_o.dtype.element_ty), mask=m_s)
|
||||
b_z += tl.sum(b_s, 0)
|
||||
|
||||
|
||||
def chunk_local_cumsum_scalar(
|
||||
g: torch.Tensor,
|
||||
chunk_size: int,
|
||||
reverse: bool = False,
|
||||
scale: float = None,
|
||||
cu_seqlens: torch.Tensor | None = None,
|
||||
output_dtype: torch.dtype | None = torch.float,
|
||||
chunk_indices: torch.LongTensor | None = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if 'head_first' in kwargs:
|
||||
raise DeprecationWarning(
|
||||
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
|
||||
)
|
||||
B, T, H = g.shape
|
||||
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
|
||||
BT = chunk_size
|
||||
if chunk_indices is None and cu_seqlens is not None:
|
||||
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
|
||||
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
|
||||
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
|
||||
grid = (NT, B * H)
|
||||
chunk_local_cumsum_scalar_kernel[grid](
|
||||
s=g_org,
|
||||
o=g,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
B=B,
|
||||
H=H,
|
||||
BT=BT,
|
||||
REVERSE=reverse,
|
||||
)
|
||||
return g
|
||||
|
||||
|
||||
def chunk_local_cumsum_vector(
|
||||
g: torch.Tensor,
|
||||
chunk_size: int,
|
||||
reverse: bool = False,
|
||||
scale: float = None,
|
||||
cu_seqlens: torch.Tensor | None = None,
|
||||
output_dtype: torch.dtype | None = torch.float,
|
||||
chunk_indices: torch.LongTensor | None = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if 'head_first' in kwargs:
|
||||
raise DeprecationWarning(
|
||||
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
|
||||
)
|
||||
B, T, H, S = g.shape
|
||||
BT = chunk_size
|
||||
if chunk_indices is None and cu_seqlens is not None:
|
||||
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
|
||||
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
|
||||
assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
|
||||
|
||||
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
|
||||
def grid(meta): return (triton.cdiv(meta['S'], meta['BS']), NT, B * H)
|
||||
# keep cummulative normalizer in fp32
|
||||
# this kernel is equivalent to
|
||||
# g = g.view(B, H, NT, BT, -1).cumsum(-2).view(B, H, T, -1)
|
||||
chunk_local_cumsum_vector_kernel[grid](
|
||||
s=g_org,
|
||||
o=g,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
B=B,
|
||||
H=H,
|
||||
S=S,
|
||||
BT=BT,
|
||||
REVERSE=reverse,
|
||||
)
|
||||
return g
|
||||
|
||||
|
||||
@input_guard
|
||||
def chunk_global_cumsum_scalar(
|
||||
s: torch.Tensor,
|
||||
reverse: bool = False,
|
||||
cu_seqlens: torch.Tensor | None = None,
|
||||
scale: float = None,
|
||||
output_dtype: torch.dtype | None = torch.float,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if 'head_first' in kwargs:
|
||||
raise DeprecationWarning(
|
||||
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
|
||||
)
|
||||
B, T, H = s.shape
|
||||
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
|
||||
|
||||
z = torch.empty_like(s, dtype=output_dtype or s.dtype)
|
||||
grid = (N * H,)
|
||||
chunk_global_cumsum_scalar_kernel[grid](
|
||||
s=s,
|
||||
o=z,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
T=T,
|
||||
B=B,
|
||||
H=H,
|
||||
REVERSE=reverse,
|
||||
)
|
||||
return z
|
||||
|
||||
|
||||
@input_guard
|
||||
def chunk_global_cumsum_vector(
|
||||
s: torch.Tensor,
|
||||
reverse: bool = False,
|
||||
cu_seqlens: torch.Tensor | None = None,
|
||||
scale: float = None,
|
||||
output_dtype: torch.dtype | None = torch.float,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if 'head_first' in kwargs:
|
||||
raise DeprecationWarning(
|
||||
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
|
||||
)
|
||||
B, T, H, S = s.shape
|
||||
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
|
||||
BS = min(32, triton.next_power_of_2(S))
|
||||
|
||||
z = torch.empty_like(s, dtype=output_dtype or s.dtype)
|
||||
grid = (triton.cdiv(S, BS), N * H)
|
||||
chunk_global_cumsum_vector_kernel[grid](
|
||||
s=s,
|
||||
o=z,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
T=T,
|
||||
B=B,
|
||||
H=H,
|
||||
S=S,
|
||||
BS=BS,
|
||||
REVERSE=reverse,
|
||||
)
|
||||
return z
|
||||
|
||||
|
||||
@input_guard
|
||||
@dispatch('utils')
|
||||
def chunk_global_cumsum(
|
||||
s: torch.Tensor,
|
||||
reverse: bool = False,
|
||||
cu_seqlens: torch.Tensor | None = None,
|
||||
scale: float = None,
|
||||
output_dtype: torch.dtype | None = torch.float,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if 'head_first' in kwargs:
|
||||
raise DeprecationWarning(
|
||||
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
|
||||
)
|
||||
if cu_seqlens is not None:
|
||||
assert s.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
|
||||
if len(s.shape) == 3:
|
||||
return chunk_global_cumsum_scalar(
|
||||
s=s,
|
||||
reverse=reverse,
|
||||
cu_seqlens=cu_seqlens,
|
||||
scale=scale,
|
||||
output_dtype=output_dtype,
|
||||
)
|
||||
elif len(s.shape) == 4:
|
||||
return chunk_global_cumsum_vector(
|
||||
s=s,
|
||||
reverse=reverse,
|
||||
cu_seqlens=cu_seqlens,
|
||||
scale=scale,
|
||||
output_dtype=output_dtype,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported input shape {s.shape}, "
|
||||
f"which should be [B, T, H] or [B, T, H, D]",
|
||||
)
|
||||
|
||||
|
||||
@input_guard
|
||||
@dispatch('utils')
|
||||
def chunk_local_cumsum(
|
||||
g: torch.Tensor,
|
||||
chunk_size: int,
|
||||
reverse: bool = False,
|
||||
scale: float = None,
|
||||
cu_seqlens: torch.Tensor | None = None,
|
||||
output_dtype: torch.dtype | None = torch.float,
|
||||
chunk_indices: torch.LongTensor | None = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if 'head_first' in kwargs:
|
||||
raise DeprecationWarning(
|
||||
"head_first has been removed. Inputs must be in `[B, T, H, ...]` format.",
|
||||
)
|
||||
if cu_seqlens is not None:
|
||||
assert g.shape[0] == 1, "Only batch size 1 is supported when cu_seqlens are provided"
|
||||
if len(g.shape) == 3:
|
||||
return chunk_local_cumsum_scalar(
|
||||
g=g,
|
||||
chunk_size=chunk_size,
|
||||
reverse=reverse,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
output_dtype=output_dtype,
|
||||
chunk_indices=chunk_indices,
|
||||
)
|
||||
elif len(g.shape) == 4:
|
||||
return chunk_local_cumsum_vector(
|
||||
g=g,
|
||||
chunk_size=chunk_size,
|
||||
reverse=reverse,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
output_dtype=output_dtype,
|
||||
chunk_indices=chunk_indices,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported input shape {g.shape}, "
|
||||
f"which should be (B, T, H) or (B, T, H, D)",
|
||||
)
|
||||
@@ -0,0 +1,183 @@
|
||||
# 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 torch
|
||||
import torch.nn.functional as F
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from kda._fla.utils import autotune_cache_kwargs, tensor_cache
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=[
|
||||
triton.Config({}, num_warps=num_warps)
|
||||
for num_warps in [4, 8, 16, 32]
|
||||
],
|
||||
key=['B'],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit
|
||||
def prepare_position_ids_kernel(
|
||||
y,
|
||||
cu_seqlens,
|
||||
B: tl.constexpr,
|
||||
):
|
||||
i_n = tl.program_id(0)
|
||||
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
|
||||
T = eos - bos
|
||||
|
||||
o = tl.arange(0, B)
|
||||
for i in range(0, tl.cdiv(T, B) * B, B):
|
||||
o_i = o + i
|
||||
tl.store(y + bos + o_i, o_i, o_i < T)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_lens(cu_seqlens: torch.LongTensor) -> torch.LongTensor:
|
||||
return torch.diff(cu_seqlens)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_lens_from_mask(mask: torch.BoolTensor) -> torch.LongTensor:
|
||||
return mask.sum(dim=-1, dtype=torch.int32)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_cu_seqlens_from_lens(
|
||||
lens: torch.LongTensor,
|
||||
dtype: torch.dtype | None = torch.int32,
|
||||
) -> torch.LongTensor:
|
||||
return F.pad(lens.cumsum(dim=0, dtype=dtype), (1, 0))
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_cu_seqlens_from_mask(
|
||||
mask: torch.BoolTensor,
|
||||
dtype: torch.dtype | None = torch.int32,
|
||||
) -> torch.LongTensor:
|
||||
return prepare_cu_seqlens_from_lens(prepare_lens_from_mask(mask), dtype)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_split_cu_seqlens(
|
||||
batch_size: int | None = None,
|
||||
seq_len: int | None = None,
|
||||
split_size: int | None = None,
|
||||
cu_seqlens: torch.LongTensor | None = None,
|
||||
dtype: torch.dtype | None = torch.int32,
|
||||
device: torch.device | None = torch.device('cpu'),
|
||||
) -> torch.LongTensor:
|
||||
"""Sub-split a (optionally packed) batch along the token axis.
|
||||
|
||||
Two calling modes:
|
||||
- **Rectangular batch**: pass `batch_size` and `seq_len`, leave
|
||||
`cu_seqlens=None`. Internally synthesizes `[0, L, 2L, ..., B*L]`.
|
||||
- **Packed varlen**: pass `cu_seqlens`. `batch_size` and `seq_len` are
|
||||
ignored (kept as optional kwargs for backward-compat with callers
|
||||
that used to pass dummies).
|
||||
|
||||
`split_size` is always required.
|
||||
|
||||
The legacy positional signature `(batch_size, seq_len, split_size, ...)`
|
||||
continues to work — the first two args retain their position but may now
|
||||
be omitted when `cu_seqlens` is supplied.
|
||||
"""
|
||||
if split_size is None:
|
||||
raise TypeError("prepare_split_cu_seqlens() requires `split_size`")
|
||||
if cu_seqlens is None:
|
||||
if batch_size is None or seq_len is None:
|
||||
raise TypeError(
|
||||
"prepare_split_cu_seqlens() requires either `cu_seqlens`, "
|
||||
"or both `batch_size` and `seq_len`"
|
||||
)
|
||||
total_tokens = batch_size * seq_len
|
||||
cu_seqlens = list(range(0, total_tokens, seq_len)) + [total_tokens]
|
||||
else:
|
||||
cu_seqlens = cu_seqlens.tolist()
|
||||
return torch.tensor(
|
||||
[
|
||||
i
|
||||
for bos, eos in zip(cu_seqlens[:-1], cu_seqlens[1:], strict=False)
|
||||
for i in range(bos, eos, split_size)
|
||||
] + [cu_seqlens[-1]],
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
|
||||
def _segmented_arange(counts: torch.LongTensor) -> tuple[torch.LongTensor, torch.LongTensor]:
|
||||
"""Expand per-segment counts into flat per-slot index tensors.
|
||||
|
||||
Given segment sizes ``counts = [c0, c1, ...]``, return two 1-D tensors of
|
||||
length ``counts.sum()`` that together label every slot with its segment and
|
||||
its position within that segment.
|
||||
|
||||
Example -- ``counts = [2, 3]`` (segment 0 spans 2 slots, segment 1 spans 3)::
|
||||
|
||||
seg_id = [0, 0, 1, 1, 1] # which segment each slot belongs to
|
||||
intra_idx = [0, 1, 0, 1, 2] # running index within that segment
|
||||
|
||||
With CUDA ``counts``, ``repeat_interleave`` reads ``counts.sum()`` on the
|
||||
host (one device sync). Pass host-side counts to avoid it.
|
||||
"""
|
||||
seg_id = torch.repeat_interleave(
|
||||
torch.arange(counts.numel(), device=counts.device, dtype=counts.dtype),
|
||||
counts,
|
||||
)
|
||||
seg_start = F.pad(counts.cumsum(0), (1, 0))[:-1]
|
||||
intra_idx = torch.arange(seg_id.shape[0], device=counts.device, dtype=counts.dtype) - seg_start[seg_id]
|
||||
return seg_id, intra_idx
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_position_ids(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
|
||||
src = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens
|
||||
_, position_ids = _segmented_arange(prepare_lens(src))
|
||||
return position_ids.to(cu_seqlens)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_sequence_ids(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
|
||||
return prepare_position_ids(cu_seqlens, cu_seqlens_cpu).eq(0).cumsum(0) - 1
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_token_indices(cu_seqlens: torch.LongTensor, cu_seqlens_cpu: torch.LongTensor | None = None) -> torch.LongTensor:
|
||||
position_ids = prepare_position_ids(cu_seqlens, cu_seqlens_cpu)
|
||||
return torch.stack([prepare_sequence_ids(cu_seqlens, cu_seqlens_cpu), position_ids], 1).to(cu_seqlens)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_chunk_indices(
|
||||
cu_seqlens: torch.LongTensor,
|
||||
chunk_size: int,
|
||||
cu_seqlens_cpu: torch.LongTensor | None = None,
|
||||
) -> torch.LongTensor:
|
||||
src = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens
|
||||
chunk_counts = (prepare_lens(src) + (chunk_size - 1)).div(chunk_size, rounding_mode='floor')
|
||||
seg_id, intra_chunk_idx = _segmented_arange(chunk_counts)
|
||||
return torch.stack([seg_id, intra_chunk_idx], 1).to(cu_seqlens)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_chunk_offsets(
|
||||
cu_seqlens: torch.LongTensor,
|
||||
chunk_size: int,
|
||||
) -> torch.LongTensor:
|
||||
return F.pad(triton.cdiv(prepare_lens(cu_seqlens), chunk_size), (1, 0), value=0).cumsum(-1)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def get_max_num_splits(
|
||||
cu_seqlens: torch.LongTensor,
|
||||
chunk_size: int,
|
||||
cu_seqlens_cpu: torch.LongTensor | None = None
|
||||
) -> int:
|
||||
if cu_seqlens_cpu is not None:
|
||||
return triton.cdiv(int(max(prepare_lens(cu_seqlens_cpu))), chunk_size)
|
||||
return triton.cdiv(int(max(prepare_lens(cu_seqlens))), chunk_size)
|
||||
@@ -0,0 +1,101 @@
|
||||
# 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
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import triton.language.extra.libdevice as tldevice
|
||||
|
||||
from kda._fla.utils import IS_GATHER_SUPPORTED, IS_NVIDIA_BLACKWELL
|
||||
|
||||
if os.environ.get('FLA_USE_FAST_OPS', '0') == '1':
|
||||
@triton.jit
|
||||
def exp(x): return tldevice.fast_expf(x.to(tl.float32))
|
||||
@triton.jit
|
||||
def exp2(x): return tldevice.exp2(x.to(tl.float32))
|
||||
@triton.jit
|
||||
def log(x): return tldevice.fast_logf(x.to(tl.float32))
|
||||
@triton.jit
|
||||
def log2(x): return tldevice.fast_log2f(x.to(tl.float32))
|
||||
@triton.jit
|
||||
def tanh(x): return tldevice.fast_tanhf(x.to(tl.float32))
|
||||
else:
|
||||
@triton.jit
|
||||
def exp(x): return tl.exp(x.to(tl.float32))
|
||||
@triton.jit
|
||||
def exp2(x): return tl.math.exp2(x.to(tl.float32))
|
||||
@triton.jit
|
||||
def log(x): return tl.log(x.to(tl.float32))
|
||||
@triton.jit
|
||||
def log2(x): return tl.log2(x.to(tl.float32))
|
||||
@triton.jit
|
||||
def tanh(x): return tldevice.tanh(x.to(tl.float32))
|
||||
|
||||
|
||||
if IS_NVIDIA_BLACKWELL:
|
||||
"""
|
||||
Compute tl.dot with Blackwell workaround.
|
||||
|
||||
On SM100 datacenter and SM120 consumer Blackwell GPUs, wraps the result in
|
||||
inline assembly to prevent the TritonGPUHoistTMEMAlloc pass from incorrectly
|
||||
fusing add and dot operations.
|
||||
See: https://github.com/fla-org/flash-linear-attention/issues/638
|
||||
|
||||
TODO: Remove this workaround once the Triton compiler bug is fixed.
|
||||
Track upstream issue at: https://github.com/triton-lang/triton/issues/8695
|
||||
"""
|
||||
@triton.jit
|
||||
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
|
||||
return tl.inline_asm_elementwise(
|
||||
asm="mov.f32 $0, $1;",
|
||||
constraints="=r,r",
|
||||
args=[tl.dot(a, b, allow_tf32=allow_tf32)],
|
||||
dtype=tl.float32,
|
||||
is_pure=True,
|
||||
pack=1,
|
||||
)
|
||||
else:
|
||||
@triton.jit
|
||||
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
|
||||
return tl.dot(a, b, allow_tf32=allow_tf32)
|
||||
|
||||
|
||||
if not IS_GATHER_SUPPORTED:
|
||||
@triton.jit
|
||||
def gather(src, index, axis, _builder=None):
|
||||
"""
|
||||
Gather operation that works when tl.gather is not supported.
|
||||
This is a fallback implementation that returns None.
|
||||
Just to make triton compiler happy.
|
||||
"""
|
||||
return None
|
||||
else:
|
||||
gather = tl.gather
|
||||
|
||||
|
||||
if hasattr(triton.language, '_experimental_make_tensor_descriptor'):
|
||||
# For Triton 3.3.x
|
||||
make_tensor_descriptor = triton.language._experimental_make_tensor_descriptor
|
||||
elif hasattr(triton.language, 'make_tensor_descriptor'):
|
||||
# For Triton 3.4.x and later
|
||||
make_tensor_descriptor = triton.language.make_tensor_descriptor
|
||||
else:
|
||||
"""
|
||||
Fallback implementation when TMA is not supported.
|
||||
Returns None to indicate TMA descriptors are unavailable.
|
||||
Just make triton compiler happy.
|
||||
"""
|
||||
@triton.jit
|
||||
def make_tensor_descriptor(
|
||||
base,
|
||||
shape,
|
||||
strides,
|
||||
block_shape,
|
||||
_builder=None,
|
||||
):
|
||||
return None
|
||||
@@ -0,0 +1,115 @@
|
||||
# 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
|
||||
|
||||
# REVISED FROM
|
||||
# https://github.com/shawntan/stickbreaking-attention/blob/main/stickbreaking_attention/sb_varlen/softplus.py
|
||||
|
||||
import triton
|
||||
from triton import language as tl
|
||||
|
||||
from kda._fla.utils import IS_NVIDIA
|
||||
|
||||
|
||||
def _generate_softplus(num_pack):
|
||||
template = """
|
||||
.reg .pred p;
|
||||
setp.gt.f32 p, ${in_reg}, 20.;
|
||||
@p mov.f32 ${out_reg}, ${in_reg};
|
||||
@!p mul.f32 ${out_reg}, ${in_reg}, 1.4426950408889634;
|
||||
@!p ex2.approx.ftz.f32 ${out_reg}, ${out_reg};
|
||||
@!p add.f32 ${out_reg}, ${out_reg}, 1.0;
|
||||
@!p lg2.approx.ftz.f32 ${out_reg}, ${out_reg};
|
||||
@!p mul.f32 ${out_reg}, ${out_reg}, 0.6931471805599453;
|
||||
"""
|
||||
out_str = ""
|
||||
|
||||
for i in range(num_pack):
|
||||
inner_str = template.format(out_reg=i, in_reg=i + num_pack)
|
||||
out_str += "{" + inner_str + "}\n"
|
||||
# flatten out because torch.compile doesn't like newlines
|
||||
out_str = " ".join(out_str.split("\n"))
|
||||
return out_str
|
||||
|
||||
|
||||
def _generate_softplus2(num_pack):
|
||||
template = """
|
||||
.reg .pred p;
|
||||
setp.gt.f32 p, ${in_reg}, 15.;
|
||||
@p mov.f32 ${out_reg}, ${in_reg};
|
||||
@!p ex2.approx.ftz.f32 ${out_reg}, ${in_reg};
|
||||
@!p add.f32 ${out_reg}, ${out_reg}, 1.0;
|
||||
@!p lg2.approx.ftz.f32 ${out_reg}, ${out_reg};
|
||||
"""
|
||||
out_str = ""
|
||||
|
||||
for i in range(num_pack):
|
||||
inner_str = template.format(out_reg=i, in_reg=i + num_pack)
|
||||
out_str += "{" + inner_str + "}\n"
|
||||
# flatten out because torch.compile doesn't like newlines
|
||||
out_str = " ".join(out_str.split("\n"))
|
||||
return out_str
|
||||
|
||||
|
||||
def _generate_constraints(num_pack):
|
||||
return ",".join("=r" for i in range(num_pack)) + "," + ",".join("r" for i in range(num_pack))
|
||||
|
||||
|
||||
_NUM_REG = 1
|
||||
s_softplus: tl.constexpr = tl.constexpr(_generate_softplus(_NUM_REG))
|
||||
s_softplus2: tl.constexpr = tl.constexpr(_generate_softplus2(_NUM_REG))
|
||||
s_constraints: tl.constexpr = tl.constexpr(_generate_constraints(_NUM_REG))
|
||||
NUM_REG: tl.constexpr = tl.constexpr(_NUM_REG)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def softplus_nv(x):
|
||||
# equivalent to:
|
||||
# return tl.where(x < 20.0, tl.math.log(1 + tl.math.exp(x)), x)
|
||||
return tl.inline_asm_elementwise(
|
||||
asm=s_softplus,
|
||||
constraints=s_constraints,
|
||||
pack=NUM_REG,
|
||||
args=[
|
||||
x,
|
||||
],
|
||||
dtype=tl.float32,
|
||||
is_pure=True,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def softplus_triton(x):
|
||||
return tl.where(x < 20.0, tl.math.log(1 + tl.math.exp(x)), x)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def softplus2_nv(x):
|
||||
# equivalent to:
|
||||
# return tl.where(x < 15.0, tl.math.log2(1 + tl.math.exp2(x)), x)
|
||||
return tl.inline_asm_elementwise(
|
||||
asm=s_softplus2,
|
||||
constraints=s_constraints,
|
||||
pack=NUM_REG,
|
||||
args=[
|
||||
x,
|
||||
],
|
||||
dtype=tl.float32,
|
||||
is_pure=True,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def softplus2_triton(x):
|
||||
return tl.where(x < 15.0, tl.math.log2(1 + tl.math.exp2(x)), x)
|
||||
|
||||
|
||||
if IS_NVIDIA:
|
||||
softplus = softplus_nv
|
||||
softplus2 = softplus2_nv
|
||||
else:
|
||||
softplus = softplus_triton
|
||||
softplus2 = softplus2_triton
|
||||
Reference in New Issue
Block a user