The model code was spread across eight crates that had grown into each other:
ggml and cuda and mlx each owned part of a tensor runtime, llama and tts and
voice2 each owned part of a model, and libs/diffusion owned everything else.
They are now one tree with an explicit shape:
libs/ai/cuda — kernels and launch surface
libs/ai/metal — Metal shaders and the shim
libs/ai/llm — the language-model runtime (sessions, lanes, contexts,
the CUDA and Metal executors, the compiled Metal path)
libs/ai/models/ — common, flux, h3, music, paint, speech, stems, vision
libs/diffusion is not deleted but demoted: what remains is the VALIDATOR
crate — several dozen `*_validate.rs` oracles that check a native
implementation against a reference, which is where they belong now that the
implementations live next door.
The functional work inside the move is mostly in the LLM runtime: N lanes that
draft while one verify batch serves all of them, per-slot prefill over a shared
folded attention arena, speculation that survives batching, and a scheduler
that reports rather than publishes. And in the CUDA build: a machine without
usable CUDA must still LINK (and say so), the default kernel arch is the
building machine's GPU, `NO_CUDA` forces the stub even where the toolkit
exists, and kernels compile in parallel with progress.
libs/video_flow is new here: classical optical flow estimation and the `mkfl`
motion-field payload — a flow field measured from a clip without a model,
which is what drives free-rate bounce-looping playback and the uprez/tween
enhance pipe.
211 lines
8.5 KiB
Text
211 lines
8.5 KiB
Text
// SPDX-License-Identifier: MIT
|
|
// Copyright (c) 2023-2026 The ggml authors
|
|
//
|
|
// Substantial portions derived from ggml / llama.cpp
|
|
// (https://github.com/ggml-org/llama.cpp), MIT licensed.
|
|
// The original copyright notice and permission notice are retained.
|
|
// See libs/ai/NOTICE and, where present, LICENSE in this directory.
|
|
//
|
|
#pragma once
|
|
// Copied from llama.cpp ggml/src/ggml-cuda/gated_delta_net.cu:3-201
|
|
// (kernel + launch_gated_delta_net). Host ggml_cuda_info trimmed to
|
|
// warp_size=32. No chunked kernel — they have none (cu:159).
|
|
#include "common.cuh"
|
|
|
|
template <int S_v, bool KDA>
|
|
__global__ void __launch_bounds__((ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v) * 4, 2)
|
|
gated_delta_net_cuda(const float * q,
|
|
const float * k,
|
|
const float * v,
|
|
const float * g,
|
|
const float * beta,
|
|
const float * curr_state,
|
|
float * dst,
|
|
int64_t H,
|
|
int64_t n_tokens,
|
|
int64_t n_seqs,
|
|
int64_t sq1,
|
|
int64_t sq2,
|
|
int64_t sq3,
|
|
int64_t sv1,
|
|
int64_t sv2,
|
|
int64_t sv3,
|
|
int64_t sb1,
|
|
int64_t sb2,
|
|
int64_t sb3,
|
|
const uint3 neqk1_magic,
|
|
const uint3 rq3_magic,
|
|
float scale,
|
|
int64_t state_ckpt_stride) {
|
|
const uint32_t h_idx = blockIdx.x;
|
|
const uint32_t sequence = blockIdx.y;
|
|
const int lane = threadIdx.x;
|
|
const int col = blockIdx.z * blockDim.y + threadIdx.y;
|
|
|
|
const uint32_t iq1 = fastmodulo(h_idx, neqk1_magic);
|
|
const uint32_t iq3 = fastdiv(sequence, rq3_magic);
|
|
|
|
const int64_t attn_score_elems = S_v * H * n_tokens * n_seqs;
|
|
float * attn_data = dst;
|
|
float * state = dst + attn_score_elems;
|
|
|
|
const int64_t state_offset = (sequence * H + h_idx) * S_v * S_v;
|
|
state += state_offset;
|
|
curr_state += state_offset + col * S_v;
|
|
attn_data += (sequence * n_tokens * H + h_idx) * S_v;
|
|
|
|
constexpr int warp_size = ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v;
|
|
static_assert(S_v % warp_size == 0, "S_v must be a multiple of warp_size");
|
|
constexpr int rows_per_lane = (S_v + warp_size - 1) / warp_size;
|
|
float s_shard[rows_per_lane];
|
|
|
|
#pragma unroll
|
|
for (int r = 0; r < rows_per_lane; r++) {
|
|
const int i = r * warp_size + lane;
|
|
s_shard[r] = curr_state[i];
|
|
}
|
|
|
|
for (int t = 0; t < n_tokens; t++) {
|
|
const float * q_t = q + iq3 * sq3 + t * sq2 + iq1 * sq1;
|
|
const float * k_t = k + iq3 * sq3 + t * sq2 + iq1 * sq1;
|
|
const float * v_t = v + sequence * sv3 + t * sv2 + h_idx * sv1;
|
|
|
|
const int64_t gb_offset = sequence * sb3 + t * sb2 + h_idx * sb1;
|
|
const float * beta_t = beta + gb_offset;
|
|
const float * g_t = g + gb_offset * (KDA ? S_v : 1);
|
|
|
|
const float beta_val = *beta_t;
|
|
|
|
float k_reg[rows_per_lane];
|
|
float q_reg[rows_per_lane];
|
|
#pragma unroll
|
|
for (int r = 0; r < rows_per_lane; r++) {
|
|
const int i = r * warp_size + lane;
|
|
k_reg[r] = k_t[i];
|
|
q_reg[r] = q_t[i];
|
|
}
|
|
|
|
if constexpr (!KDA) {
|
|
const float g_val = expf(*g_t);
|
|
|
|
float kv_shard = 0.0f;
|
|
#pragma unroll
|
|
for (int r = 0; r < rows_per_lane; r++) {
|
|
kv_shard += s_shard[r] * k_reg[r];
|
|
}
|
|
float kv_col = warp_reduce_sum<warp_size>(kv_shard);
|
|
|
|
float delta_col = (v_t[col] - g_val * kv_col) * beta_val;
|
|
|
|
float attn_partial = 0.0f;
|
|
#pragma unroll
|
|
for (int r = 0; r < rows_per_lane; r++) {
|
|
s_shard[r] = g_val * s_shard[r] + k_reg[r] * delta_col;
|
|
attn_partial += s_shard[r] * q_reg[r];
|
|
}
|
|
|
|
float attn_col = warp_reduce_sum<warp_size>(attn_partial);
|
|
|
|
if (lane == 0) {
|
|
attn_data[col] = attn_col * scale;
|
|
}
|
|
} else {
|
|
float kv_shard = 0.0f;
|
|
#pragma unroll
|
|
for (int r = 0; r < rows_per_lane; r++) {
|
|
const int i = r * warp_size + lane;
|
|
kv_shard += expf(g_t[i]) * s_shard[r] * k_reg[r];
|
|
}
|
|
|
|
float kv_col = warp_reduce_sum<warp_size>(kv_shard);
|
|
|
|
float delta_col = (v_t[col] - kv_col) * beta_val;
|
|
|
|
float attn_partial = 0.0f;
|
|
#pragma unroll
|
|
for (int r = 0; r < rows_per_lane; r++) {
|
|
const int i = r * warp_size + lane;
|
|
s_shard[r] = expf(g_t[i]) * s_shard[r] + k_reg[r] * delta_col;
|
|
attn_partial += s_shard[r] * q_reg[r];
|
|
}
|
|
|
|
float attn_col = warp_reduce_sum<warp_size>(attn_partial);
|
|
|
|
if (lane == 0) {
|
|
attn_data[col] = attn_col * scale;
|
|
}
|
|
}
|
|
|
|
// Speculative verification needs the state after EVERY token, not
|
|
// just the last one, so a rejected draft can be undone by resuming
|
|
// from an earlier row. Emitting them from this one call keeps the
|
|
// checkpointed graph at the node count of an ordinary decode.
|
|
if (state_ckpt_stride != 0) {
|
|
#pragma unroll
|
|
for (int r = 0; r < rows_per_lane; r++) {
|
|
const int i = r * warp_size + lane;
|
|
state[t * state_ckpt_stride + col * S_v + i] = s_shard[r];
|
|
}
|
|
}
|
|
|
|
attn_data += S_v * H;
|
|
}
|
|
|
|
if (state_ckpt_stride == 0) {
|
|
#pragma unroll
|
|
for (int r = 0; r < rows_per_lane; r++) {
|
|
const int i = r * warp_size + lane;
|
|
state[col * S_v + i] = s_shard[r];
|
|
}
|
|
}
|
|
}
|
|
|
|
template <bool KDA>
|
|
static void launch_gated_delta_net(
|
|
const float * q_d, const float * k_d, const float * v_d,
|
|
const float * g_d, const float * b_d, const float * s_d,
|
|
float * dst_d,
|
|
int64_t S_v, int64_t H, int64_t n_tokens, int64_t n_seqs,
|
|
int64_t sq1, int64_t sq2, int64_t sq3,
|
|
int64_t sv1, int64_t sv2, int64_t sv3,
|
|
int64_t sb1, int64_t sb2, int64_t sb3,
|
|
int64_t neqk1, int64_t rq3,
|
|
float scale, int64_t state_ckpt_stride, cudaStream_t stream) {
|
|
const int warp_size = 32;
|
|
const int num_warps = 4;
|
|
dim3 grid_dims((unsigned) H, (unsigned) n_seqs, (unsigned) ((S_v + num_warps - 1) / num_warps));
|
|
dim3 block_dims(warp_size <= S_v ? warp_size : (int) S_v, num_warps, 1);
|
|
|
|
const uint3 neqk1_magic = init_fastdiv_values((uint32_t) neqk1);
|
|
const uint3 rq3_magic = init_fastdiv_values((uint32_t) rq3);
|
|
|
|
switch (S_v) {
|
|
case 16:
|
|
gated_delta_net_cuda<16, KDA><<<grid_dims, block_dims, 0, stream>>>(
|
|
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H,
|
|
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
|
|
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_ckpt_stride);
|
|
break;
|
|
case 32:
|
|
gated_delta_net_cuda<32, KDA><<<grid_dims, block_dims, 0, stream>>>(
|
|
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H,
|
|
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
|
|
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_ckpt_stride);
|
|
break;
|
|
case 64:
|
|
gated_delta_net_cuda<64, KDA><<<grid_dims, block_dims, 0, stream>>>(
|
|
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H,
|
|
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
|
|
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_ckpt_stride);
|
|
break;
|
|
case 128:
|
|
gated_delta_net_cuda<128, KDA><<<grid_dims, block_dims, 0, stream>>>(
|
|
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H,
|
|
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
|
|
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_ckpt_stride);
|
|
break;
|
|
default:
|
|
GGML_ABORT("unsupported GDN S_v");
|
|
break;
|
|
}
|
|
}
|