File size: 8,122 Bytes
f145768
 
6709b3f
f145768
 
 
 
 
 
6709b3f
f145768
6709b3f
 
 
 
 
 
f145768
6709b3f
 
 
 
f145768
6709b3f
 
 
f145768
6709b3f
 
 
 
f145768
6709b3f
 
 
 
 
f145768
6709b3f
 
 
f145768
6709b3f
 
f145768
 
 
 
 
 
07e0355
 
 
 
 
 
ee0b350
 
 
 
07e0355
6709b3f
 
 
 
 
 
 
 
 
5532501
 
 
 
6709b3f
 
 
 
 
 
 
 
 
 
f145768
 
07e0355
f145768
 
 
07e0355
 
 
ee0b350
 
07e0355
 
 
 
 
ee0b350
 
 
6709b3f
5532501
 
 
6709b3f
 
f145768
 
 
 
 
 
 
 
d8005f6
 
f145768
 
 
 
 
 
 
 
 
 
 
6709b3f
 
 
 
f145768
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6709b3f
f145768
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d8005f6
f145768
 
 
 
 
6709b3f
f145768
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
"""KDA (Kernel-based Decomposed Attention) NKI kernels for AWS Neuron / Trainium.

Model-agnostic NKI implementation of KDA linear attention, compatible with any
HuggingFace Transformers model whose attention layer follows the KDA algorithm
described in the flash-linear-attention (fla-core) library.

References:
    - Algorithm: https://github.com/fla-org/flash-linear-attention (KDA)

Entry points (see README for full signatures and the input contract):

  Inference forward:
    - kda_recurrent_fwd / kda_recurrent_fwd_state -- decode per-token recurrence
    - kda_chunk_step -- prefill per-chunk (scalar-mean
                                                       gate approximation; see warning)
    - kda_chunk_step_exact / _multihead -- numerically exact prefill
    - kda_decode_batch -- batched multi-request decode

  Training (differentiable):
    - kda_recurrent, kda_chunked, kda_chunked_fused
    - raw fwd/bwd kernels: kda_recurrent_fwd_v2, kda_chunk_step_v2,
      kda_recurrent_bwd, kda_chunk_bwd, kda_fused_chunked_fwd, kda_fused_chunked_bwd

Input contract:
    Callers pass q, k already L2-normed, with q scaled by 1/sqrt(dk) (fla-core
    convention). The kernels compute all decay-related scaling internally from g.

Constraints:
    - head_k_dim == head_v_dim == 128 (NeuronCore SBUF partition width)
    - chunk_size == 128 for the chunked kernels; S divisible by 128
    - float32 inputs

Functional warning:
    kda_chunk_step (and its training wrappers) use a scalar-mean gate-decay
    approximation that is only accurate for small gate decay. For non-trivial
    gate decay, use kda_chunk_step_exact. The recurrent kernels are exact in all
    regimes. See the README for details.

Requirements:
    - PyTorch Native (device="neuron"), torch-neuronx 2.11+, PyTorch 2.11+
    - NKI >= 0.4.0

License: Apache-2.0. Inference/training runtime kernel package (not a model).
The fla-core algorithm reference is MIT-licensed and compatible.
"""

# Re-export the kernel entry points at package level
from .nki_kda import kda_recurrent_fwd, kda_recurrent_fwd_state
from .nki_kda_chunked import kda_chunk_step

# Training (v1.2): forward-with-saved-intermediates + backward kernels, and the
# differentiable torch.autograd.Function wrappers.
from .kda_recurrent_fwd_v2 import kda_recurrent_fwd_v2
from .kda_chunk_step_v2 import kda_chunk_step_v2
from .kda_recurrent_bwd import kda_recurrent_bwd
from .kda_chunk_bwd import kda_chunk_bwd
from .kda_autograd import kda_recurrent, kda_chunked, kda_chunked_fused
# Fused multi-chunk kernels (single NKI launch for all chunks) -- v1.3
from .kda_fused_chunked_fwd import kda_fused_chunked_fwd
from .kda_fused_chunked_bwd import kda_fused_chunked_bwd

# Forward optimizations:
# - kda_recurrent_fwd / _state use a fused decode body (key-fold): the decode
# wall-clock at S=128 is ~527 us and parity is 1.0000.
# - kda_decode_batch: batched multi-(request, head) decode in one launch;
# amortizes launch overhead across a serving batch.
# - kda_chunk_step_exact: numerically exact per-channel prefill
# (cos_sim >= 0.9999999 in all gate regimes) vs the scalar-mean approximation
# in kda_chunk_step. ~1.05x the latency of the approximation.
# - kda_chunk_step_exact_multihead: head-interleaved exact prefill.
from .nki_kda_decode_batch import kda_decode_batch
from .nki_kda_chunked_exact import kda_chunk_step_exact
from .nki_kda_chunked_exact_multihead import kda_chunk_step_exact_multihead

