Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
50 changes: 49 additions & 1 deletion skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api
Original file line number Diff line number Diff line change
Expand Up @@ -1246,6 +1246,35 @@ public final class sk/ainet/lang/memory/ViewsKt {
public static final fun viewOrNull (Lsk/ainet/lang/tensor/Tensor;)Lsk/ainet/lang/memory/TensorView;
}

public final class sk/ainet/lang/memory/WindowedAttention {
public static final field INSTANCE Lsk/ainet/lang/memory/WindowedAttention;
public final fun decodeStep ([FLsk/ainet/lang/memory/WindowedKV;Lsk/ainet/lang/memory/WindowedKV;[FF)V
public static synthetic fun decodeStep$default (Lsk/ainet/lang/memory/WindowedAttention;[FLsk/ainet/lang/memory/WindowedKV;Lsk/ainet/lang/memory/WindowedKV;[FFILjava/lang/Object;)V
public final fun decodeStepGathered ([FLsk/ainet/lang/memory/WindowedKV;Lsk/ainet/lang/memory/WindowedKV;[FLsk/ainet/lang/memory/Scope;Lsk/ainet/lang/memory/trace/TraceSink;F)V
public static synthetic fun decodeStepGathered$default (Lsk/ainet/lang/memory/WindowedAttention;[FLsk/ainet/lang/memory/WindowedKV;Lsk/ainet/lang/memory/WindowedKV;[FLsk/ainet/lang/memory/Scope;Lsk/ainet/lang/memory/trace/TraceSink;FILjava/lang/Object;)V
}

public final class sk/ainet/lang/memory/WindowedKV {
public static final field Companion Lsk/ainet/lang/memory/WindowedKV$Companion;
public static final field POSITION_AXIS I
public fun <init> (Lsk/ainet/lang/memory/TensorView;Lsk/ainet/lang/memory/TensorView;)V
public synthetic fun <init> (Lsk/ainet/lang/memory/TensorView;Lsk/ainet/lang/memory/TensorView;ILkotlin/jvm/internal/DefaultConstructorMarker;)V
public final fun gather (Lsk/ainet/lang/memory/Scope;Lsk/ainet/lang/memory/trace/TraceSink;Lsk/ainet/lang/tensor/TensorId;)Lsk/ainet/lang/memory/TensorView;
public static synthetic fun gather$default (Lsk/ainet/lang/memory/WindowedKV;Lsk/ainet/lang/memory/Scope;Lsk/ainet/lang/memory/trace/TraceSink;Lsk/ainet/lang/tensor/TensorId;ILjava/lang/Object;)Lsk/ainet/lang/memory/TensorView;
public final fun get (III)F
public final fun getHead ()Lsk/ainet/lang/memory/TensorView;
public final fun getHeadDim ()I
public final fun getHeads ()I
public final fun getLength ()I
public final fun getParts ()Ljava/util/List;
public final fun getTail ()Lsk/ainet/lang/memory/TensorView;
public final fun getWrapped ()Z
public fun toString ()Ljava/lang/String;
}

public final class sk/ainet/lang/memory/WindowedKV$Companion {
}

public final class sk/ainet/lang/memory/plan/ActualMemory {
public static final field Companion Lsk/ainet/lang/memory/plan/ActualMemory$Companion;
public fun <init> (Ljava/util/Map;Ljava/util/Map;Ljava/util/Map;Ljava/util/Map;)V
Expand Down Expand Up @@ -6625,6 +6654,8 @@ public final class sk/ainet/lang/tensor/storage/CompressedKvAttention {
public synthetic fun <init> (Lsk/ainet/lang/tensor/storage/KvCacheStore;Lsk/ainet/lang/tensor/storage/CompressedKvAttention$DequantStrategy;ILkotlin/jvm/internal/DefaultConstructorMarker;)V
public final fun isKeyCompressed ()Z
public final fun isValueCompressed ()Z
public final fun keyWindowForAttention (III)Lsk/ainet/lang/memory/WindowedKV;
public static synthetic fun keyWindowForAttention$default (Lsk/ainet/lang/tensor/storage/CompressedKvAttention;IIIILjava/lang/Object;)Lsk/ainet/lang/memory/WindowedKV;
public final fun loadKeyStorageRaw (III)Lsk/ainet/lang/tensor/storage/TensorStorage;
public static synthetic fun loadKeyStorageRaw$default (Lsk/ainet/lang/tensor/storage/CompressedKvAttention;IIIILjava/lang/Object;)Lsk/ainet/lang/tensor/storage/TensorStorage;
public final fun loadKeysForAttention (III)[F
Expand All @@ -6634,6 +6665,8 @@ public final class sk/ainet/lang/tensor/storage/CompressedKvAttention {
public final fun loadValuesForAttention (III)[F
public static synthetic fun loadValuesForAttention$default (Lsk/ainet/lang/tensor/storage/CompressedKvAttention;IIIILjava/lang/Object;)[F
public final fun storeKeyValue (I[F[F)V
public final fun valueWindowForAttention (III)Lsk/ainet/lang/memory/WindowedKV;
public static synthetic fun valueWindowForAttention$default (Lsk/ainet/lang/tensor/storage/CompressedKvAttention;IIIILjava/lang/Object;)Lsk/ainet/lang/memory/WindowedKV;
}

public final class sk/ainet/lang/tensor/storage/CompressedKvAttention$DequantStrategy : java/lang/Enum {
Expand Down Expand Up @@ -6671,7 +6704,8 @@ public final class sk/ainet/lang/tensor/storage/DefaultBufferResolver : sk/ainet
public final class sk/ainet/lang/tensor/storage/DefaultKvCacheStore : sk/ainet/lang/tensor/storage/KvCacheStore {
public fun <init> (Lsk/ainet/lang/tensor/storage/KvCacheConfig;)V
public fun <init> (Lsk/ainet/lang/tensor/storage/KvCacheConfig;Lsk/ainet/lang/memory/ModelScope;)V
public synthetic fun <init> (Lsk/ainet/lang/tensor/storage/KvCacheConfig;Lsk/ainet/lang/memory/ModelScope;ILkotlin/jvm/internal/DefaultConstructorMarker;)V
public fun <init> (Lsk/ainet/lang/tensor/storage/KvCacheConfig;Lsk/ainet/lang/memory/ModelScope;Z)V
public synthetic fun <init> (Lsk/ainet/lang/tensor/storage/KvCacheConfig;Lsk/ainet/lang/memory/ModelScope;ZILkotlin/jvm/internal/DefaultConstructorMarker;)V
public fun appendToken (I[F[F)V
public fun clear ()V
public fun evict (I)V
Expand All @@ -6685,14 +6719,18 @@ public final class sk/ainet/lang/tensor/storage/DefaultKvCacheStore : sk/ainet/l
public fun getNumLayers ()I
public fun getPlacement ()Lsk/ainet/lang/tensor/storage/Placement;
public final fun getPreallocatedBytes ()J
public final fun getSlidingWindow ()Z
public fun getValueBytesPerElement ()D
public fun getValueEncoding ()Lsk/ainet/lang/tensor/storage/TensorEncoding;
public fun getValueFormat ()Lsk/ainet/lang/memory/Format;
public final fun getWindowStart ()I
public fun keyWindow (III)Lsk/ainet/lang/memory/WindowedKV;
public fun memoryReport ()Lsk/ainet/lang/tensor/storage/KvCacheMemoryReport;
public fun readKeyStorage (III)Lsk/ainet/lang/tensor/storage/TensorStorage;
public fun readKeys (III)[F
public fun readValueStorage (III)Lsk/ainet/lang/tensor/storage/TensorStorage;
public fun readValues (III)[F
public fun valueWindow (III)Lsk/ainet/lang/memory/WindowedKV;
}

public final class sk/ainet/lang/tensor/storage/DeviceKind : java/lang/Enum {
Expand Down Expand Up @@ -6811,6 +6849,8 @@ public abstract interface class sk/ainet/lang/tensor/storage/KvCacheStore {
public fun getValueBytesPerElement ()D
public abstract fun getValueEncoding ()Lsk/ainet/lang/tensor/storage/TensorEncoding;
public fun getValueFormat ()Lsk/ainet/lang/memory/Format;
public fun keyWindow (III)Lsk/ainet/lang/memory/WindowedKV;
public static synthetic fun keyWindow$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)Lsk/ainet/lang/memory/WindowedKV;
public abstract fun memoryReport ()Lsk/ainet/lang/tensor/storage/KvCacheMemoryReport;
public abstract fun readKeyStorage (III)Lsk/ainet/lang/tensor/storage/TensorStorage;
public static synthetic fun readKeyStorage$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)Lsk/ainet/lang/tensor/storage/TensorStorage;
Expand All @@ -6820,6 +6860,8 @@ public abstract interface class sk/ainet/lang/tensor/storage/KvCacheStore {
public static synthetic fun readValueStorage$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)Lsk/ainet/lang/tensor/storage/TensorStorage;
public abstract fun readValues (III)[F
public static synthetic fun readValues$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)[F
public fun valueWindow (III)Lsk/ainet/lang/memory/WindowedKV;
public static synthetic fun valueWindow$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)Lsk/ainet/lang/memory/WindowedKV;
}

public final class sk/ainet/lang/tensor/storage/KvCacheStore$Companion {
Expand All @@ -6835,10 +6877,14 @@ public final class sk/ainet/lang/tensor/storage/KvCacheStore$DefaultImpls {
public static fun getKeyFormat (Lsk/ainet/lang/tensor/storage/KvCacheStore;)Lsk/ainet/lang/memory/Format;
public static fun getValueBytesPerElement (Lsk/ainet/lang/tensor/storage/KvCacheStore;)D
public static fun getValueFormat (Lsk/ainet/lang/tensor/storage/KvCacheStore;)Lsk/ainet/lang/memory/Format;
public static fun keyWindow (Lsk/ainet/lang/tensor/storage/KvCacheStore;III)Lsk/ainet/lang/memory/WindowedKV;
public static synthetic fun keyWindow$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)Lsk/ainet/lang/memory/WindowedKV;
public static synthetic fun readKeyStorage$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)Lsk/ainet/lang/tensor/storage/TensorStorage;
public static synthetic fun readKeys$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)[F
public static synthetic fun readValueStorage$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)Lsk/ainet/lang/tensor/storage/TensorStorage;
public static synthetic fun readValues$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)[F
public static fun valueWindow (Lsk/ainet/lang/tensor/storage/KvCacheStore;III)Lsk/ainet/lang/memory/WindowedKV;
public static synthetic fun valueWindow$default (Lsk/ainet/lang/tensor/storage/KvCacheStore;IIIILjava/lang/Object;)Lsk/ainet/lang/memory/WindowedKV;
}

public final class sk/ainet/lang/tensor/storage/LogicalDType : java/lang/Enum {
Expand Down Expand Up @@ -7350,11 +7396,13 @@ public final class sk/ainet/lang/tensor/storage/TurboQuantKvCacheStore : sk/aine
public fun getValueBytesPerElement ()D
public fun getValueEncoding ()Lsk/ainet/lang/tensor/storage/TensorEncoding;
public fun getValueFormat ()Lsk/ainet/lang/memory/Format;
public fun keyWindow (III)Lsk/ainet/lang/memory/WindowedKV;
public fun memoryReport ()Lsk/ainet/lang/tensor/storage/KvCacheMemoryReport;
public fun readKeyStorage (III)Lsk/ainet/lang/tensor/storage/TensorStorage;
public fun readKeys (III)[F
public fun readValueStorage (III)Lsk/ainet/lang/tensor/storage/TensorStorage;
public fun readValues (III)[F
public fun valueWindow (III)Lsk/ainet/lang/memory/WindowedKV;
}

public abstract interface annotation class sk/ainet/lang/tensor/storage/Weights : java/lang/annotation/Annotation {
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
package sk.ainet.lang.memory

import kotlin.math.exp
import kotlin.math.max
import kotlin.math.sqrt

/**
* Attention over a [WindowedKV] — the reference kernel that **iterates the pair** instead of
* requiring one contiguous window (SKEEP-003 §4.6, M2-F5/M2-A4).
*
* The softmax is computed online (running maximum and running denominator, the flash-attention
* recurrence), which is what makes the pair workable: positions are consumed in order, one run
* after the other, and nothing has to exist as a single array. It also means the kernel allocates
* **nothing** per token — the running accumulator is the caller's output row.
*
* A kernel that cannot do this calls [WindowedKV.gather] instead and gets one contiguous view plus
* one traced adapter; both paths are asserted to agree.
*/
@ExperimentalMemoryApi
public object WindowedAttention {

/**
* One decode step: `softmax(q·kᵀ · scale) @ v` for every head, over the window's positions in
* order.
*
* @param query `[heads * headDim]` — the current token's query rows
* @param keys the key window; `[heads, positions, headDim]` across its runs
* @param values the value window, the same shape as [keys]
* @param out `[heads * headDim]`, overwritten with the attention output
* @param scale multiplier on the scores; `0` means the usual `1/sqrt(headDim)`
*/
public fun decodeStep(
query: FloatArray,
keys: WindowedKV,
values: WindowedKV,
out: FloatArray,
scale: Float = 0f,
) {
val heads = keys.heads
val headDim = keys.headDim
require(keys.length == values.length) { "key and value windows differ: ${keys.length} vs ${values.length}" }
require(values.heads == heads && values.headDim == headDim) { "key and value windows must have the same geometry" }
require(query.size == heads * headDim) { "query must be [heads * headDim] = ${heads * headDim}, was ${query.size}" }
require(out.size == heads * headDim) { "out must be [heads * headDim] = ${heads * headDim}, was ${out.size}" }
require(keys.length > 0) { "cannot attend over an empty window" }
val s = if (scale != 0f) scale else 1f / sqrt(headDim.toFloat())

for (h in 0 until heads) {
val base = h * headDim
for (d in 0 until headDim) out[base + d] = 0f
var runningMax = Float.NEGATIVE_INFINITY
var denominator = 0f
var position = 0
// Runs in order, oldest first: head, then tail when the ring wrapped.
for (partIndex in keys.parts.indices) {
val k = keys.parts[partIndex]
val v = values.parts[partIndex]
require(k.shape[WindowedKV.POSITION_AXIS] == v.shape[WindowedKV.POSITION_AXIS]) {
"key and value runs must line up (run $partIndex)"
}
for (p in 0 until k.shape[WindowedKV.POSITION_AXIS]) {
var score = 0f
for (d in 0 until headDim) score += query[base + d] * k.get(h, p, d)
score *= s

val newMax = max(runningMax, score)
val correction = if (runningMax == Float.NEGATIVE_INFINITY) 0f else exp(runningMax - newMax)
val weight = exp(score - newMax)
denominator = denominator * correction + weight
for (d in 0 until headDim) {
out[base + d] = out[base + d] * correction + weight * v.get(h, p, d)
}
runningMax = newMax
position++
}
}
for (d in 0 until headDim) out[base + d] = out[base + d] / denominator
}
}

/**
* [decodeStep] through the gather adapter: the window is copied into [scope] as one contiguous
* view first, for kernels that cannot iterate a pair. Same numbers, one allocation and one
* traced adapter per call.
*/
public fun decodeStepGathered(
query: FloatArray,
keys: WindowedKV,
values: WindowedKV,
out: FloatArray,
scope: Scope,
sink: sk.ainet.lang.memory.trace.TraceSink = sk.ainet.lang.memory.trace.NoopTraceSink,
scale: Float = 0f,
) {
decodeStep(query, WindowedKV(keys.gather(scope, sink)), WindowedKV(values.gather(scope, sink)), out, scale)
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
package sk.ainet.lang.memory

import sk.ainet.lang.memory.trace.NoopTraceSink
import sk.ainet.lang.memory.trace.TraceEvent
import sk.ainet.lang.memory.trace.TraceSink
import sk.ainet.lang.tensor.Shape
import sk.ainet.lang.tensor.TensorId

/**
* An attention window over a KV ring, as the **one or two pieces it physically is** (SKEEP-003
* §4.6, decision #4; M2-F5).
*
* A ring that has wrapped holds its newest `n` positions in two runs — `[from, capacity)` and
* `[0, wrapEnd)` — and the choice has always been to copy them together before attention or to
* grow the buffer forever. This is the third option: hand the kernel the pair. Both halves are
* ordinary [TensorView]s over the same [Storage], so nothing is copied and nothing is special-cased;
* a kernel that cannot iterate a pair calls [gather], which makes the copy *visible* as one
* adapter in the trace instead of an invisible per-token allocation.
*
* Position order is `head` then `tail`: oldest to newest.
*
* @property head the first (older) run — the only one when the window does not wrap
* @property tail the second (newer) run, present only after the ring wraps
*/
@ExperimentalMemoryApi
public class WindowedKV(
public val head: TensorView,
public val tail: TensorView? = null,
) {
init {
val t = tail
if (t != null) {
require(t.shape.rank == head.shape.rank) { "both halves must have the same rank" }
require(t.format == head.format) { "both halves must have the same format" }
for (axis in 0 until head.shape.rank) {
if (axis == POSITION_AXIS) continue
require(t.shape[axis] == head.shape[axis]) { "halves differ on axis $axis: ${head.shape} vs ${t.shape}" }
}
}
}

/** The halves in position order, oldest first. */
public val parts: List<TensorView> get() = if (tail == null) listOf(head) else listOf(head, tail)

/** Positions in the window, across both halves. */
public val length: Int get() = head.shape[POSITION_AXIS] + (tail?.shape?.get(POSITION_AXIS) ?: 0)

/** True when the ring wrapped inside this window — the case the pair exists for. */
public val wrapped: Boolean get() = tail != null

/** `[heads, dim]` — the shape of the window with its position axis removed. */
public val heads: Int get() = head.shape[0]
public val headDim: Int get() = head.shape[head.shape.rank - 1]

/**
* The window as one contiguous view in [scope] — the gather adapter (§5.1).
*
* For a kernel that cannot iterate the pair. The copy is real, so it is *traced*: one
* [TraceEvent.AdapterInserted] per call, which is what makes "the zero-copy path allocates
* nothing per token" an assertion rather than a hope.
*/
public fun gather(scope: Scope, sink: TraceSink = NoopTraceSink, id: TensorId? = null): TensorView {
val total = length
val shape = Shape(heads, total, headDim)
val storage = scope.allocateFloats(heads * total * headDim, id)
val out = storage.floats!!
val base = storage.arrayOffset
var written = 0
for (part in parts) {
val positions = part.shape[POSITION_AXIS]
for (h in 0 until heads) {
for (p in 0 until positions) {
val dst = base + (h.toLong() * total + written + p).toInt() * headDim
for (d in 0 until headDim) out[dst + d] = part.get(h, p, d)
}
}
written += positions
}
if (sink.isEnabled) {
sink.emit(
TraceEvent.AdapterInserted(
kind = "gather-kv-window",
from = head.format,
to = Format.dense(head.format.dtype),
bytes = heads.toLong() * total * headDim * head.format.dtype.sizeInBytes,
target = id ?: head.id,
scope = scope.kind,
),
)
}
return TensorView.dense(storage, shape, head.format.dtype, id ?: head.id)
}

/** The decoded value at (`head`, absolute window position, `dim`), crossing the halves. */
public fun get(head: Int, position: Int, dim: Int): Float {
require(position in 0 until length) { "position $position outside the window (length $length)" }
val first = this.head.shape[POSITION_AXIS]
return if (position < first) this.head.get(head, position, dim)
else tail!!.get(head, position - first, dim)
}

override fun toString(): String =
"WindowedKV(${heads}h × $length × $headDim, ${if (wrapped) "wrapped: ${head.shape[POSITION_AXIS]}+${tail!!.shape[POSITION_AXIS]}" else "contiguous"})"

public companion object {
/** KV windows are `[heads, positions, headDim]`. */
public const val POSITION_AXIS: Int = 1
}
}
Loading
Loading