Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
38 lines
898 B
Python
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,
|
|
)
|