# v1.5 -- exact chunked BACKWARD (matching the exact forward). The prior chunked
# backward differentiates the scalar-mean approximation: its dg gradient is
# essentially uncorrelated with the true gradient (cos_sim ~0.09) in ALL regimes,
# and it NaNs at large gate decay. kda_chunk_step_exact_bwd is the gradient of the
# exact 16-sub-chunk + WY forward -- cos_sim 1.0 vs fla-core autograd for
# dq/dk/dv/dg/dbeta in all gate regimes (g=0.01..2.0), no overflow. Use this
# backward whenever you use kda_chunk_step_exact on the forward / for training with
# non-trivial gating.
from .nki_kda_chunked_exact_bwd import kda_chunk_step_exact_bwd


__all__ = [
    # inference forward kernels (v1.0/v1.1)
    "kda_recurrent_fwd",
    "kda_recurrent_fwd_state",
    "kda_chunk_step",
    # training: differentiable wrappers (v1.2)
    "kda_recurrent",
    "kda_chunked",
    # training: fused single-launch chunked wrapper (v1.3)
    "kda_chunked_fused",
    # training: raw fwd-v2 + backward kernels (v1.2)
    "kda_recurrent_fwd_v2",
    "kda_chunk_step_v2",
    "kda_recurrent_bwd",
    "kda_chunk_bwd",
    # training: raw fused kernels (v1.3)
    "kda_fused_chunked_fwd",
    "kda_fused_chunked_bwd",
    # v1.4 fused decode body forward optimizations 
    "kda_decode_batch",
    "kda_chunk_step_exact",
    "kda_chunk_step_exact_multihead",
    # v1.5 exact chunked backward
    "kda_chunk_step_exact_bwd",
]


# =============================================================================
# HF Transformers `KernelConfig` wrapper stub
# =============================================================================
#
# The wrapping below exposes the kernels through the HF `KernelConfig` API.
# For the exact `Kda` layer name and weight schema of the target model,
# see that model's `configuration_kda.py` and
# `modeling_kda.py` in transformers (or the equivalent -- naming is not
# yet finalized in upstream transformers as of 2026-08).
#
# Downstream users who need a full model integration (projections, conv1d,
# gating, RMS norm, out_proj) should follow the pattern in
# `jburtoft/qwen35-deltanet-neuron-kernels:NeuronGatedDeltaNet` (the
# sibling kernel package for the DeltaNet family). This KDA package
# publishes ONLY the raw kernels; the surrounding layer glue is model-
# specific and lives with the model wrapper, not with the kernel package.
#
# If a user wants the full-layer wrapper, they can either:
# 1. Build it from the target model's `modeling_kda.py` plus these
# three kernel calls.
# 2. Wait for a v1.1 release of this package that includes the wrapper
# once the upstream transformers KDA integration is finalized.
#
# =============================================================================

# Placeholder KernelConfig-compatible class -- documents the intended API surface.
# Real full-layer wrapper is a v1.1 deliverable.

import torch
import torch.nn as nn


class NeuronKDA(nn.Module):
    """Placeholder for the full HF `Kda` layer replacement.

    v1.0 of this package exposes only the raw NKI kernels
    (`kda_recurrent_fwd`, `kda_recurrent_fwd_state`, `kda_chunk_step`).
    Downstream users assemble the surrounding layer arithmetic
    (projections, conv1d, gating, RMS norm) themselves.

    A future v1.1 will provide a drop-in `NeuronKDA(nn.Module)` matching the
    upstream transformers `Kda` layer's forward signature. Until then, see
    the module docstring above and the example in `tests/example_usage.py`.
    """

    def __init__(self, config, layer_idx: int):
        super.__init__
        raise NotImplementedError(
            "NeuronKDA full-layer wrapper is planned for v1.1. "
            "For v1.0, use the raw kernels: `kda_recurrent_fwd`, "
            "`kda_recurrent_fwd_state`, `kda_chunk_step`. See the package "
            "docstring for the wrapper contract."
        )

    def forward(self, hidden_states, cache_params=None, attention_mask=None, **kwargs):
        raise NotImplementedError


class NeuronKDALayout(nn.Module):
    """Placeholder for the weight-layout class for the full HF `Kda` layer.

    Structurally analogous to `NeuronGatedDeltaNetLayout` in the sibling
    deltanet kernel package. Will be filled in for v1.1 once the upstream
    the upstream `Kda` layer's weight schema is finalized.
    """

    conversion_mapping = []

    def __init__(self, config, layer_idx: int):
        super.__init__
        raise NotImplementedError(
            "NeuronKDALayout is planned for v1.1. See package docstring."
        )


class layers:
    """Layer registry for HF `KernelConfig`. Populated in v1.1."""
    NeuronKDA = NeuronKDA
    NeuronKDALayout = NeuronKDALayout