Files
K3/kda/ops/triton/chunk_fwd.py
T
dela 584f7e9e73 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.
2026-08-25 14:43:17 +08:00

38 lines
898 B
Python

"""Vendored FLA chunk KDA forward, returning ``(o, ht)`` like the public op."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.chunk import chunk_kda
from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd as fla_chunk_kda_fwd
__all__ = ["chunk_kda_fwd", "fla_chunk_kda_fwd"]
def chunk_kda_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
chunk_size: int = 64,
**kwargs,
):
"""Chunked KDA forward with FLA kernels. Returns ``(o, ht)``."""
return chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=chunk_size,
**kwargs,
)