spark_model/layers/ops/
token_overlay.rs1use anyhow::Result;
14use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
15use spark_runtime::kernel_args::KernelLaunch;
16
17use crate::layers::try_kernel;
18
19#[derive(Clone, Copy)]
23pub struct OverlayKernels {
24 pub rowdiff: KernelHandle,
25 pub embed_overlay: KernelHandle,
26 pub lmhead_overlay_bf16: KernelHandle,
27 pub lmhead_overlay_f32: KernelHandle,
28}
29
30impl Default for OverlayKernels {
31 fn default() -> Self {
35 Self {
36 rowdiff: KernelHandle(0),
37 embed_overlay: KernelHandle(0),
38 lmhead_overlay_bf16: KernelHandle(0),
39 lmhead_overlay_f32: KernelHandle(0),
40 }
41 }
42}
43
44impl OverlayKernels {
45 pub fn new(gpu: &dyn GpuBackend) -> Self {
46 Self {
47 rowdiff: try_kernel(gpu, "token_overlay", "embed_rowdiff_bf16"),
48 embed_overlay: try_kernel(gpu, "token_overlay", "embed_overlay_routed_bf16"),
49 lmhead_overlay_bf16: try_kernel(gpu, "token_overlay", "lmhead_overlay_routed_bf16"),
50 lmhead_overlay_f32: try_kernel(gpu, "token_overlay", "lmhead_overlay_routed_f32"),
51 }
52 }
53}
54
55#[allow(clippy::too_many_arguments)]
57pub fn embed_rowdiff(
58 gpu: &dyn GpuBackend,
59 kernel: KernelHandle,
60 base: DevicePtr, served: DevicePtr, flags: DevicePtr, rows: u32,
64 h: u32,
65 thresh: f32,
66 stream: u64,
67) -> Result<()> {
68 KernelLaunch::new(gpu, kernel)
69 .grid([rows.div_ceil(256), 1, 1])
70 .block([256, 1, 1])
71 .arg_ptr(base)
72 .arg_ptr(served)
73 .arg_ptr(flags)
74 .arg_u32(rows)
75 .arg_u32(h)
76 .arg_f32(thresh)
77 .launch(stream)
78}
79
80#[allow(clippy::too_many_arguments)]
87pub fn embed_overlay_routed(
88 gpu: &dyn GpuBackend,
89 kernel: KernelHandle,
90 ids: DevicePtr, seq_slot: DevicePtr, active: i32,
93 slot_map_tab: DevicePtr, rows_tab: DevicePtr, n_tab: DevicePtr, out: DevicePtr, num_tokens: u32,
98 h: u32,
99 vocab: u32, stream: u64,
101) -> Result<()> {
102 KernelLaunch::new(gpu, kernel)
103 .grid([num_tokens, 1, 1])
104 .block([256, 1, 1])
105 .arg_ptr(ids)
106 .arg_ptr(seq_slot)
107 .arg_i32(active)
108 .arg_ptr(slot_map_tab)
109 .arg_ptr(rows_tab)
110 .arg_ptr(n_tab)
111 .arg_ptr(out)
112 .arg_u32(h)
113 .arg_u32(vocab)
114 .launch(stream)
115}
116
117#[allow(clippy::too_many_arguments)]
121pub fn lmhead_overlay_routed(
122 gpu: &dyn GpuBackend,
123 kernel: KernelHandle, hidden: DevicePtr, seq_slot: DevicePtr, active: i32,
127 rows_tab: DevicePtr, ids_tab: DevicePtr, n_tab: DevicePtr, logits: DevicePtr, m: u32,
132 max_n_override: u32,
133 h: u32,
134 vocab: u32,
135 stream: u64,
136) -> Result<()> {
137 KernelLaunch::new(gpu, kernel)
138 .grid([m, max_n_override, 1])
139 .block([32, 1, 1])
140 .arg_ptr(hidden)
141 .arg_ptr(seq_slot)
142 .arg_i32(active)
143 .arg_ptr(rows_tab)
144 .arg_ptr(ids_tab)
145 .arg_ptr(n_tab)
146 .arg_ptr(logits)
147 .arg_u32(h)
148 .arg_u32(vocab)
149 .launch(stream)
150}