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,7 @@
|
||||
"""Local Triton KDA kernels vendored from FLA chunk_{fwd,intra,bwd,wy,gate}."""
|
||||
|
||||
from .chunk import ChunkKDAFunction, chunk_kda
|
||||
from .chunk_fwd import chunk_kda_fwd
|
||||
from .gate import kda_gate_fwd
|
||||
|
||||
__all__ = ["ChunkKDAFunction", "chunk_kda", "chunk_kda_fwd", "kda_gate_fwd"]
|
||||
@@ -0,0 +1,5 @@
|
||||
"""FLA ``chunk_kda`` surface used by ``ops.api`` backend='triton'."""
|
||||
|
||||
from kda._fla.ops.kda.chunk import ChunkKDAFunction, chunk_kda
|
||||
|
||||
__all__ = ["ChunkKDAFunction", "chunk_kda"]
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Vendored FLA chunk KDA backward."""
|
||||
|
||||
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
|
||||
|
||||
__all__ = ["chunk_kda_bwd"]
|
||||
@@ -0,0 +1,37 @@
|
||||
"""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,
|
||||
)
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Vendored FLA KDA gate fusion (standard + safe gate + chunk cumsum)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from kda._fla.ops.kda.gate import (
|
||||
kda_gate_bwd,
|
||||
kda_gate_chunk_cumsum,
|
||||
kda_gate_fwd as _kda_gate_fwd,
|
||||
)
|
||||
|
||||
DEFAULT_LOWER_BOUND = -5.0
|
||||
|
||||
|
||||
def kda_gate_fwd(
|
||||
g: torch.Tensor,
|
||||
A_log: torch.Tensor,
|
||||
dt_bias: torch.Tensor | None = None,
|
||||
lower_bound: float | None = DEFAULT_LOWER_BOUND,
|
||||
):
|
||||
return _kda_gate_fwd(
|
||||
g,
|
||||
A_log=A_log,
|
||||
dt_bias=dt_bias,
|
||||
lower_bound=lower_bound,
|
||||
output_dtype=g.dtype,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_LOWER_BOUND",
|
||||
"kda_gate_bwd",
|
||||
"kda_gate_chunk_cumsum",
|
||||
"kda_gate_fwd",
|
||||
]
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Vendored FLA WY recompute used by the chunk KDA backward."""
|
||||
|
||||
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
|
||||
|
||||
__all__ = ["recompute_w_u_fwd"]
|
||||
Reference in New Issue
Block a user