1use cortiq_core::CmfModel;
14use std::cell::Cell;
15use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering};
16use std::sync::{Arc, OnceLock};
17
18pub struct EmbryoGraphModel {
27 pub id: u64,
28 pub hidden: usize,
29 pub intermediate: usize,
30 pub vocab: usize,
31 pub layers: usize,
32 pub phase_heads: usize,
33 pub nphase: usize,
34 pub phase_dv: usize,
35 pub anchor_q_heads: usize,
36 pub anchor_kv_heads: usize,
37 pub anchor_head_dim: usize,
38 pub rotary_dim: usize,
39 pub max_seq: usize,
40 pub cluster_count: usize,
41 pub cluster_size: usize,
42 pub phase_state_len: usize,
43 pub state_stride: usize,
44 pub kv_stride: usize,
45 pub norm_gemma: bool,
46 pub phase_mass: f32,
47 pub weights: Vec<f32>,
48 pub meta: Vec<u32>,
49 pub lm_head: Vec<f32>,
50 pub clusters: Vec<f32>,
51 pub final_norm: Vec<f32>,
52 pub inv_freq: Vec<f32>,
53 pub bounded: bool,
58 pub kv_layers: usize,
62 pub state_layers: usize,
65 pub anchor_window: usize,
66 pub anchor_sink: usize,
67 pub phase_layers: usize,
69 pub gdn_layers: usize,
73 pub gdn_heads: usize,
74 pub gdn_k_heads: usize,
75 pub gdn_dk: usize,
76 pub gdn_dv: usize,
77 pub gdn_kk: usize,
78}
79
80impl EmbryoGraphModel {
81 pub fn gdn_c_dim(&self) -> usize {
83 2 * self.gdn_k_heads * self.gdn_dk + self.gdn_heads * self.gdn_dv
84 }
85}
86
87pub const EMBRYO_META_HEADER: usize = 32;
90
91pub const EMBRYO_CHUNK_MAX: usize = 64;
95
96pub fn forward_embryo_graph_chunk(
101 model: &Arc<EmbryoGraphModel>,
102 kv_id: u64,
103 rows: &[f32],
104 position: usize,
105 n: usize,
106 logits: &mut Vec<f32>,
107) -> bool {
108 #[cfg(feature = "gpu")]
109 if backend() == Backend::Wgpu {
110 return crate::gpu_wgpu::forward_embryo_graph_chunk(model, kv_id, rows, position, n, logits);
111 }
112 let _ = (model, kv_id, rows, position, n, logits);
113 false
114}
115
116pub fn embryo_device_state_bytes(kv_id: u64) -> Option<(u64, u64)> {
122 #[cfg(feature = "gpu")]
123 if backend() == Backend::Wgpu {
124 return crate::gpu_wgpu::embryo_device_state_bytes(kv_id);
125 }
126 let _ = kv_id;
127 None
128}
129
130pub fn counts_as_placement_weight(model: &CmfModel, name: &str) -> bool {
137 model.header.genome.is_none() || cortiq_core::knowledge::is_trunk_tensor(name)
138}
139
140pub fn placement_weight_bytes(model: &CmfModel) -> u64 {
144 if model.header.genome.is_none() {
145 return model.primary_bytes().len() as u64;
146 }
147 model
148 .tensors
149 .iter()
150 .filter(|t| counts_as_placement_weight(model, &t.name))
151 .map(|t| t.nbytes)
152 .sum()
153}
154
155pub fn embryo_device_next_position(kv_id: u64) -> Option<usize> {
159 #[cfg(feature = "gpu")]
160 if backend() == Backend::Wgpu {
161 return crate::gpu_wgpu::embryo_device_next_position(kv_id);
162 }
163 let _ = kv_id;
164 None
165}
166
167thread_local! {
168 static CUR_LAYER: Cell<i64> = const { Cell::new(-1) };
172 static CPU_ONLY: Cell<bool> = const { Cell::new(false) };
177 static PROBE_COLD: Cell<bool> = const { Cell::new(false) };
181}
182
183pub struct CpuScopeGuard(bool);
189
190impl Drop for CpuScopeGuard {
191 fn drop(&mut self) {
192 CPU_ONLY.with(|c| c.set(self.0));
193 }
194}
195
196pub fn enter_cpu_scope() -> CpuScopeGuard {
197 let previous = CPU_ONLY.with(|c| c.replace(true));
198 CpuScopeGuard(previous)
199}
200
201pub fn cpu_scope<R>(f: impl FnOnce() -> R) -> R {
203 let _restore = enter_cpu_scope();
204 f()
205}
206
207pub(crate) fn inherit_cpu_scope() -> impl Fn() -> Option<CpuScopeGuard> + Copy {
212 let on = CPU_ONLY.get();
213 move || on.then(enter_cpu_scope)
214}
215
216pub fn probe_set_device(label: &str) {
221 let _ = DEVICE_LABEL.set(label.to_string());
222}
223
224fn device_label() -> &'static str {
225 DEVICE_LABEL.get().map(String::as_str).unwrap_or("unknown")
226}
227
228static DEVICE_LABEL: std::sync::OnceLock<String> = std::sync::OnceLock::new();
229
230static CACHE_DIR: std::sync::OnceLock<std::path::PathBuf> = std::sync::OnceLock::new();
239
240pub fn set_cache_dir(dir: std::path::PathBuf) {
242 let _ = CACHE_DIR.set(dir);
243}
244
245pub fn cache_dir_pub() -> std::path::PathBuf {
247 cache_dir()
248}
249
250fn cache_dir() -> std::path::PathBuf {
251 if let Some(d) = CACHE_DIR.get() {
252 return d.clone();
253 }
254 match std::env::var_os("TMPDIR") {
255 Some(t) => std::path::PathBuf::from(t),
256 None => std::env::temp_dir(),
257 }
258}
259
260fn probe_cache_path() -> Option<std::path::PathBuf> {
263 match std::env::var("CMF_PROBE_CACHE") {
264 Ok(v) if v == "0" => None,
265 Ok(v) => Some(std::path::PathBuf::from(v)),
266 Err(_) => Some(cache_dir().join("cortiq-gpu-probe.tsv")),
267 }
268}
269
270fn probe_cache_key_named(class: &str) -> String {
274 format!(
275 "{}\t{}\t{}",
276 env!("CARGO_PKG_VERSION"),
277 device_label(),
278 class
279 )
280}
281
282const CLASS_NAMES: [&str; 7] = [
283 "ffn",
284 "matvec",
285 "matmat",
286 "qkv-batch",
287 "matmat-wide",
288 "lm-head",
289 "gemm-nt",
290];
291
292fn probe_cache_load() {
301 static ONCE: std::sync::Once = std::sync::Once::new();
302 ONCE.call_once(|| {
303 let Some(path) = probe_cache_path() else {
304 return;
305 };
306 if cfg!(test) && std::env::var("CMF_PROBE_CACHE").is_err() {
311 return;
312 }
313 let Ok(text) = std::fs::read_to_string(&path) else {
314 return;
315 };
316 probe_cache_adopt(&text);
317 });
318}
319
320fn probe_cache_adopt(text: &str) {
324 for line in text.lines() {
325 let Some((key, verdict)) = line.rsplit_once('\t') else {
326 continue;
327 };
328 let winner = match verdict.trim() {
329 "gpu" => 1u8,
330 "cpu" => 2u8,
331 _ => continue,
332 };
333 for (i, name) in CLASS_NAMES.iter().enumerate() {
334 if probe_cache_key_named(name) == key {
335 let _ = PROBES[i].state.compare_exchange(
336 0,
337 winner,
338 Ordering::Relaxed,
339 Ordering::Relaxed,
340 );
341 tracing::debug!("gpu probe [{name}]: remembered → {verdict}");
342 }
343 }
344 }
345}
346
347fn probe_cache_store(c: OpClass, winner: u8) {
350 let Some(path) = probe_cache_path() else {
351 return;
352 };
353 let line = format!(
354 "{}\t{}\n",
355 probe_cache_key_named(CLASS_NAMES[c as usize]),
356 if winner == 1 { "gpu" } else { "cpu" }
357 );
358 use std::io::Write;
359 if let Ok(mut f) = std::fs::OpenOptions::new()
360 .create(true)
361 .append(true)
362 .open(&path)
363 {
364 let _ = f.write_all(line.as_bytes());
365 }
366}
367
368pub fn cold_epoch() -> u64 {
374 COLD_EPOCH.load(std::sync::atomic::Ordering::Relaxed)
375}
376static COLD_EPOCH: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
377
378pub(crate) fn probe_note_cold() {
379 COLD_EPOCH.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
380 PROBE_COLD.with(|c| c.set(true));
381}
382
383pub(crate) fn probe_was_cold() -> bool {
387 PROBE_COLD.with(|c| c.get())
388}
389
390pub fn set_layer(l: i64) {
392 CUR_LAYER.with(|c| c.set(l));
393}
394
395pub fn cur_layer() -> i64 {
397 CUR_LAYER.with(|c| c.get())
398}
399
400pub fn automatic_layer_prefix(
403 model: &Arc<CmfModel>,
404 num_layers: usize,
405 physical_layers: usize,
406) -> Option<usize> {
407 match backend() {
408 #[cfg(feature = "gpu")]
409 Backend::Wgpu => {
410 crate::gpu_wgpu::automatic_layer_prefix(model, num_layers, physical_layers)
411 }
412 _ => None,
413 }
414}
415
416fn layer_ranges() -> &'static Option<Vec<(i64, i64)>> {
419 static R: OnceLock<Option<Vec<(i64, i64)>>> = OnceLock::new();
420 R.get_or_init(|| {
421 let s = std::env::var("CMF_GPU_LAYERS").ok()?;
422 let mut v = Vec::new();
423 for part in s.split(',') {
424 let part = part.trim();
425 match part.split_once('-') {
426 Some((a, b)) => v.push((a.trim().parse().ok()?, b.trim().parse().ok()?)),
427 None => {
428 let x: i64 = part.parse().ok()?;
429 v.push((x, x));
430 }
431 }
432 }
433 Some(v)
434 })
435}
436
437fn layer_allowed() -> bool {
438 match layer_ranges() {
439 None => true,
440 Some(ranges) => {
441 let cur = CUR_LAYER.with(|c| c.get());
442 cur < 0 || ranges.iter().any(|(a, b)| cur >= *a && cur <= *b)
443 }
444 }
445}
446
447pub fn enabled_here() -> bool {
451 !CPU_ONLY.with(|c| c.get()) && enabled() && layer_allowed()
452}
453
454pub fn q2tp_gpu_opt_in() -> bool {
460 std::env::var("CMF_Q2TP_GPU").as_deref() == Ok("1")
461}
462
463#[derive(Clone, Copy)]
475pub enum OpClass {
476 Ffn = 0,
478 Matvec = 1,
480 Matmat = 2,
482 Batch = 3,
484 MatmatWide = 4,
490 MatvecHead = 5,
497 GemmNt = 6,
504}
505
506pub fn matvec_class(rows: usize, cols: usize) -> OpClass {
510 if rows * cols >= 67_108_864 {
511 OpClass::MatvecHead
512 } else {
513 OpClass::Matvec
514 }
515}
516
517pub enum ProbeArm {
519 Gpu,
521 CpuTimed,
523 Cpu,
525}
526
527const PROBE_SAMPLES: u32 = 6;
529
530const PROBE_DECLINE_LIMIT: u32 = 16;
534
535const PROBE_WARMUP: u32 = 1;
537
538struct Probe {
539 state: AtomicU8,
541 flip: AtomicU32,
542 gpu_ns: AtomicU64,
543 gpu_n: AtomicU32,
544 declines: AtomicU32,
554 gpu_burn: AtomicU32,
567 cpu_ns: AtomicU64,
568 cpu_n: AtomicU32,
569 gpu_min: AtomicU64,
576 cpu_min: AtomicU64,
577}
578
579impl Probe {
580 const fn new() -> Self {
581 Self {
582 state: AtomicU8::new(0),
583 flip: AtomicU32::new(0),
584 gpu_ns: AtomicU64::new(0),
585 gpu_n: AtomicU32::new(0),
586 declines: AtomicU32::new(0),
587 gpu_burn: AtomicU32::new(PROBE_WARMUP),
588 cpu_ns: AtomicU64::new(0),
589 cpu_n: AtomicU32::new(0),
590 gpu_min: AtomicU64::new(u64::MAX),
591 cpu_min: AtomicU64::new(u64::MAX),
592 }
593 }
594}
595
596static PROBES: [Probe; 7] = [
597 Probe::new(),
598 Probe::new(),
599 Probe::new(),
600 Probe::new(),
601 Probe::new(),
602 Probe::new(),
603 Probe::new(),
604];
605
606static TRUST_GPU: AtomicBool = AtomicBool::new(false);
613
614pub fn trust_gpu() -> GpuTrust {
616 let was = TRUST_GPU.swap(true, Ordering::Relaxed);
617 GpuTrust(was)
618}
619
620pub struct GpuTrust(bool);
621
622impl Drop for GpuTrust {
623 fn drop(&mut self) {
624 TRUST_GPU.store(self.0, Ordering::Relaxed);
625 }
626}
627
628fn probe_on_for(c: OpClass) -> bool {
629 if TRUST_GPU.load(Ordering::Relaxed) && matches!(c, OpClass::MatmatWide | OpClass::Ffn) {
635 return false;
636 }
637 probe_on()
638}
639
640pub fn probe_enabled() -> bool {
644 probe_on()
645}
646
647fn probe_on() -> bool {
648 static ON: OnceLock<bool> = OnceLock::new();
649 *ON.get_or_init(|| {
650 std::env::var("CMF_GPU_PROBE")
651 .map(|v| v != "0" && v != "off")
652 .unwrap_or(true)
653 })
654}
655
656pub fn q1_force() -> bool {
661 #[cfg(target_os = "macos")]
662 {
663 backend() == Backend::Metal
664 }
665 #[cfg(not(target_os = "macos"))]
666 {
667 false
668 }
669}
670
671pub fn fused_block_trusted() -> bool {
690 #[cfg(target_os = "macos")]
691 if backend() == Backend::Metal {
692 return true;
693 }
694 wgpu_graph_default()
695}
696
697pub fn weight_is_resident(model: &Arc<CmfModel>, idx: usize) -> bool {
709 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
710 {
711 return crate::gpu_wgpu::weight_is_resident(model, idx);
712 }
713 #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
714 {
715 let _ = (model, idx);
716 true
717 }
718}
719
720pub fn probe_arm_cold_prefers_gpu(c: OpClass, weights_resident: bool) -> ProbeArm {
721 if !weights_resident && probe_deciding(c) {
722 return ProbeArm::Gpu;
723 }
724 probe_arm(c)
725}
726
727pub fn probe_arm(c: OpClass) -> ProbeArm {
728 PROBE_COLD.with(|f| f.set(false));
733 if !probe_on_for(c) {
734 return ProbeArm::Gpu;
735 }
736 probe_cache_load();
737 let p = &PROBES[c as usize];
738 match p.state.load(Ordering::Relaxed) {
739 1 => ProbeArm::Gpu,
740 2 => ProbeArm::Cpu,
741 _ => {
742 if p.flip.fetch_add(1, Ordering::Relaxed) % 2 == 0 {
743 ProbeArm::Gpu
744 } else {
745 ProbeArm::CpuTimed
746 }
747 }
748 }
749}
750
751pub fn probe_note_decline(c: OpClass) {
755 let p = &PROBES[c as usize];
756 if p.state.load(Ordering::Relaxed) != 0 {
757 return;
758 }
759 let n = p.declines.fetch_add(1, Ordering::Relaxed) + 1;
760 if n >= PROBE_DECLINE_LIMIT
761 && p.state
762 .compare_exchange(0, 2, Ordering::Relaxed, Ordering::Relaxed)
763 .is_ok()
764 {
765 tracing::info!(
766 "gpu probe [{}]: device declined {n} times → cpu",
767 CLASS_NAMES[c as usize]
768 );
769 }
770}
771
772pub fn probe_record(c: OpClass, gpu: bool, dur: std::time::Duration) {
775 probe_record_into(
776 &PROBES[c as usize],
777 CLASS_NAMES[c as usize],
778 Some(c),
779 gpu,
780 dur,
781 )
782}
783
784fn probe_record_into(
787 p: &Probe,
788 class_name: &str,
789 cache: Option<OpClass>,
790 gpu: bool,
791 dur: std::time::Duration,
792) {
793 if p.state.load(Ordering::Relaxed) != 0 {
794 return;
795 }
796 if gpu && PROBE_COLD.with(|f| f.replace(false)) {
797 return; }
799 if gpu {
800 let left = p.gpu_burn.load(Ordering::Relaxed);
804 if left > 0 {
805 p.gpu_burn.store(left - 1, Ordering::Relaxed);
806 return; }
808 }
809 let ns = dur.as_nanos().min(u64::MAX as u128) as u64;
810 if gpu {
811 p.gpu_ns.fetch_add(ns, Ordering::Relaxed);
812 p.gpu_n.fetch_add(1, Ordering::Relaxed);
813 p.gpu_min.fetch_min(ns, Ordering::Relaxed);
814 } else {
815 p.cpu_ns.fetch_add(ns, Ordering::Relaxed);
816 p.cpu_n.fetch_add(1, Ordering::Relaxed);
817 p.cpu_min.fetch_min(ns, Ordering::Relaxed);
818 }
819 let (gn, cn) = (
820 p.gpu_n.load(Ordering::Relaxed),
821 p.cpu_n.load(Ordering::Relaxed),
822 );
823 if gn >= 2 && cn >= 2 {
824 let g = p.gpu_min.load(Ordering::Relaxed) as f64;
828 let cp = p.cpu_min.load(Ordering::Relaxed) as f64;
829 if (gn < PROBE_SAMPLES || cn < PROBE_SAMPLES) && g < cp * 2.0 && cp < g * 2.0 {
839 return;
840 }
841 let winner = if g <= cp { 1 } else { 2 };
842 if p.state
843 .compare_exchange(0, winner, Ordering::Relaxed, Ordering::Relaxed)
844 .is_ok()
845 {
846 tracing::info!(
847 "gpu probe [{}]: gpu {:.2} ms vs cpu {:.2} ms per op → {}",
848 class_name,
849 g / 1e6,
850 cp / 1e6,
851 if winner == 1 { "gpu" } else { "cpu" },
852 );
853 if let Some(c) = cache {
854 probe_cache_store(c, winner);
855 }
856 }
857 }
858}
859
860pub fn probe_deciding(c: OpClass) -> bool {
863 probe_on_for(c) && PROBES[c as usize].state.load(Ordering::Relaxed) == 0
864}
865
866#[allow(unused_variables)]
876pub fn q8_resident_or_upload(model: &Arc<CmfModel>, idx: usize) -> bool {
877 static PROBE_UPLOADS: AtomicU32 = AtomicU32::new(0);
878 let may_upload = PROBE_UPLOADS.load(Ordering::Relaxed) < 4;
879 let resident = match backend() {
880 #[cfg(target_os = "macos")]
881 Backend::Metal => crate::gpu_metal::q8_resident_or_upload(model, idx, may_upload),
882 #[cfg(feature = "gpu")]
883 Backend::Wgpu => crate::gpu_wgpu::q8_resident_or_upload(model, idx, may_upload),
884 Backend::None => false,
885 };
886 if !resident && may_upload {
887 PROBE_UPLOADS.fetch_add(1, Ordering::Relaxed);
888 }
889 resident
890}
891
892#[cfg(test)]
894pub(crate) fn probe_reset() {
895 for p in &PROBES {
896 p.state.store(0, Ordering::Relaxed);
897 p.flip.store(0, Ordering::Relaxed);
898 p.gpu_ns.store(0, Ordering::Relaxed);
899 p.gpu_n.store(0, Ordering::Relaxed);
900 p.cpu_ns.store(0, Ordering::Relaxed);
901 p.cpu_n.store(0, Ordering::Relaxed);
902 }
903}
904
905#[cfg(test)]
909static PROBE_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
910
911#[cfg(test)]
912fn probe_test_guard() -> std::sync::MutexGuard<'static, ()> {
913 PROBE_TEST_LOCK
914 .lock()
915 .unwrap_or_else(std::sync::PoisonError::into_inner)
916}
917
918#[cfg(test)]
919mod probe_tests {
920 use super::*;
921 use std::time::Duration;
922
923 #[test]
924 fn cpu_only_whole_operator_dispatch_inherits_and_restores_scope() {
925 let pool = crate::pool::Pool::with_spin(3, 0);
926 cpu_scope(|| {
927 let inherit = inherit_cpu_scope();
928 pool.run_rows(64, &|_, _| {
929 let _guard = inherit();
930 assert!(CPU_ONLY.get());
931 });
932 });
933 pool.run_rows(64, &|_, _| assert!(!CPU_ONLY.get()));
934 }
935
936 #[test]
939 fn probe_alternates_discards_cold_and_decides() {
940 let _probe_guard = probe_test_guard();
941 probe_reset();
942 assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::Gpu));
944 assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::CpuTimed));
945
946 probe_note_cold();
950 probe_record(OpClass::Ffn, true, Duration::from_secs(1000));
951 for _ in 0..PROBE_SAMPLES {
952 probe_record(OpClass::Ffn, true, Duration::from_millis(1));
953 probe_record(OpClass::Ffn, false, Duration::from_millis(4));
954 }
955 assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::Gpu));
956
957 for _ in 0..PROBE_SAMPLES {
959 probe_record(OpClass::Matmat, true, Duration::from_millis(4));
960 probe_record(OpClass::Matmat, false, Duration::from_millis(1));
961 }
962 assert!(matches!(probe_arm(OpClass::Matmat), ProbeArm::Cpu));
963
964 cpu_scope(|| CPU_ONLY.with(|c| assert!(c.get())));
966 CPU_ONLY.with(|c| assert!(!c.get()));
967 cpu_scope(|| {
968 cpu_scope(|| CPU_ONLY.with(|c| assert!(c.get())));
969 CPU_ONLY.with(|c| assert!(c.get()));
970 });
971 let _ = std::panic::catch_unwind(|| cpu_scope(|| panic!("scope test")));
972 CPU_ONLY.with(|c| assert!(!c.get()));
973 probe_reset();
974 }
975
976 #[test]
977 fn a_remembered_verdict_is_adopted_and_a_stranger_is_not() {
978 let _probe_guard = probe_test_guard();
979 let mine = probe_cache_key_named("gemm-nt");
991 let state = || {
992 PROBES[OpClass::GemmNt as usize]
993 .state
994 .load(Ordering::Relaxed)
995 };
996
997 probe_cache_adopt("SomeOtherGPU/Vulkan\tgemm-nt\tgpu\n");
999 assert_eq!(state(), 0);
1000 let older = mine.replacen(env!("CARGO_PKG_VERSION"), "0.0.0-old", 1);
1002 assert_ne!(older, mine);
1003 probe_cache_adopt(&format!("{older}\tgpu\n"));
1004 assert_eq!(state(), 0);
1005 probe_cache_adopt(&format!("{mine}\tcpu\n"));
1007 assert_eq!(state(), 2);
1008
1009 PROBES[OpClass::GemmNt as usize]
1010 .state
1011 .store(0, Ordering::Relaxed);
1012 }
1013}
1014
1015pub const GPU_MIN_ROWS: usize = 65_536;
1018
1019pub fn min_rows() -> usize {
1026 if let Some(v) = std::env::var("CMF_GPU_MIN_ROWS")
1027 .ok()
1028 .and_then(|v| v.parse().ok())
1029 {
1030 return v;
1031 }
1032 if discrete() { 4096 } else { GPU_MIN_ROWS }
1033}
1034
1035pub fn discrete() -> bool {
1037 match backend() {
1038 #[cfg(feature = "gpu")]
1039 Backend::Wgpu => crate::gpu_wgpu::is_discrete(),
1040 #[cfg(target_os = "macos")]
1041 Backend::Metal => false, Backend::None => false,
1043 }
1044}
1045
1046pub struct MoeJob<'a> {
1050 pub gate: (usize, usize, usize, &'a [f32]),
1051 pub up: (usize, usize, usize, &'a [f32]),
1052 pub down: (usize, usize, usize, &'a [f32]),
1053 pub xs_gate: Vec<f32>,
1054 pub xs_up: Vec<f32>,
1055 pub down_col: &'a [f32],
1056 pub w: f32,
1057 pub q1: bool,
1060 pub q4t: bool,
1063 pub q4tp: bool,
1067 pub gu_q2: bool,
1071 pub swiglu_limit: f32,
1076}
1077
1078pub struct BatchJob<'a> {
1080 pub idx: usize,
1081 pub rows: usize,
1082 pub cols: usize,
1083 pub row_scale: &'a [f32],
1084 pub xs: Vec<f32>,
1085 pub layout: BatchLayout,
1089}
1090
1091#[derive(Clone, Copy, PartialEq, Eq, Debug)]
1094pub enum BatchLayout {
1095 Q8,
1096 Q1,
1097 Q4t,
1098 Q4tp,
1099}
1100
1101#[derive(Clone, Copy, PartialEq, Eq)]
1102enum Backend {
1103 None,
1104 #[cfg(target_os = "macos")]
1105 Metal,
1106 #[cfg(feature = "gpu")]
1107 Wgpu,
1108}
1109
1110fn backend() -> Backend {
1111 #[cfg(feature = "gpu")]
1112 if crate::gpu_wgpu::selected() {
1113 return if crate::gpu_wgpu::enabled() {
1114 Backend::Wgpu
1115 } else {
1116 Backend::None
1117 };
1118 }
1119 #[cfg(target_os = "macos")]
1120 if crate::gpu_metal::enabled() {
1121 return Backend::Metal;
1122 }
1123 Backend::None
1124}
1125
1126pub fn backend_available() -> bool {
1132 #[cfg(target_os = "macos")]
1133 {
1134 true
1136 }
1137 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
1138 {
1139 static AVAIL: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1140 *AVAIL.get_or_init(crate::gpu_wgpu::adapter_probe)
1141 }
1142 #[cfg(all(not(feature = "gpu"), not(target_os = "macos")))]
1143 {
1144 false
1145 }
1146}
1147
1148static GPU_PAUSED: AtomicBool = AtomicBool::new(false);
1154
1155pub fn pause_gpu() -> GpuPause {
1157 GPU_PAUSED.store(true, Ordering::Relaxed);
1158 GpuPause(())
1159}
1160
1161pub struct GpuPause(());
1162
1163impl Drop for GpuPause {
1164 fn drop(&mut self) {
1165 GPU_PAUSED.store(false, Ordering::Relaxed);
1166 }
1167}
1168
1169pub fn enabled() -> bool {
1170 !GPU_PAUSED.load(Ordering::Relaxed) && backend() != Backend::None
1171}
1172
1173pub fn wgpu_active() -> bool {
1187 #[cfg(feature = "gpu")]
1188 {
1189 matches!(backend(), Backend::Wgpu)
1190 }
1191 #[cfg(not(feature = "gpu"))]
1192 {
1193 false
1194 }
1195}
1196
1197pub fn default_device() -> usize {
1204 static D: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
1205 *D.get_or_init(|| {
1206 std::env::var("CMF_GPU_ADAPTER")
1207 .ok()
1208 .and_then(|v| v.trim().parse::<usize>().ok())
1209 .unwrap_or(0)
1210 })
1211}
1212
1213thread_local! {
1214 static CUR_DEV: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
1215}
1216
1217pub fn current_device() -> usize {
1219 CUR_DEV.with(|c| c.get()).unwrap_or_else(default_device)
1220}
1221
1222pub fn set_current_device(i: usize) {
1226 CUR_DEV.with(|c| c.set(Some(i)));
1227}
1228
1229pub fn with_device<R>(dev: usize, f: impl FnOnce() -> R) -> R {
1231 let prev = CUR_DEV.with(|c| c.replace(Some(dev)));
1232 let r = f();
1233 CUR_DEV.with(|c| c.set(prev));
1234 r
1235}
1236
1237pub fn device_count() -> usize {
1240 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
1241 {
1242 return crate::gpu_wgpu::adapter_count();
1243 }
1244 #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
1245 {
1246 usize::from(backend_available())
1247 }
1248}
1249
1250pub fn vram_budget() -> u64 {
1254 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
1255 {
1256 return crate::gpu_wgpu::device_vram_budget();
1257 }
1258 #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
1259 {
1260 if backend_available() { u64::MAX } else { 0 }
1261 }
1262}
1263
1264pub fn resident_bytes() -> u64 {
1268 #[cfg(feature = "gpu")]
1269 {
1270 if backend() == Backend::Wgpu {
1271 return crate::gpu_wgpu::resident_bytes();
1272 }
1273 }
1274 0
1275}
1276
1277pub fn o1_device_stats(kv_id: u64) -> (usize, u64) {
1281 #[cfg(feature = "gpu")]
1282 {
1283 if backend() == Backend::Wgpu {
1284 return crate::gpu_wgpu::o1_device_stats(kv_id);
1285 }
1286 }
1287 let _ = kv_id;
1288 (0, 0)
1289}
1290
1291pub fn upload_bytes() -> u64 {
1295 #[cfg(feature = "gpu")]
1296 {
1297 return crate::gpu_wgpu::UPLOAD_BYTES.load(std::sync::atomic::Ordering::Relaxed);
1298 }
1299 #[cfg(not(feature = "gpu"))]
1300 0
1301}
1302
1303pub fn upload_bandwidth_probe(block: usize, rounds: usize) -> Option<f64> {
1307 #[cfg(feature = "gpu")]
1308 {
1309 return crate::gpu_wgpu::upload_bandwidth_probe(block, rounds);
1310 }
1311 let _ = (block, rounds);
1312 None
1313}
1314
1315#[derive(Clone, Copy, PartialEq, Eq, Debug)]
1327pub enum GraphPhase {
1328 Prefill,
1329 Decode,
1330}
1331
1332pub fn wgpu_graph_on(phase: GraphPhase) -> bool {
1340 match std::env::var("CMF_GPU_WGPU_GRAPH").ok().as_deref() {
1341 Some("0") => false,
1342 Some("prefill") => phase == GraphPhase::Prefill,
1343 Some(_) => true,
1344 None => {
1345 if wgpu_graph_default() {
1346 return true;
1347 }
1348 let _ = phase;
1353 false
1354 }
1355 }
1356}
1357
1358pub fn wgpu_graph_default() -> bool {
1359 #[cfg(feature = "gpu")]
1360 {
1361 matches!(backend(), Backend::Wgpu)
1367 && (crate::gpu_wgpu::discrete_active()
1368 || (cfg!(target_os = "macos") && crate::gpu_wgpu::adapter_up()))
1369 }
1370 #[cfg(not(feature = "gpu"))]
1371 {
1372 false
1373 }
1374}
1375
1376#[allow(clippy::too_many_arguments, unused_variables)]
1378pub fn q8_matvec_range(
1379 model: &Arc<CmfModel>,
1380 idx: usize,
1381 row0: usize,
1382 row_scale: &[f32],
1383 xs: &[f32],
1384 rows: usize,
1385 cols: usize,
1386 out: &mut [f32],
1387) -> bool {
1388 match backend() {
1389 #[cfg(target_os = "macos")]
1390 Backend::Metal => {
1391 crate::gpu_metal::q8_matvec_range(model, idx, row0, row_scale, xs, rows, cols, out)
1392 }
1393 #[cfg(feature = "gpu")]
1394 Backend::Wgpu => {
1395 crate::gpu_wgpu::q8_matvec_range(model, idx, row0, row_scale, xs, rows, cols, out)
1396 }
1397 Backend::None => false,
1398 }
1399}
1400
1401pub(crate) fn q82_short_rows(
1403 model: &Arc<CmfModel>,
1404 idx: usize,
1405 xs: &[f32],
1406 b: usize,
1407 rows: usize,
1408 cols: usize,
1409 out: &mut [f32],
1410) -> bool {
1411 #[cfg(feature = "gpu")]
1412 if enabled_here() && backend() == Backend::Wgpu {
1413 return crate::gpu_wgpu::q82_short_rows(model, idx, xs, b, rows, cols, out);
1414 }
1415 let _ = (model, idx, xs, b, rows, cols, out);
1416 false
1417}
1418
1419#[allow(clippy::too_many_arguments, unused_variables)]
1422#[allow(clippy::too_many_arguments)]
1426pub fn q8_matmat_2f(
1427 model: &Arc<CmfModel>,
1428 idx: usize,
1429 row_scale: &[f32],
1430 col_field: &[f32],
1431 xs: &[f32],
1432 b: usize,
1433 rows: usize,
1434 cols: usize,
1435 out: &mut [f32],
1436) -> bool {
1437 #[allow(unreachable_patterns)]
1438 match backend() {
1439 #[cfg(feature = "gpu")]
1440 Backend::Wgpu => {
1441 crate::gpu_wgpu::q8_matmat_2f(model, idx, row_scale, col_field, xs, b, rows, cols, out)
1442 }
1443 _ => false,
1444 }
1445}
1446
1447pub fn q8_matmat(
1448 model: &Arc<CmfModel>,
1449 idx: usize,
1450 row_scale: &[f32],
1451 pre: &[f32],
1452 b: usize,
1453 rows: usize,
1454 cols: usize,
1455 out: &mut [f32],
1456) -> bool {
1457 match backend() {
1458 #[cfg(target_os = "macos")]
1459 Backend::Metal => {
1460 crate::gpu_metal::q8_matmat(model, idx, row_scale, pre, b, rows, cols, out)
1461 }
1462 #[cfg(feature = "gpu")]
1463 Backend::Wgpu => crate::gpu_wgpu::q8_matmat(model, idx, row_scale, pre, b, rows, cols, out),
1464 Backend::None => false,
1465 }
1466}
1467
1468#[allow(unused_variables)]
1471pub fn q1_matvec(
1472 model: &Arc<CmfModel>,
1473 idx: usize,
1474 xs: &[f32],
1475 rows: usize,
1476 cols: usize,
1477 out: &mut [f32],
1478) -> bool {
1479 match backend() {
1480 #[cfg(target_os = "macos")]
1481 Backend::Metal => crate::gpu_metal::q1_matvec(model, idx, xs, rows, cols, out),
1482 #[cfg(feature = "gpu")]
1483 Backend::Wgpu => crate::gpu_wgpu::q1_matvec(model, idx, xs, rows, cols, out),
1484 Backend::None => false,
1485 }
1486}
1487
1488#[allow(clippy::too_many_arguments)]
1492pub fn attn_dropin(
1493 model: &Arc<CmfModel>,
1494 kv_id: u64,
1495 layer: usize,
1496 normed: &[f32],
1497 wq_idx: usize,
1498 wk_idx: usize,
1499 wv_idx: usize,
1500 wo_idx: usize,
1501 q_norm: Option<&[f32]>,
1502 k_norm: Option<&[f32]>,
1503 late_qk_norm: bool,
1504 invf: &[f32],
1505 nh: usize,
1506 nkv: usize,
1507 hd: usize,
1508 rd: usize,
1509 hidden: usize,
1510 pos: usize,
1511 cap: usize,
1512 gemma: bool,
1513 eps: f32,
1514 cpu_k: &[Vec<f32>],
1515 cpu_v: &[Vec<f32>],
1516 out: &mut [f32],
1517) -> bool {
1518 match backend() {
1519 #[cfg(feature = "gpu")]
1520 Backend::Wgpu => crate::gpu_wgpu::attn_dropin_gpu(
1521 model, kv_id, layer, normed, wq_idx, wk_idx, wv_idx, wo_idx, q_norm, k_norm,
1522 late_qk_norm, invf, nh, nkv, hd, rd, hidden, pos, cap, gemma, eps, cpu_k, cpu_v, out,
1523 ),
1524 #[allow(unused_variables)]
1525 _ => false,
1526 }
1527}
1528
1529#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1533pub enum GraphPrismOp {
1534 None,
1535 Forward,
1536 InverseEmbedding,
1537}
1538
1539pub struct GraphW<'a> {
1543 pub idx: usize,
1544 pub kind: u8,
1545 pub row_scale: &'a [f32],
1546 pub data: &'a [f32],
1547 pub prism: GraphPrismOp,
1548 pub affine: bool,
1549}
1550
1551pub enum GraphAttn<'a> {
1554 Full {
1555 wq: GraphW<'a>,
1556 wk: GraphW<'a>,
1557 wv: GraphW<'a>,
1558 wo: GraphW<'a>,
1559 q_norm: Option<&'a [f32]>,
1560 k_norm: Option<&'a [f32]>,
1561 late_qk_norm: bool,
1563 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
1565 output_gate: bool,
1568 cpu_k: &'a [Vec<f32>],
1569 cpu_v: &'a [Vec<f32>],
1570 cpu_base: usize,
1574 geom: Option<GraphAttnGeom<'a>>,
1581 head_gate: Option<GraphW<'a>>,
1586 },
1587 Gdn {
1588 qkv: GraphW<'a>,
1589 z: GraphW<'a>,
1590 a: GraphW<'a>,
1591 b: GraphW<'a>,
1592 out: GraphW<'a>,
1593 conv1d: &'a [f32],
1594 a_log: &'a [f32],
1595 dt_bias: &'a [f32],
1596 norm: &'a [f32],
1597 nv: usize,
1598 nk: usize,
1599 dk: usize,
1600 dv: usize,
1601 kk: usize,
1602 cpu_state: &'a [f32],
1607 },
1608 ShortConv {
1615 inp: GraphW<'a>,
1617 out: GraphW<'a>,
1619 taps: &'a [f32],
1622 kernel: usize,
1623 cpu_state: &'a [f32],
1627 },
1628}
1629
1630#[derive(Clone, Copy)]
1635pub struct GraphAttnGeom<'a> {
1636 pub nkv: usize,
1638 pub dv: usize,
1640 pub rd: usize,
1642 pub invf: &'a [f32],
1644 pub window: Option<usize>,
1647 pub sink: Option<&'a [f32]>,
1650}
1651
1652#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1656pub enum GraphAct {
1657 Silu,
1658 GeluErf,
1661}
1662
1663impl GraphAct {
1664 pub fn code(self) -> u32 {
1665 match self {
1666 Self::Silu => 0,
1667 Self::GeluErf => 1,
1668 }
1669 }
1670}
1671
1672pub struct GraphLayer<'a> {
1674 pub input_norm: &'a [f32],
1675 pub attn: GraphAttn<'a>,
1676 pub post_norm: &'a [f32],
1677 pub ffn: GraphFfn<'a>,
1678}
1679
1680pub enum GraphFfn<'a> {
1685 AttentionOnly,
1688 Dense {
1689 gate: GraphW<'a>,
1690 up: GraphW<'a>,
1691 down: GraphW<'a>,
1692 act: GraphAct,
1695 },
1696 Moe {
1697 router: GraphW<'a>,
1699 shared_gate: GraphW<'a>,
1701 experts: Vec<(usize, usize, usize)>,
1705 n_exp: usize,
1707 top_k: usize,
1708 inter: usize,
1709 norm_topk: bool,
1710 q4tp: bool,
1716 gu_q2: bool,
1720 sigmoid: bool,
1724 bias: Option<&'a [f32]>,
1727 has_shared: bool,
1731 shared_gated: bool,
1736 route_scale: f32,
1739 },
1740}
1741
1742#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1745pub enum TokenGraphOutcome {
1746 Declined,
1748 Completed,
1750 Failed,
1752}
1753
1754#[allow(clippy::too_many_arguments)]
1759pub fn forward_token_graph(
1760 model: &Arc<CmfModel>,
1761 kv_id: u64,
1762 layers: &[GraphLayer],
1763 o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1766 o1_epoch: u64,
1767 invf: &[f32],
1768 h: &mut [f32],
1769 nh: usize,
1770 nkv: usize,
1771 hd: usize,
1772 attn_scale: f32,
1773 rd: usize,
1774 hidden: usize,
1775 inter: usize,
1776 position: usize,
1777 cap: usize,
1778 gemma: bool,
1779 eps: f32,
1780 lm_head: Option<(&GraphW, usize)>,
1781 final_norm: &[f32],
1782 logits: &mut Vec<f32>,
1783 loop_norm_at: &[usize],
1784 steps: usize,
1785 embed: Option<(&GraphW, usize, f32)>,
1786 ids_out: Option<&mut Vec<u32>>,
1787 layers_run: Option<&mut usize>,
1790 layer_base: usize,
1794 hidden_too: bool,
1796) -> TokenGraphOutcome {
1797 match backend() {
1798 #[cfg(feature = "gpu")]
1799 Backend::Wgpu => crate::gpu_wgpu::forward_token_graph(
1800 model,
1801 kv_id,
1802 layers,
1803 o1,
1804 o1_epoch,
1805 invf,
1806 h,
1807 nh,
1808 nkv,
1809 hd,
1810 attn_scale,
1811 rd,
1812 hidden,
1813 inter,
1814 position,
1815 cap,
1816 gemma,
1817 eps,
1818 lm_head,
1819 final_norm,
1820 logits,
1821 loop_norm_at,
1822 steps,
1823 embed,
1824 ids_out,
1825 layers_run,
1826 layer_base,
1827 hidden_too,
1828 ),
1829 #[allow(unused_variables)]
1830 _ => {
1831 let _ = (
1832 attn_scale,
1833 lm_head,
1834 final_norm,
1835 logits,
1836 loop_norm_at,
1837 layers_run,
1838 layer_base,
1839 hidden_too,
1840 );
1841 TokenGraphOutcome::Declined
1842 }
1843 }
1844}
1845
1846pub fn forward_embryo_graph(
1851 model: &Arc<EmbryoGraphModel>,
1852 kv_id: u64,
1853 hidden: &[f32],
1854 position: usize,
1855 logits: &mut Vec<f32>,
1856) -> bool {
1857 #[cfg(feature = "gpu")]
1858 if backend() == Backend::Wgpu {
1859 return crate::gpu_wgpu::forward_embryo_graph(model, kv_id, hidden, position, logits);
1860 }
1861 let _ = (model, kv_id, hidden, position, logits);
1862 false
1863}
1864
1865#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1869pub enum BatchGraphOutcome {
1870 Declined,
1873 Completed,
1875 Failed,
1879}
1880
1881pub struct SpecTail<'a> {
1882 pub lm: GraphW<'a>,
1883 pub lm_rows: usize,
1884 pub final_norm: &'a [f32],
1885 pub logits_out: &'a mut Vec<f32>,
1886}
1887
1888#[allow(clippy::too_many_arguments)]
1892pub fn forward_batch_graph(
1893 model: &Arc<CmfModel>,
1894 kv_id: u64,
1895 layers: &[GraphLayer],
1896 invf: &[f32],
1897 h: &mut [f32],
1898 nh: usize,
1899 nkv: usize,
1900 hd: usize,
1901 rd: usize,
1902 hidden: usize,
1903 inter: usize,
1904 positions: &[usize],
1905 cap: usize,
1906 gemma: bool,
1907 eps: f32,
1908 attn_scale: f32,
1909 k: usize,
1910 o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1913 o1_epoch: u64,
1914 spec: Option<SpecTail<'_>>,
1915 layers_run: Option<&mut usize>,
1921) -> BatchGraphOutcome {
1922 forward_batch_graph_at(model, kv_id, 0, layers, invf, h, nh, nkv, hd, rd, hidden, inter, positions, cap, gemma, eps, attn_scale, k, o1, o1_epoch, spec, layers_run)
1923}
1924
1925thread_local! {
1926 static MIMO_ATTN_SCRATCH: std::cell::Cell<bool> = const { std::cell::Cell::new(true) };
1927}
1928
1929pub(crate) fn mimo_attention_scratch_enabled() -> bool {
1930 MIMO_ATTN_SCRATCH.with(std::cell::Cell::get)
1931}
1932
1933#[doc(hidden)]
1936pub fn mimo_attention_scratch_scope<R>(enabled: bool, f: impl FnOnce() -> R) -> R {
1937 struct Restore(bool);
1938 impl Drop for Restore {
1939 fn drop(&mut self) {
1940 MIMO_ATTN_SCRATCH.with(|v| v.set(self.0));
1941 }
1942 }
1943 let _restore = Restore(MIMO_ATTN_SCRATCH.with(|v| v.replace(enabled)));
1944 f()
1945}
1946
1947thread_local! {
1948 static MIMO_Q8_SHORT: std::cell::Cell<bool> = const { std::cell::Cell::new(true) };
1949}
1950
1951pub(crate) fn mimo_q8_short_enabled() -> bool {
1952 MIMO_Q8_SHORT.with(std::cell::Cell::get)
1953}
1954
1955#[doc(hidden)]
1957pub fn mimo_q8_short_scope<R>(enabled: bool, f: impl FnOnce() -> R) -> R {
1958 struct Restore(bool);
1959 impl Drop for Restore {
1960 fn drop(&mut self) {
1961 MIMO_Q8_SHORT.with(|v| v.set(self.0));
1962 }
1963 }
1964 let _restore = Restore(MIMO_Q8_SHORT.with(|v| v.replace(enabled)));
1965 f()
1966}
1967
1968#[allow(clippy::too_many_arguments)]
1970pub fn forward_batch_graph_at(
1971 model: &Arc<CmfModel>,
1972 kv_id: u64,
1973 layer_base: usize,
1974 layers: &[GraphLayer],
1975 invf: &[f32],
1976 h: &mut [f32],
1977 nh: usize,
1978 nkv: usize,
1979 hd: usize,
1980 rd: usize,
1981 hidden: usize,
1982 inter: usize,
1983 positions: &[usize],
1984 cap: usize,
1985 gemma: bool,
1986 eps: f32,
1987 attn_scale: f32,
1988 k: usize,
1989 o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1992 o1_epoch: u64,
1993 spec: Option<SpecTail<'_>>,
1994 layers_run: Option<&mut usize>,
2000) -> BatchGraphOutcome {
2001 match backend() {
2002 #[cfg(feature = "gpu")]
2003 Backend::Wgpu => crate::gpu_wgpu::forward_batch_graph_at(
2004 model, kv_id, layer_base, layers, invf, h, nh, nkv, hd, rd, hidden, inter, positions,
2005 cap, gemma, eps, attn_scale, k, o1, o1_epoch, spec, layers_run,
2006 ),
2007 #[allow(unreachable_patterns)]
2008 _ => {
2009 let _ = (o1, o1_epoch, spec, layers_run);
2010 BatchGraphOutcome::Declined
2011 }
2012 }
2013}
2014
2015pub fn gdn_spec_restore(kv_id: u64, slot: usize, base_pos: usize, expected_layers: usize) -> bool {
2020 #[cfg(feature = "gpu")]
2021 if backend() == Backend::Wgpu {
2022 return crate::gpu_wgpu::gdn_spec_restore(kv_id, slot, base_pos, expected_layers);
2023 }
2024 #[allow(unreachable_code)]
2025 {
2026 let _ = (kv_id, slot, base_pos, expected_layers);
2027 false
2028 }
2029}
2030
2031pub fn graph_kv_set_stored(kv_id: u64, layer: usize, stored: usize) -> bool {
2037 #[cfg(feature = "gpu")]
2038 if backend() == Backend::Wgpu {
2039 return crate::gpu_wgpu::kv_mirror_set_stored(kv_id, layer, stored);
2040 }
2041 #[cfg(target_os = "macos")]
2042 if backend() == Backend::Metal {
2043 crate::gpu_metal::kv_mirror_set_stored(kv_id, layer, stored);
2044 return true;
2045 }
2046 false
2047}
2048
2049pub fn graph_kv_stored(_kv_id: u64, _layer: usize) -> Option<usize> {
2053 #[cfg(feature = "gpu")]
2054 if backend() == Backend::Wgpu {
2055 return crate::gpu_wgpu::kv_mirror_stored(_kv_id, _layer);
2056 }
2057 None
2058}
2059
2060pub fn graph_state_resident(_kv_id: u64, _layer: usize) -> bool {
2063 #[cfg(feature = "gpu")]
2064 if backend() == Backend::Wgpu {
2065 return crate::gpu_wgpu::graph_state_resident(_kv_id, _layer);
2066 }
2067 false
2068}
2069
2070pub fn graph_kv_read_rows(
2074 _kv_id: u64,
2075 _reqs: &[(usize, usize, usize)],
2076 _nkv: usize,
2077 _hd: usize,
2078) -> Option<Vec<(Vec<f32>, Vec<f32>)>> {
2079 #[cfg(feature = "gpu")]
2080 if backend() == Backend::Wgpu {
2081 return crate::gpu_wgpu::kv_mirror_read_rows(_kv_id, _reqs, _nkv, _hd);
2082 }
2083 None
2084}
2085
2086pub fn graph_kv_pull_host(
2092 _kv_id: u64,
2093 _layer: usize,
2094 _from: usize,
2095 _to: usize,
2096 _nkv: usize,
2097 _hd: usize,
2098) -> Option<(Vec<f32>, Vec<f32>, usize)> {
2099 #[cfg(feature = "gpu")]
2100 if backend() == Backend::Wgpu {
2101 return crate::gpu_wgpu::kv_mirror_pull_host(_kv_id, _layer, _from, _to, _nkv, _hd);
2102 }
2103 None
2104}
2105
2106pub fn graph_kv_reset(_kv_id: u64) {
2108 #[cfg(feature = "gpu")]
2109 if backend() == Backend::Wgpu {
2110 crate::gpu_wgpu::kv_mirror_reset(_kv_id);
2111 crate::gpu_wgpu::embryo_graph_reset(_kv_id);
2112 }
2113}
2114
2115pub fn q1t_matvec(
2119 model: &Arc<CmfModel>,
2120 idx: usize,
2121 xs: &[f32],
2122 rows: usize,
2123 cols: usize,
2124 out: &mut [f32],
2125) -> bool {
2126 match backend() {
2127 #[cfg(target_os = "macos")]
2128 Backend::Metal => {
2129 if metal_q1t_enabled() {
2130 crate::gpu_metal::q1t_matvec(model, idx, xs, rows, cols, out)
2131 } else {
2132 false
2133 }
2134 }
2135 #[cfg(feature = "gpu")]
2136 Backend::Wgpu => crate::gpu_wgpu::q1t_matvec(model, idx, xs, rows, cols, out),
2137 Backend::None => false,
2138 }
2139}
2140
2141#[allow(unused_variables)]
2144pub fn q4b_matvec(
2145 model: &Arc<CmfModel>,
2146 idx: usize,
2147 xs: &[f32],
2148 rows: usize,
2149 cols: usize,
2150 out: &mut [f32],
2151) -> bool {
2152 match backend() {
2153 #[cfg(target_os = "macos")]
2154 Backend::Metal => false,
2155 #[cfg(feature = "gpu")]
2156 Backend::Wgpu => crate::gpu_wgpu::q4b_matvec(model, idx, xs, rows, cols, out),
2157 Backend::None => false,
2158 }
2159}
2160
2161pub fn q1t_matmat(
2164 model: &Arc<CmfModel>,
2165 idx: usize,
2166 xs: &[f32],
2167 b: usize,
2168 rows: usize,
2169 cols: usize,
2170 out: &mut [f32],
2171) -> bool {
2172 match backend() {
2173 #[cfg(target_os = "macos")]
2174 Backend::Metal => crate::gpu_metal::q1t_matmat(model, idx, xs, b, rows, cols, out),
2178 #[cfg(feature = "gpu")]
2179 Backend::Wgpu => crate::gpu_wgpu::q1t_matmat(model, idx, xs, b, rows, cols, out),
2180 Backend::None => false,
2181 }
2182}
2183
2184#[cfg(target_os = "macos")]
2188pub(crate) fn metal_q1t_enabled() -> bool {
2189 std::env::var("CMF_METAL_Q1T")
2190 .map(|v| v != "0" && !v.eq_ignore_ascii_case("off"))
2191 .unwrap_or(true)
2192}
2193
2194pub fn q1_matmat(
2196 model: &Arc<CmfModel>,
2197 idx: usize,
2198 xs: &[f32],
2199 b: usize,
2200 rows: usize,
2201 cols: usize,
2202 out: &mut [f32],
2203) -> bool {
2204 match backend() {
2205 #[cfg(feature = "gpu")]
2206 Backend::Wgpu => crate::gpu_wgpu::q1_matmat(model, idx, xs, b, rows, cols, out),
2207 #[allow(unused_variables)]
2208 _ => false,
2209 }
2210}
2211
2212static MM_KILL: AtomicBool = AtomicBool::new(false);
2217pub(crate) fn mm_killed() -> bool {
2218 MM_KILL.load(Ordering::Relaxed)
2219}
2220pub(crate) fn mm_kill() {
2221 MM_KILL.store(true, Ordering::Relaxed);
2222}
2223
2224static MM_STRIKES: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
2231const MM_STRIKES_TO_KILL: u32 = 3;
2232static MM_ARMED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
2239
2240pub fn mm_kill_arm(on: bool) {
2243 MM_ARMED.store(on, Ordering::Relaxed);
2244 if on {
2245 MM_STRIKES.store(0, Ordering::Relaxed);
2246 }
2247}
2248
2249pub(crate) fn mm_budget_check(
2256 what: &str,
2257 el: std::time::Duration,
2258 budget: std::time::Duration,
2259 exempt: bool,
2260) {
2261 if el <= budget {
2262 MM_STRIKES.store(0, Ordering::Relaxed);
2263 return;
2264 }
2265 if exempt || !MM_ARMED.load(Ordering::Relaxed) {
2266 return;
2267 }
2268 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2269 let on = *ON.get_or_init(|| std::env::var("CMF_MM_KILL").as_deref() != Ok("0"));
2270 let n = MM_STRIKES.fetch_add(1, Ordering::Relaxed) + 1;
2271 if !on {
2272 tracing::info!(
2273 "gpu {what} took {el:?} (budget {budget:?}) — over budget, CMF_MM_KILL=0 keeps the device"
2274 );
2275 return;
2276 }
2277 if n >= MM_STRIKES_TO_KILL {
2278 tracing::warn!(
2279 "gpu {what} took {el:?} (budget {budget:?}), {n} in a row — \
2280 device contended, CPU for the rest of the process (CMF_MM_KILL=0 to override)"
2281 );
2282 mm_kill();
2283 } else {
2284 tracing::info!(
2285 "gpu {what} took {el:?} (budget {budget:?}) — strike {n} of {MM_STRIKES_TO_KILL}"
2286 );
2287 }
2288}
2289
2290#[allow(unused_variables, clippy::too_many_arguments)]
2295pub fn chunk_attend(
2296 q: &[f32],
2297 k: &[&[f32]],
2298 v: &[&[f32]],
2299 b: usize,
2300 s0: usize,
2301 nh: usize,
2302 nkv: usize,
2303 hd: usize,
2304 scale: f32,
2305 out: &mut [f32],
2306) -> bool {
2307 match backend() {
2308 #[cfg(feature = "gpu")]
2309 Backend::Wgpu => crate::gpu_wgpu::chunk_attend(q, k, v, b, s0, nh, nkv, hd, scale, out),
2310 #[allow(unreachable_patterns)]
2311 _ => false,
2312 }
2313}
2314
2315#[allow(unused_variables, clippy::too_many_arguments)]
2318pub fn chunk_attend_win(
2319 q: &[f32],
2320 k: &[&[f32]],
2321 v: &[&[f32]],
2322 b: usize,
2323 s0: usize,
2324 nh: usize,
2325 nkv: usize,
2326 hd: usize,
2327 scale: f32,
2328 window: usize,
2329 out: &mut [f32],
2330) -> bool {
2331 match backend() {
2332 #[cfg(feature = "gpu")]
2333 Backend::Wgpu => {
2334 crate::gpu_wgpu::chunk_attend_win(q, k, v, b, s0, nh, nkv, hd, scale, window, out)
2335 }
2336 #[allow(unreachable_patterns)]
2337 _ => false,
2338 }
2339}
2340
2341#[allow(unused_variables)]
2344pub fn gemm_many_keep(
2345 model: &Arc<CmfModel>,
2346 idxs: &[(usize, usize)],
2347 xs: &[f32],
2348 b: usize,
2349 cols: usize,
2350 out: &mut [f32],
2351) -> bool {
2352 match backend() {
2353 #[cfg(feature = "gpu")]
2354 Backend::Wgpu => crate::gpu_wgpu::gemm_many_keep(model, idxs, xs, b, cols, out),
2355 #[allow(unreachable_patterns)]
2356 _ => false,
2357 }
2358}
2359
2360thread_local! {
2361 static PREFILL_FAST_GEMM: Cell<bool> = const { Cell::new(false) };
2365}
2366
2367pub struct PrefillFastGemmGuard(bool);
2369
2370impl Drop for PrefillFastGemmGuard {
2371 fn drop(&mut self) {
2372 PREFILL_FAST_GEMM.with(|c| c.set(self.0));
2373 }
2374}
2375
2376pub fn enter_prefill_fast_gemm() -> PrefillFastGemmGuard {
2380 PrefillFastGemmGuard(PREFILL_FAST_GEMM.with(|c| c.replace(true)))
2381}
2382
2383pub fn prefill_fast_gemm() -> bool {
2385 PREFILL_FAST_GEMM.with(|c| c.get())
2386}
2387
2388pub const F16_MAX: f32 = 65504.0;
2391
2392pub fn abs_max_or_inf(xs: &[f32]) -> f32 {
2396 let mut acc = [0f32; 16];
2397 let mut bad = [false; 16];
2398 let mut it = xs.chunks_exact(16);
2399 for c in &mut it {
2400 for ((a, f), &v) in acc.iter_mut().zip(bad.iter_mut()).zip(c) {
2401 let v = v.abs();
2402 *f |= !(v < f32::INFINITY);
2404 *a = if v > *a { v } else { *a };
2405 }
2406 }
2407 let mut m = 0f32;
2408 for &v in it.remainder() {
2409 let v = v.abs();
2410 if !(v < f32::INFINITY) {
2411 return f32::INFINITY;
2412 }
2413 m = m.max(v);
2414 }
2415 if bad.iter().any(|&f| f) {
2416 return f32::INFINITY;
2417 }
2418 acc.iter().fold(m, |m, &a| m.max(a))
2419}
2420
2421#[derive(Clone, Copy, Debug, PartialEq, Eq)]
2427pub struct PrefillMirror {
2428 pub kv_id: u64,
2429 pub layer: usize,
2430 pub limit: usize,
2431}
2432
2433thread_local! {
2434 static PREFILL_MIRROR: Cell<Option<PrefillMirror>> = const { Cell::new(None) };
2438}
2439
2440pub struct PrefillMirrorGuard(Option<PrefillMirror>);
2442
2443impl Drop for PrefillMirrorGuard {
2444 fn drop(&mut self) {
2445 PREFILL_MIRROR.with(|c| c.set(self.0));
2446 }
2447}
2448
2449pub fn enter_prefill_mirror(t: PrefillMirror) -> PrefillMirrorGuard {
2452 PrefillMirrorGuard(PREFILL_MIRROR.with(|c| c.replace(Some(t))))
2453}
2454
2455pub fn prefill_mirror() -> Option<PrefillMirror> {
2457 PREFILL_MIRROR.with(|c| c.get())
2458}
2459
2460#[allow(unused_variables, clippy::too_many_arguments)]
2473pub fn chunk_attend_mirror(
2474 t: PrefillMirror,
2475 cpu_k: &[Vec<f32>],
2476 cpu_v: &[Vec<f32>],
2477 base: usize,
2478 q: &[f32],
2479 b: usize,
2480 s0: usize,
2481 nh: usize,
2482 nkv: usize,
2483 hd: usize,
2484 scale: f32,
2485 ring: Option<usize>,
2486 window: usize,
2487 operand_max: f32,
2488 out: &mut [f32],
2489) -> bool {
2490 match backend() {
2491 #[cfg(feature = "gpu")]
2492 Backend::Wgpu => crate::gpu_wgpu::chunk_attend_mirror(
2497 t.kv_id,
2498 t.layer,
2499 t.limit,
2500 cpu_k,
2501 cpu_v,
2502 base,
2503 q,
2504 b,
2505 s0,
2506 nh,
2507 nkv,
2508 hd,
2509 scale,
2510 ring,
2511 window,
2512 std::env::var("CMF_PREFILL_ATTN_COOP").as_deref() != Ok("0"),
2513 operand_max,
2514 out,
2515 ),
2516 #[allow(unreachable_patterns)]
2517 _ => false,
2518 }
2519}
2520
2521#[allow(unused_variables, clippy::too_many_arguments)]
2529pub fn chunk_attend_mirror_wo(
2530 t: PrefillMirror,
2531 cpu_k: &[Vec<f32>],
2532 cpu_v: &[Vec<f32>],
2533 base: usize,
2534 q: &[f32],
2535 b: usize,
2536 s0: usize,
2537 nh: usize,
2538 nkv: usize,
2539 hd: usize,
2540 scale: f32,
2541 ring: Option<usize>,
2542 window: usize,
2543 operand_max: f32,
2544 gains: Option<&[f32]>,
2545 wo: (&Arc<CmfModel>, usize),
2546 hidden: usize,
2547 out: &mut [f32],
2548) -> bool {
2549 match backend() {
2550 #[cfg(feature = "gpu")]
2551 Backend::Wgpu => crate::gpu_wgpu::chunk_attend_mirror_wo(
2552 t.kv_id,
2553 t.layer,
2554 t.limit,
2555 cpu_k,
2556 cpu_v,
2557 base,
2558 q,
2559 b,
2560 s0,
2561 nh,
2562 nkv,
2563 hd,
2564 scale,
2565 ring,
2566 window,
2567 std::env::var("CMF_PREFILL_ATTN_COOP").as_deref() != Ok("0"),
2568 operand_max,
2569 gains,
2570 wo,
2571 hidden,
2572 out,
2573 ),
2574 #[allow(unreachable_patterns)]
2575 _ => false,
2576 }
2577}
2578
2579pub fn chunk_attend_windowed() -> bool {
2581 match backend() {
2582 #[cfg(feature = "gpu")]
2583 Backend::Wgpu => true,
2584 #[allow(unreachable_patterns)]
2585 _ => false,
2586 }
2587}
2588
2589#[allow(unused_variables, clippy::too_many_arguments)]
2593pub fn q4t_qkv(
2594 model: &Arc<CmfModel>,
2595 wq: usize,
2596 wk: usize,
2597 wv: usize,
2598 xs: &[f32],
2599 b: usize,
2600 cols: usize,
2601 rq: usize,
2602 rk: usize,
2603 rv: usize,
2604 out: &mut [f32],
2605) -> bool {
2606 match backend() {
2607 #[cfg(feature = "gpu")]
2608 Backend::Wgpu => crate::gpu_wgpu::q4t_qkv(model, wq, wk, wv, xs, b, cols, rq, rk, rv, out),
2609 #[allow(unreachable_patterns)]
2610 _ => false,
2611 }
2612}
2613
2614#[allow(unused_variables, clippy::too_many_arguments)]
2616#[allow(clippy::too_many_arguments, unused_variables)]
2620pub fn q4tp_ffn_packed(
2621 model: &Arc<CmfModel>,
2622 w1: usize,
2623 w2: usize,
2624 xs: &[f32],
2625 b: usize,
2626 hidden: usize,
2627 inter: usize,
2628 bias: Option<&[f32]>,
2629 out: &mut [f32],
2630) -> bool {
2631 match backend() {
2632 #[cfg(feature = "gpu")]
2633 Backend::Wgpu => {
2634 crate::gpu_wgpu::ffn_packed(model, w1, w2, xs, b, hidden, inter, bias, out)
2635 }
2636 #[allow(unreachable_patterns)]
2637 _ => false,
2638 }
2639}
2640
2641pub fn q4tp_ffn(
2642 model: &Arc<CmfModel>,
2643 w1: usize,
2644 w3: usize,
2645 w2: usize,
2646 xs: &[f32],
2647 b: usize,
2648 hidden: usize,
2649 inter: usize,
2650 out: &mut [f32],
2651) -> bool {
2652 match backend() {
2653 #[cfg(target_os = "macos")]
2654 Backend::Metal => crate::gpu_metal::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2655 #[cfg(feature = "gpu")]
2656 Backend::Wgpu => crate::gpu_wgpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2657 #[allow(unreachable_patterns)]
2658 _ => false,
2659 }
2660}
2661
2662#[allow(clippy::too_many_arguments, unused_variables)]
2666pub fn q4_ffn_act(
2667 model: &Arc<CmfModel>,
2668 w1: usize,
2669 w3: usize,
2670 w2: usize,
2671 xs: &[f32],
2672 b: usize,
2673 hidden: usize,
2674 inter: usize,
2675 q4tp: bool,
2676 act: GraphAct,
2677 out: &mut [f32],
2678) -> bool {
2679 match act {
2680 GraphAct::Silu if q4tp => q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2681 GraphAct::Silu => q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2682 GraphAct::GeluErf => match backend() {
2683 #[cfg(feature = "gpu")]
2684 Backend::Wgpu => crate::gpu_wgpu::q4_ffn_act(
2685 model,
2686 w1,
2687 w3,
2688 w2,
2689 xs,
2690 b,
2691 hidden,
2692 inter,
2693 q4tp,
2694 act.code(),
2695 out,
2696 ),
2697 #[allow(unreachable_patterns)]
2698 _ => false,
2699 },
2700 }
2701}
2702
2703#[allow(clippy::too_many_arguments, unused_variables)]
2708pub fn ffn_act_keep(
2709 model: &Arc<CmfModel>,
2710 w1: usize,
2711 w3: usize,
2712 w2: usize,
2713 xs: &[f32],
2714 b: usize,
2715 hidden: usize,
2716 inter: usize,
2717 act: GraphAct,
2718 out: &mut [f32],
2719) -> bool {
2720 match backend() {
2721 #[cfg(feature = "gpu")]
2722 Backend::Wgpu => crate::gpu_wgpu::ffn_act_keep(
2723 model,
2724 w1,
2725 w3,
2726 w2,
2727 xs,
2728 b,
2729 hidden,
2730 inter,
2731 act.code(),
2732 out,
2733 ),
2734 #[allow(unreachable_patterns)]
2735 _ => false,
2736 }
2737}
2738
2739#[allow(clippy::too_many_arguments, unused_variables)]
2744pub fn q4tp_gelu_ffn(
2745 model: &Arc<CmfModel>,
2746 w_in: usize,
2747 w_out: usize,
2748 xs: &[f32],
2749 b: usize,
2750 hidden: usize,
2751 inter: usize,
2752 bias_in: &[f32],
2753 bias_out: &[f32],
2754 out: &mut [f32],
2755) -> bool {
2756 match backend() {
2757 #[cfg(feature = "gpu")]
2758 Backend::Wgpu => crate::gpu_wgpu::q4tp_gelu_ffn(
2759 model, w_in, w_out, xs, b, hidden, inter, bias_in, bias_out, out,
2760 ),
2761 #[allow(unreachable_patterns)]
2762 _ => false,
2763 }
2764}
2765
2766pub fn q4t_ffn(
2767 model: &Arc<CmfModel>,
2768 w1: usize,
2769 w3: usize,
2770 w2: usize,
2771 xs: &[f32],
2772 b: usize,
2773 hidden: usize,
2774 inter: usize,
2775 out: &mut [f32],
2776) -> bool {
2777 match backend() {
2778 #[cfg(target_os = "macos")]
2779 Backend::Metal => crate::gpu_metal::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2780 #[cfg(feature = "gpu")]
2781 Backend::Wgpu => crate::gpu_wgpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2782 #[allow(unreachable_patterns)]
2783 _ => false,
2784 }
2785}
2786
2787pub struct DitBlockArgs<'a> {
2792 pub n: usize,
2793 pub hidden: usize,
2794 pub inter: usize,
2795 pub nh: usize,
2796 pub nkv: usize,
2797 pub hd: usize,
2798 pub eps: f32,
2799 pub rope_cos: &'a [f32],
2800 pub rope_sin: &'a [f32],
2801 pub norm1: &'a [f32],
2802 pub norm2: &'a [f32],
2803 pub ffn_norm1: &'a [f32],
2804 pub ffn_norm2: &'a [f32],
2805 pub norm_q: &'a [f32],
2806 pub norm_k: &'a [f32],
2807 pub s_msa: &'a [f32],
2808 pub gate_msa: &'a [f32],
2809 pub s_mlp: &'a [f32],
2810 pub gate_mlp: &'a [f32],
2811 pub wq: usize,
2812 pub wk: usize,
2813 pub wv: usize,
2814 pub wo: usize,
2815 pub w1: usize,
2816 pub w3: usize,
2817 pub w2: usize,
2818 pub q4tp: bool,
2822 pub resident_in: bool,
2825 pub resident_out: bool,
2829}
2830
2831pub fn dit_chain_supported() -> bool {
2835 #[cfg(feature = "gpu")]
2836 {
2837 return matches!(backend(), Backend::Wgpu) && fused_dit_block_available();
2838 }
2839 #[allow(unreachable_code)]
2840 false
2841}
2842
2843pub fn dit_state_fetch(_x: &mut [f32]) -> bool {
2846 #[cfg(feature = "gpu")]
2847 {
2848 if matches!(backend(), Backend::Wgpu) {
2849 return crate::gpu_wgpu::dit_state_fetch(_x);
2850 }
2851 }
2852 false
2853}
2854
2855#[allow(unused_variables)]
2859#[allow(unused_variables, clippy::too_many_arguments)]
2863pub fn dit_qkv(
2864 model: &Arc<CmfModel>,
2865 wq: usize,
2866 wk: usize,
2867 wv: usize,
2868 xs: &[f32],
2869 b: usize,
2870 hidden: usize,
2871 qrows: usize,
2872 kvrows: usize,
2873 q_out: &mut [f32],
2874 k_out: &mut [f32],
2875 v_out: &mut [f32],
2876) -> bool {
2877 match backend() {
2878 #[cfg(feature = "gpu")]
2879 Backend::Wgpu => crate::gpu_wgpu::q4tp_qkv(
2880 model, wq, wk, wv, xs, b, hidden, qrows, kvrows, q_out, k_out, v_out,
2881 ),
2882 #[allow(unreachable_patterns)]
2883 _ => false,
2884 }
2885}
2886
2887pub struct QwenImageAttentionArgs<'a> {
2895 pub image: &'a [f32],
2896 pub text: &'a [f32],
2897 pub image_tokens: usize,
2898 pub text_tokens: usize,
2899 pub heads: usize,
2900 pub head_dim: usize,
2901 pub image_q: usize,
2902 pub image_k: usize,
2903 pub image_v: usize,
2904 pub text_q: usize,
2905 pub text_k: usize,
2906 pub text_v: usize,
2907 pub image_out: usize,
2908 pub text_out: usize,
2909 pub image_q_norm: &'a [f32],
2910 pub image_k_norm: &'a [f32],
2911 pub text_q_norm: &'a [f32],
2912 pub text_k_norm: &'a [f32],
2913 pub image_cos: &'a [f32],
2914 pub image_sin: &'a [f32],
2915 pub text_cos: &'a [f32],
2916 pub text_sin: &'a [f32],
2917 pub image_q_bias: &'a [f32],
2918 pub image_k_bias: &'a [f32],
2919 pub image_v_bias: &'a [f32],
2920 pub text_q_bias: &'a [f32],
2921 pub text_k_bias: &'a [f32],
2922 pub text_v_bias: &'a [f32],
2923 pub image_out_bias: &'a [f32],
2924 pub text_out_bias: &'a [f32],
2925 pub image_proj: &'a mut [f32],
2926 pub text_proj: &'a mut [f32],
2927}
2928
2929#[allow(clippy::too_many_fields)]
2934pub struct QwenImageChainBlock<'a> {
2935 pub image_mod: &'a [f32],
2936 pub text_mod: &'a [f32],
2937 pub image_q: usize,
2938 pub image_k: usize,
2939 pub image_v: usize,
2940 pub text_q: usize,
2941 pub text_k: usize,
2942 pub text_v: usize,
2943 pub image_out: usize,
2944 pub text_out: usize,
2945 pub image_q_norm: &'a [f32],
2946 pub image_k_norm: &'a [f32],
2947 pub text_q_norm: &'a [f32],
2948 pub text_k_norm: &'a [f32],
2949 pub image_q_bias: &'a [f32],
2950 pub image_k_bias: &'a [f32],
2951 pub image_v_bias: &'a [f32],
2952 pub text_q_bias: &'a [f32],
2953 pub text_k_bias: &'a [f32],
2954 pub text_v_bias: &'a [f32],
2955 pub image_out_bias: &'a [f32],
2956 pub text_out_bias: &'a [f32],
2957 pub image_attn_gate: &'a [f32],
2958 pub text_attn_gate: &'a [f32],
2959 pub image_mlp_in: usize,
2960 pub image_mlp_out: usize,
2961 pub text_mlp_in: usize,
2962 pub text_mlp_out: usize,
2963 pub image_mlp_in_bias: &'a [f32],
2964 pub image_mlp_out_bias: &'a [f32],
2965 pub text_mlp_in_bias: &'a [f32],
2966 pub text_mlp_out_bias: &'a [f32],
2967}
2968
2969#[allow(clippy::too_many_fields)]
2975pub struct QwenImageBlockArgs<'a> {
2976 pub image: &'a mut [f32],
2979 pub text: &'a mut [f32],
2980 pub image_norm: &'a [f32],
2981 pub text_norm: &'a [f32],
2982 pub image_tokens: usize,
2983 pub text_tokens: usize,
2984 pub heads: usize,
2985 pub head_dim: usize,
2986 pub image_cos: &'a [f32],
2987 pub image_sin: &'a [f32],
2988 pub text_cos: &'a [f32],
2989 pub text_sin: &'a [f32],
2990 pub image_q: usize,
2991 pub image_k: usize,
2992 pub image_v: usize,
2993 pub text_q: usize,
2994 pub text_k: usize,
2995 pub text_v: usize,
2996 pub image_out: usize,
2997 pub text_out: usize,
2998 pub image_q_norm: &'a [f32],
2999 pub image_k_norm: &'a [f32],
3000 pub text_q_norm: &'a [f32],
3001 pub text_k_norm: &'a [f32],
3002 pub image_q_bias: &'a [f32],
3003 pub image_k_bias: &'a [f32],
3004 pub image_v_bias: &'a [f32],
3005 pub text_q_bias: &'a [f32],
3006 pub text_k_bias: &'a [f32],
3007 pub text_v_bias: &'a [f32],
3008 pub image_out_bias: &'a [f32],
3009 pub text_out_bias: &'a [f32],
3010 pub image_attn_gate: &'a [f32],
3011 pub text_attn_gate: &'a [f32],
3012 pub image_mlp_in: usize,
3013 pub image_mlp_out: usize,
3014 pub text_mlp_in: usize,
3015 pub text_mlp_out: usize,
3016 pub image_mlp_in_bias: &'a [f32],
3017 pub image_mlp_out_bias: &'a [f32],
3018 pub text_mlp_in_bias: &'a [f32],
3019 pub text_mlp_out_bias: &'a [f32],
3020 pub image_mlp_mod: &'a [f32],
3021 pub text_mlp_mod: &'a [f32],
3022 pub image_mlp_gate: &'a [f32],
3023 pub text_mlp_gate: &'a [f32],
3024}
3025
3026pub struct QwenImageChainArgs<'a> {
3031 pub image: &'a mut [f32],
3032 pub text: &'a mut [f32],
3033 pub image_tokens: usize,
3034 pub text_tokens: usize,
3035 pub heads: usize,
3036 pub head_dim: usize,
3037 pub image_cos: &'a [f32],
3038 pub image_sin: &'a [f32],
3039 pub text_cos: &'a [f32],
3040 pub text_sin: &'a [f32],
3041 pub blocks: &'a [QwenImageChainBlock<'a>],
3042}
3043
3044#[allow(unused_variables)]
3045pub fn qwen_image_attention(
3046 model: &Arc<CmfModel>,
3047 args: &mut QwenImageAttentionArgs<'_>,
3048) -> bool {
3049 match backend() {
3050 #[cfg(feature = "gpu")]
3051 Backend::Wgpu => crate::gpu_wgpu::qwen_image_attention(model, args),
3052 #[allow(unreachable_patterns)]
3053 _ => false,
3054 }
3055}
3056
3057#[allow(unused_variables)]
3058pub fn qwen_image_block(model: &Arc<CmfModel>, args: &mut QwenImageBlockArgs<'_>) -> bool {
3059 match backend() {
3060 #[cfg(feature = "gpu")]
3061 Backend::Wgpu => crate::gpu_wgpu::qwen_image_block(model, args),
3062 #[allow(unreachable_patterns)]
3063 _ => false,
3064 }
3065}
3066
3067#[allow(unused_variables)]
3071pub fn qwen_image_chain(model: &Arc<CmfModel>, args: &mut QwenImageChainArgs<'_>) -> bool {
3072 match backend() {
3073 #[cfg(feature = "gpu")]
3074 Backend::Wgpu => crate::gpu_wgpu::qwen_image_chain(model, args),
3075 #[allow(unreachable_patterns)]
3076 _ => false,
3077 }
3078}
3079
3080#[allow(unused_variables, clippy::too_many_arguments)]
3087pub fn qwen_image_mlp_inplace(
3088 model: &Arc<CmfModel>,
3089 w_in: usize,
3090 w_out: usize,
3091 data: &mut [f32],
3092 batch: usize,
3093 hidden: usize,
3094 inter: usize,
3095 bias_in: &[f32],
3096 bias_out: &[f32],
3097 modulation: &[f32],
3098 gate: &[f32],
3099) -> bool {
3100 match backend() {
3101 #[cfg(feature = "gpu")]
3102 Backend::Wgpu => crate::gpu_wgpu::qwen_image_mlp_inplace(
3103 model,
3104 w_in,
3105 w_out,
3106 data,
3107 batch,
3108 hidden,
3109 inter,
3110 bias_in,
3111 bias_out,
3112 modulation,
3113 gate,
3114 ),
3115 #[allow(unreachable_patterns)]
3116 _ => false,
3117 }
3118}
3119
3120pub fn fused_dit_block_available() -> bool {
3124 #[cfg(target_os = "macos")]
3125 {
3126 matches!(backend(), Backend::Metal) && fused_block_trusted()
3127 }
3128 #[cfg(not(target_os = "macos"))]
3129 {
3130 false
3131 }
3132}
3133
3134pub fn dit_block(model: &Arc<CmfModel>, a: &DitBlockArgs, x: &mut [f32]) -> bool {
3135 dit_block_seg(model, a, &[a.n], x)
3136}
3137
3138pub fn dit_block_seg(
3142 model: &Arc<CmfModel>,
3143 a: &DitBlockArgs,
3144 segs: &[usize],
3145 x: &mut [f32],
3146) -> bool {
3147 match backend() {
3148 #[cfg(target_os = "macos")]
3149 Backend::Metal if segs.len() <= 1 => crate::gpu_metal::dit_block(model, a, x),
3150 #[cfg(feature = "gpu")]
3157 Backend::Wgpu
3158 if match std::env::var("CMF_DIT_FUSED").ok().as_deref() {
3159 Some("0") => false,
3160 Some(_) => true,
3161 None => crate::gpu_wgpu::discrete_active(),
3162 } =>
3163 {
3164 crate::gpu_wgpu::dit_block_seg(model, a, segs, x)
3165 }
3166 #[allow(unreachable_patterns)]
3167 _ => false,
3168 }
3169}
3170
3171pub struct VaeResnetArgs<'a> {
3175 pub groups: usize,
3176 pub ic: usize,
3177 pub oc: usize,
3178 pub h: usize,
3179 pub w: usize,
3180 pub n1w: &'a [f32],
3181 pub n1b: &'a [f32],
3182 pub c1w: &'a [f32],
3183 pub c1b: &'a [f32],
3184 pub c1k: usize,
3185 pub n2w: &'a [f32],
3186 pub n2b: &'a [f32],
3187 pub c2w: &'a [f32],
3188 pub c2b: &'a [f32],
3189 pub c2k: usize,
3190 pub shortcut: Option<(&'a [f32], &'a [f32], usize)>,
3191}
3192
3193#[allow(unused_variables)]
3196pub fn vae_resnet(a: &VaeResnetArgs, x: &[f32], out: &mut [f32]) -> bool {
3197 match backend() {
3198 #[cfg(target_os = "macos")]
3199 Backend::Metal => crate::gpu_metal::vae_resnet(a, x, out),
3200 _ => false,
3201 }
3202}
3203
3204#[allow(unused_variables, clippy::too_many_arguments)]
3207pub fn vae_upsample_conv(
3208 w: &[f32],
3209 bias: &[f32],
3210 x: &[f32],
3211 ic: usize,
3212 oc: usize,
3213 h: usize,
3214 w_img: usize,
3215 k: usize,
3216 out: &mut [f32],
3217) -> bool {
3218 match backend() {
3219 #[cfg(target_os = "macos")]
3220 Backend::Metal => crate::gpu_metal::vae_upsample_conv(w, bias, x, ic, oc, h, w_img, k, out),
3221 #[cfg(feature = "gpu")]
3222 Backend::Wgpu => crate::gpu_wgpu::vae_upsample_conv(w, bias, x, ic, oc, h, w_img, k, out),
3223 #[allow(unreachable_patterns)]
3224 _ => false,
3225 }
3226}
3227
3228#[allow(unused_variables, clippy::too_many_arguments)]
3231pub fn vae_conv2d(
3232 w: &[f32],
3233 bias: &[f32],
3234 x: &[f32],
3235 ic: usize,
3236 oc: usize,
3237 h: usize,
3238 w_img: usize,
3239 k: usize,
3240 out: &mut [f32],
3241) -> bool {
3242 match backend() {
3243 #[cfg(target_os = "macos")]
3244 Backend::Metal => crate::gpu_metal::vae_conv2d(w, bias, x, ic, oc, h, w_img, k, out),
3245 #[cfg(feature = "gpu")]
3246 Backend::Wgpu => crate::gpu_wgpu::vae_conv2d(w, bias, x, ic, oc, h, w_img, k, out),
3247 #[allow(unreachable_patterns)]
3248 _ => false,
3249 }
3250}
3251
3252#[allow(unused_variables, clippy::too_many_arguments)]
3256#[allow(unused_variables)]
3260#[allow(clippy::too_many_arguments)]
3261#[allow(clippy::too_many_arguments, unused_variables)]
3264pub fn dit_qkv_attention(
3265 model: &Arc<CmfModel>,
3266 qkv_idx: usize,
3267 xn: &[f32],
3268 n: usize,
3269 hidden: usize,
3270 nh: usize,
3271 hd: usize,
3272 scale: f32,
3273 nr: (&[f32], &[f32], &[f32], f32),
3274 out: &mut [f32],
3275) -> bool {
3276 match backend() {
3277 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3278 Backend::Wgpu => crate::gpu_wgpu::dit_qkv_attention(
3279 model, qkv_idx, xn, n, hidden, nh, hd, scale, nr, out,
3280 ),
3281 #[allow(unreachable_patterns)]
3282 _ => false,
3283 }
3284}
3285
3286#[allow(clippy::too_many_arguments)]
3289pub fn dit_qkv_attn_out(
3290 model: &Arc<CmfModel>,
3291 qkv_idx: usize,
3292 out_idx: usize,
3293 xn: &[f32],
3294 n: usize,
3295 hidden: usize,
3296 nh: usize,
3297 hd: usize,
3298 scale: f32,
3299 nr: (&[f32], &[f32], &[f32], f32),
3300 proj: &mut [f32],
3301) -> bool {
3302 match backend() {
3303 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3304 Backend::Wgpu => crate::gpu_wgpu::dit_qkv_attn_out(
3305 model, qkv_idx, out_idx, xn, n, hidden, nh, hd, scale, nr, proj,
3306 ),
3307 #[allow(unreachable_patterns)]
3308 _ => false,
3309 }
3310}
3311
3312#[allow(clippy::too_many_arguments)]
3314pub fn vae_qkv_attn_out(
3315 model: &Arc<CmfModel>,
3316 qkv_idx: usize,
3317 out_idx: usize,
3318 xn: &[f32],
3319 n: usize,
3320 dim: usize,
3321 nh: usize,
3322 hd: usize,
3323 scale: f32,
3324 angles: &[f32],
3325 eps: f32,
3326 qkv_bias: &[f32],
3327 proj: &mut [f32],
3328) -> bool {
3329 match backend() {
3330 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3331 Backend::Wgpu => crate::gpu_wgpu::vae_qkv_attn_out(
3332 model, qkv_idx, out_idx, xn, n, dim, nh, hd, scale, angles, eps, qkv_bias, proj,
3333 ),
3334 #[allow(unreachable_patterns)]
3335 _ => false,
3336 }
3337}
3338
3339#[allow(clippy::too_many_arguments)]
3340pub fn vae_attention_packed(
3341 qkv: &[f32],
3342 nh: usize,
3343 n: usize,
3344 hd: usize,
3345 scale: f32,
3346 angles: &[f32],
3347 eps: f32,
3348 out: &mut [f32],
3349) -> bool {
3350 vae_attention_packed_layout(qkv, nh, n, hd, scale, angles, eps, out, 1)
3351}
3352
3353#[allow(clippy::too_many_arguments)]
3354pub fn vae_attention_packed_layout(
3355 qkv: &[f32],
3356 nh: usize,
3357 n: usize,
3358 hd: usize,
3359 scale: f32,
3360 angles: &[f32],
3361 eps: f32,
3362 out: &mut [f32],
3363 layout: u32,
3364) -> bool {
3365 match backend() {
3366 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3367 Backend::Wgpu => crate::gpu_wgpu::vae_attention_packed_layout(
3368 qkv, nh, n, hd, scale, angles, eps, out, layout,
3369 ),
3370 #[allow(unreachable_patterns)]
3371 _ => false,
3372 }
3373}
3374
3375#[allow(clippy::too_many_arguments)]
3376pub fn dit_split_only(
3377 qkv: &[f32],
3378 nh: usize,
3379 n: usize,
3380 hd: usize,
3381 layout: u32,
3382 norm: Option<(&[f32], f32)>,
3383 out_q: &mut [f32],
3384) -> bool {
3385 match backend() {
3386 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3387 Backend::Wgpu => crate::gpu_wgpu::dit_split_only(qkv, nh, n, hd, layout, norm, out_q),
3388 #[allow(unreachable_patterns)]
3389 _ => false,
3390 }
3391}
3392
3393pub fn gemm_nt_f32_transient(
3401 x: &[f32],
3402 w: &[f32],
3403 y: &mut [f32],
3404 n: usize,
3405 k: usize,
3406 m: usize,
3407) -> bool {
3408 match backend() {
3409 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3410 Backend::Wgpu => crate::gpu_wgpu::gemm_nt_f32_transient(x, w, y, n, k, m),
3411 #[allow(unreachable_patterns)]
3412 _ => false,
3413 }
3414}
3415
3416pub fn gemm_nt_f32(x: &[f32], w: &[f32], y: &mut [f32], n: usize, k: usize, m: usize) -> bool {
3417 match backend() {
3418 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3419 Backend::Wgpu => crate::gpu_wgpu::gemm_nt_f32(x, w, y, n, k, m),
3420 #[allow(unreachable_patterns)]
3421 _ => false,
3422 }
3423}
3424
3425#[allow(clippy::too_many_arguments)]
3428pub fn music3_ffn(
3429 model: &std::sync::Arc<CmfModel>,
3430 idx_in: usize,
3431 idx_out: usize,
3432 h: &[f32],
3433 bias_in: &[f32],
3434 n: usize,
3435 hs: usize,
3436 inter: usize,
3437 out: &mut [f32],
3438) -> bool {
3439 match backend() {
3440 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3441 Backend::Wgpu => {
3442 crate::gpu_wgpu::music3_ffn(model, idx_in, idx_out, h, bias_in, n, hs, inter, out)
3443 }
3444 #[allow(unreachable_patterns)]
3445 _ => false,
3446 }
3447}
3448
3449#[allow(clippy::too_many_arguments)]
3453pub fn conv1d_gemm(
3454 x: &[f32],
3455 w: &[f32],
3456 ic: usize,
3457 oc: usize,
3458 n: usize,
3459 k: usize,
3460 pad: usize,
3461 dil: usize,
3462 out_n: usize,
3463 yt: &mut [f32],
3464) -> bool {
3465 match backend() {
3466 #[cfg(target_os = "macos")]
3467 Backend::Metal => crate::gpu_metal::conv1d_gemm(x, w, ic, oc, n, k, pad, dil, out_n, yt),
3468 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3469 Backend::Wgpu => crate::gpu_wgpu::conv1d_gemm(x, w, ic, oc, n, k, pad, dil, out_n, yt),
3470 #[allow(unreachable_patterns)]
3471 _ => false,
3472 }
3473}
3474
3475#[allow(clippy::too_many_arguments)]
3477pub fn vae_conv2d_coop(
3478 w: &[f32],
3479 bias: Option<&[f32]>,
3480 x: &[f32],
3481 ic: usize,
3482 oc: usize,
3483 h: usize,
3484 wi: usize,
3485 k: usize,
3486 out: &mut [f32],
3487) -> bool {
3488 match backend() {
3489 #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3490 Backend::Wgpu => crate::gpu_wgpu::vae_conv2d_coop(w, bias, x, ic, oc, h, wi, k, out),
3491 #[allow(unreachable_patterns)]
3492 _ => false,
3493 }
3494}
3495
3496pub fn dit_attention_packed(
3497 qkv: &[f32],
3498 nh: usize,
3499 n: usize,
3500 hd: usize,
3501 scale: f32,
3502 nr: Option<(&[f32], &[f32], &[f32], f32)>,
3505 out: &mut [f32],
3506) -> bool {
3507 match backend() {
3508 #[cfg(feature = "gpu")]
3515 Backend::Wgpu => crate::gpu_wgpu::dit_attention_packed(qkv, nh, n, hd, scale, nr, out),
3516 #[allow(unreachable_patterns)]
3517 _ => false,
3518 }
3519}
3520
3521pub fn dit_attention_packed_available() -> bool {
3529 #[allow(unreachable_patterns)]
3530 match backend() {
3531 #[cfg(feature = "gpu")]
3532 Backend::Wgpu => crate::gpu_wgpu::dit_attention_packed_ready(),
3533 _ => false,
3534 }
3535}
3536
3537pub fn dit_attention(
3538 qh: &[f32],
3539 kh: &[f32],
3540 vh: &[f32],
3541 nh: usize,
3542 nkv: usize,
3543 n: usize,
3544 hd: usize,
3545 scale: f32,
3546 out: &mut [f32],
3547) -> bool {
3548 match backend() {
3549 #[cfg(target_os = "macos")]
3550 Backend::Metal => crate::gpu_metal::dit_attention(qh, kh, vh, nh, nkv, n, hd, scale, out),
3551 #[cfg(feature = "gpu")]
3552 Backend::Wgpu => crate::gpu_wgpu::dit_attention(qh, kh, vh, nh, nkv, n, hd, scale, out),
3553 #[allow(unreachable_patterns)]
3554 _ => false,
3555 }
3556}
3557
3558#[allow(unused_variables)]
3563pub fn q4tp_matmat(
3564 model: &Arc<CmfModel>,
3565 idx: usize,
3566 xs: &[f32],
3567 b: usize,
3568 rows: usize,
3569 cols: usize,
3570 out: &mut [f32],
3571) -> bool {
3572 match backend() {
3573 #[cfg(target_os = "macos")]
3574 Backend::Metal => crate::gpu_metal::q4tp_matmat(model, idx, xs, b, rows, cols, out),
3575 #[cfg(feature = "gpu")]
3576 Backend::Wgpu => crate::gpu_wgpu::q4tp_matmat(model, idx, xs, b, rows, cols, out),
3577 #[allow(unreachable_patterns)]
3578 _ => false,
3579 }
3580}
3581
3582pub fn q2tp_matmat(
3585 model: &Arc<CmfModel>,
3586 idx: usize,
3587 xs: &[f32],
3588 b: usize,
3589 rows: usize,
3590 cols: usize,
3591 out: &mut [f32],
3592) -> bool {
3593 match backend() {
3594 #[cfg(target_os = "macos")]
3595 Backend::Metal => crate::gpu_metal::q2tp_matmat(model, idx, xs, b, rows, cols, out),
3596 #[cfg(feature = "gpu")]
3597 Backend::Wgpu => crate::gpu_wgpu::q2tp_matmat(model, idx, xs, b, rows, cols, out),
3598 #[allow(unreachable_patterns)]
3599 _ => false,
3600 }
3601}
3602
3603pub fn q2tp_affine_matmat(
3606 model: &Arc<CmfModel>,
3607 idx: usize,
3608 xs: &[f32],
3609 b: usize,
3610 rows: usize,
3611 cols: usize,
3612 out: &mut [f32],
3613) -> bool {
3614 match backend() {
3615 #[cfg(target_os = "macos")]
3616 Backend::Metal => crate::gpu_metal::q2tp_affine_matmat(model, idx, xs, b, rows, cols, out),
3617 #[cfg(feature = "gpu")]
3618 Backend::Wgpu => crate::gpu_wgpu::q2tp_affine_matmat(model, idx, xs, b, rows, cols, out),
3619 #[allow(unreachable_patterns)]
3620 _ => false,
3621 }
3622}
3623
3624pub fn q2tp_matvec(
3626 model: &Arc<CmfModel>,
3627 idx: usize,
3628 xs: &[f32],
3629 rows: usize,
3630 cols: usize,
3631 out: &mut [f32],
3632) -> bool {
3633 match backend() {
3634 #[cfg(target_os = "macos")]
3635 Backend::Metal => crate::gpu_metal::q2tp_matvec(model, idx, xs, rows, cols, out),
3636 #[cfg(feature = "gpu")]
3637 Backend::Wgpu => crate::gpu_wgpu::q2tp_matvec(model, idx, xs, rows, cols, out),
3638 #[allow(unreachable_patterns)]
3639 _ => false,
3640 }
3641}
3642
3643pub fn q2tp_affine_matvec(
3647 model: &Arc<CmfModel>,
3648 idx: usize,
3649 xs: &[f32],
3650 rows: usize,
3651 cols: usize,
3652 out: &mut [f32],
3653) -> bool {
3654 match backend() {
3655 #[cfg(target_os = "macos")]
3656 Backend::Metal => crate::gpu_metal::q2tp_affine_matvec(model, idx, xs, rows, cols, out),
3657 #[cfg(feature = "gpu")]
3658 Backend::Wgpu => crate::gpu_wgpu::q2tp_affine_matvec(model, idx, xs, rows, cols, out),
3659 #[allow(unreachable_patterns)]
3660 _ => false,
3661 }
3662}
3663
3664pub fn q4tp_matvec(
3669 model: &Arc<CmfModel>,
3670 idx: usize,
3671 xs: &[f32],
3672 rows: usize,
3673 cols: usize,
3674 out: &mut [f32],
3675) -> bool {
3676 match backend() {
3677 #[cfg(target_os = "macos")]
3678 Backend::Metal => crate::gpu_metal::q4tp_matvec_for_test(model, idx, xs, rows, cols, out),
3679 #[cfg(feature = "gpu")]
3680 Backend::Wgpu => crate::gpu_wgpu::q4tp_matvec(model, idx, xs, rows, cols, out),
3681 #[allow(unreachable_patterns)]
3682 _ => false,
3683 }
3684}
3685
3686pub fn q4t_matvec(
3692 model: &Arc<CmfModel>,
3693 idx: usize,
3694 xs: &[f32],
3695 rows: usize,
3696 cols: usize,
3697 out: &mut [f32],
3698) -> bool {
3699 match backend() {
3700 #[cfg(target_os = "macos")]
3701 Backend::Metal => crate::gpu_metal::q4t_matvec_for_test(model, idx, xs, rows, cols, out),
3702 #[allow(unreachable_patterns)]
3703 _ => false,
3704 }
3705}
3706
3707pub fn q4t_matmat(
3708 model: &Arc<CmfModel>,
3709 idx: usize,
3710 xs: &[f32],
3711 b: usize,
3712 rows: usize,
3713 cols: usize,
3714 out: &mut [f32],
3715) -> bool {
3716 match backend() {
3717 #[cfg(target_os = "macos")]
3718 Backend::Metal => crate::gpu_metal::q4t_matmat(model, idx, xs, b, rows, cols, out),
3719 #[cfg(feature = "gpu")]
3720 Backend::Wgpu => crate::gpu_wgpu::q4t_matmat(model, idx, xs, b, rows, cols, out),
3721 #[allow(unreachable_patterns)]
3722 _ => false,
3723 }
3724}
3725
3726#[cfg(target_os = "macos")]
3728pub use crate::gpu_metal::{
3729 AttnDeviceParams, AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GpuMoe, GraphDims, MetalFfn,
3730 O1AttnParams, TokenGraph, kv_mirror_drop, kv_mirror_read_last, kv_mirror_take_imp,
3731};
3732
3733#[cfg(target_os = "macos")]
3735pub fn gdn_block(
3736 model: &Arc<CmfModel>,
3737 layers: &[GdnGpuLayer],
3738 states: &mut [&mut [f32]],
3739 cfg: &GdnGpuCfg,
3740 h: &mut [f32],
3741) -> bool {
3742 match backend() {
3743 Backend::Metal => crate::gpu_metal::gdn_block(model, layers, states, cfg, h),
3744 _ => false,
3745 }
3746}
3747
3748#[allow(unused_variables)]
3750pub fn moe_block(model: &Arc<CmfModel>, jobs: &[MoeJob], out: &mut [f32]) -> bool {
3751 match backend() {
3752 #[cfg(target_os = "macos")]
3753 Backend::Metal => crate::gpu_metal::moe_block(model, jobs, out),
3754 #[cfg(feature = "gpu")]
3755 Backend::Wgpu => crate::gpu_wgpu::moe_block(model, jobs, out),
3756 Backend::None => false,
3757 }
3758}
3759
3760#[allow(unused_variables)]
3762pub fn matvec_batch(model: &Arc<CmfModel>, jobs: &[BatchJob], out: &mut [&mut [f32]]) -> bool {
3763 match backend() {
3764 #[cfg(target_os = "macos")]
3765 Backend::Metal => crate::gpu_metal::matvec_batch(model, jobs, out),
3766 #[cfg(feature = "gpu")]
3767 Backend::Wgpu => crate::gpu_wgpu::matvec_batch(model, jobs, out),
3768 Backend::None => false,
3769 }
3770}
3771
3772static GRAPH_RACE_STATE: AtomicU8 = AtomicU8::new(0); static GRAPH_RACE_FLIP: AtomicU32 = AtomicU32::new(0);
3788static GRAPH_RACE_ARM_GRAPH: AtomicU8 = AtomicU8::new(0); static GRAPH_RACE_TOK: AtomicU32 = AtomicU32::new(0); static GRAPH_NS: [AtomicU64; 2] = [AtomicU64::new(0), AtomicU64::new(0)]; static GRAPH_N: [AtomicU32; 2] = [AtomicU32::new(0), AtomicU32::new(0)];
3792
3793const GRAPH_RACE_SAMPLES: u32 = 4;
3795
3796pub fn graph_race_begin_generation() {
3815 #[cfg(feature = "gpu")]
3820 {
3821 static FLUSHED: std::sync::Once = std::sync::Once::new();
3833 static FIRST: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
3834 if FIRST.swap(false, Ordering::Relaxed) {
3835 } else {
3837 FLUSHED.call_once(crate::gpu_wgpu::pipeline_cache_flush);
3838 }
3839 }
3840 GRAPH_RACE_TOK.store(0, Ordering::Relaxed);
3841 if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3842 return;
3843 }
3844 let (gn, cn) = (
3845 GRAPH_N[1].load(Ordering::Relaxed),
3846 GRAPH_N[0].load(Ordering::Relaxed),
3847 );
3848 if gn >= GRAPH_RACE_SAMPLES && cn >= GRAPH_RACE_SAMPLES {
3849 let g_avg = GRAPH_NS[1].load(Ordering::Relaxed) / gn as u64;
3850 let c_avg = GRAPH_NS[0].load(Ordering::Relaxed) / cn as u64;
3851 let verdict = if g_avg < c_avg { 1 } else { 2 };
3852 GRAPH_RACE_STATE.store(verdict, Ordering::Relaxed);
3853 tracing::info!(
3854 "wgpu graph race: graph {:.2} ms/tok vs normal {:.2} ms/tok -> {}",
3855 g_avg as f64 / 1e6,
3856 c_avg as f64 / 1e6,
3857 if verdict == 1 { "graph" } else { "normal path" }
3858 );
3859 return;
3860 }
3861 let flip = GRAPH_RACE_FLIP.fetch_add(1, Ordering::Relaxed);
3862 GRAPH_RACE_ARM_GRAPH.store((flip % 2 == 1) as u8, Ordering::Relaxed);
3863}
3864
3865pub fn graph_race_use_graph(trusted: bool) -> bool {
3869 if trusted {
3870 return true;
3871 }
3872 match GRAPH_RACE_STATE.load(Ordering::Relaxed) {
3873 1 => true,
3874 2 => false,
3875 _ => GRAPH_RACE_ARM_GRAPH.load(Ordering::Relaxed) == 1,
3876 }
3877}
3878
3879pub fn graph_race_first_token_hopeless(dur: std::time::Duration) -> bool {
3884 if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3885 return false;
3886 }
3887 let first = GRAPH_RACE_TOK.load(Ordering::Relaxed) == 0;
3888 let cn = GRAPH_N[0].load(Ordering::Relaxed);
3889 if !first || cn == 0 {
3890 return false;
3891 }
3892 let c_avg = GRAPH_NS[0].load(Ordering::Relaxed) / cn as u64;
3893 let ns = dur.as_nanos() as u64;
3894 if ns > 1_000_000_000 && ns > 4 * c_avg {
3895 GRAPH_RACE_STATE.store(2, Ordering::Relaxed);
3896 tracing::info!(
3897 "wgpu graph race: first graph token {:.0} ms vs normal {:.2} ms/tok — hopeless, normal path wins",
3898 ns as f64 / 1e6,
3899 c_avg as f64 / 1e6
3900 );
3901 return true;
3902 }
3903 false
3904}
3905
3906pub fn graph_race_record(used_graph: bool, dur: std::time::Duration) {
3910 if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3911 return;
3912 }
3913 let tok = GRAPH_RACE_TOK.fetch_add(1, Ordering::Relaxed);
3914 if tok == 0 {
3915 return;
3916 }
3917 let i = used_graph as usize;
3918 GRAPH_NS[i].fetch_add(dur.as_nanos() as u64, Ordering::Relaxed);
3919 GRAPH_N[i].fetch_add(1, Ordering::Relaxed);
3920}
3921
3922pub(crate) fn fp_bytes(data: &[u8]) -> u64 {
3932 #[inline]
3933 fn fnv(mut h: u64, bytes: &[u8]) -> u64 {
3934 let (chunks, tail) = bytes.split_at(bytes.len() & !7);
3935 for c in chunks.chunks_exact(8) {
3936 h ^= u64::from_le_bytes(c.try_into().unwrap());
3937 h = h.wrapping_mul(0x100_0000_01b3);
3938 }
3939 for &b in tail {
3940 h ^= b as u64;
3941 h = h.wrapping_mul(0x100_0000_01b3);
3942 }
3943 h
3944 }
3945 let mut h = 0xcbf2_9ce4_8422_2325u64 ^ (data.len() as u64);
3946 if data.len() <= 4096 {
3947 return fnv(h, data);
3948 }
3949 let step = (data.len() - 64) / 63;
3950 for i in 0..64 {
3951 h = fnv(h, &data[i * step..i * step + 64]);
3952 }
3953 h
3954}
3955
3956pub(crate) fn fp_f32(data: &[f32]) -> u64 {
3959 let bytes = unsafe { std::slice::from_raw_parts(data.as_ptr() as *const u8, data.len() * 4) };
3960 fp_bytes(bytes)
3961}
3962
3963#[cfg(test)]
3964mod fp_tests {
3965 use super::fp_bytes;
3966
3967 #[test]
3972 fn fp_bytes_sees_a_change_anywhere_in_a_sampled_slice() {
3973 let n = 1 << 20; let base: Vec<u8> = (0..n).map(|i| (i * 31 + 7) as u8).collect();
3975 let h0 = fp_bytes(&base);
3976 assert_eq!(h0, fp_bytes(&base), "fingerprint must be deterministic");
3977 let mut dense = base.clone();
3980 for b in dense.iter_mut() {
3981 *b = b.wrapping_add(1);
3982 }
3983 assert_ne!(
3984 h0,
3985 fp_bytes(&dense),
3986 "a fully different tensor slipped through"
3987 );
3988 assert_ne!(h0, fp_bytes(&base[..n - 64]));
3991 let mut small = vec![3u8; 4096];
3994 let hs = fp_bytes(&small);
3995 small[2048] ^= 1;
3996 assert_ne!(hs, fp_bytes(&small), "full hash missed a one-byte change");
3997 for n in [4097usize, 5000, 64 * 64, 1 << 16] {
3999 let v = vec![9u8; n];
4000 let _ = fp_bytes(&v); }
4002 }
4003}
4004
4005pub fn bake_release() {
4009 #[cfg(feature = "gpu")]
4010 crate::gpu_wgpu::bake_release();
4011}
4012
4013pub fn bake_precision_strict(on: bool) {
4017 #[cfg(feature = "gpu")]
4018 crate::gpu_wgpu::bake_precision_strict(on);
4019 #[cfg(not(feature = "gpu"))]
4020 let _ = on;
4021}
4022
4023pub fn hostprof_encode_done(t0: std::time::Instant) {
4029 use std::sync::atomic::{AtomicU64, Ordering};
4030 static ENC: AtomicU64 = AtomicU64::new(0);
4031 static N: AtomicU64 = AtomicU64::new(0);
4032 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
4033 return;
4034 }
4035 ENC.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
4036 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
4037 if n % 100 == 0 {
4038 eprintln!(
4039 "hostprof: encode {:.2} ms/token over {n} tokens",
4040 ENC.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
4041 );
4042 }
4043}
4044
4045pub fn hostprof_total(t0: std::time::Instant) {
4046 use std::sync::atomic::{AtomicU64, Ordering};
4047 static TOT: AtomicU64 = AtomicU64::new(0);
4048 static N: AtomicU64 = AtomicU64::new(0);
4049 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
4050 return;
4051 }
4052 TOT.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
4053 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
4054 if n % 100 == 0 {
4055 eprintln!(
4056 "hostprof: total {:.2} ms/token over {n} tokens",
4057 TOT.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
4058 );
4059 }
4060}
4061
4062pub fn stageprof(stage: u32, dt: std::time::Duration) {
4066 use std::sync::atomic::{AtomicU64, Ordering};
4067 static NS: [AtomicU64; 4] = [
4068 AtomicU64::new(0),
4069 AtomicU64::new(0),
4070 AtomicU64::new(0),
4071 AtomicU64::new(0),
4072 ];
4073 static N: AtomicU64 = AtomicU64::new(0);
4074 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
4075 return;
4076 }
4077 NS[stage as usize % 4].fetch_add(dt.as_nanos() as u64, Ordering::Relaxed);
4078 if stage == 1 {
4079 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
4080 if n % 200 == 0 {
4081 eprintln!(
4082 "stageprof: planning {:.2} ms/tok | gdn-item {:.2} ms/tok | attn-item {:.2} ms/tok ({n} tok)",
4083 NS[1].load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
4084 NS[2].load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
4085 NS[3].load(Ordering::Relaxed) as f64 / n as f64 / 1e6
4086 );
4087 }
4088 }
4089}
4090
4091pub fn weight_bytes_dispatched() -> u64 {
4094 let mut total = 0u64;
4095 #[cfg(target_os = "macos")]
4096 {
4097 total += crate::gpu_metal::WEIGHT_BYTES.load(std::sync::atomic::Ordering::Relaxed);
4098 }
4099 #[cfg(feature = "gpu")]
4100 {
4101 total += crate::gpu_wgpu::WEIGHT_BYTES.load(std::sync::atomic::Ordering::Relaxed);
4102 }
4103 total
4104}
4105
4106pub fn weight_bytes_by() -> [u64; 6] {
4109 #[cfg(target_os = "macos")]
4110 {
4111 let mut o = [0u64; 6];
4112 for (i, a) in crate::gpu_metal::WEIGHT_BYTES_BY.iter().enumerate() {
4113 o[i] = a.load(std::sync::atomic::Ordering::Relaxed);
4114 }
4115 return o;
4116 }
4117 #[allow(unreachable_code)]
4118 [0; 6]
4119}
4120
4121#[cfg(test)]
4122mod f16_guard_tests {
4123 use super::*;
4124
4125 #[test]
4128 fn abs_max_or_inf_is_the_magnitude_maximum() {
4129 let mut xs: Vec<f32> = (0..37).map(|i| (i as f32 - 18.0) * 0.5).collect();
4130 assert_eq!(abs_max_or_inf(&xs), 9.0);
4131 xs[36] = -70000.0; assert_eq!(abs_max_or_inf(&xs), 70000.0);
4133 xs[3] = f32::NAN; assert_eq!(abs_max_or_inf(&xs), f32::INFINITY);
4135 xs[3] = 0.0;
4136 xs[35] = f32::NEG_INFINITY;
4137 assert_eq!(abs_max_or_inf(&xs), f32::INFINITY);
4138 assert_eq!(abs_max_or_inf(&[]), 0.0);
4139 assert_eq!(abs_max_or_inf(&[-0.0, -1e-30]), 1e-30);
4140 }
4141}
4142
4143#[cfg(test)]
4144mod probe_warmup_tests {
4145 use super::*;
4146 use std::time::Duration;
4147
4148 fn ms(v: f64) -> Duration {
4149 Duration::from_nanos((v * 1e6) as u64)
4150 }
4151
4152 #[test]
4157 fn one_cold_first_sample_does_not_lose_the_class() {
4158 let p = Probe::new();
4159 probe_record_into(&p, "gemm-nt", None, true, ms(117.01));
4161 probe_record_into(&p, "gemm-nt", None, true, ms(1.1));
4162 probe_record_into(&p, "gemm-nt", None, true, ms(1.0));
4163 probe_record_into(&p, "gemm-nt", None, false, ms(3.19));
4164 probe_record_into(&p, "gemm-nt", None, false, ms(3.20));
4165 assert_eq!(
4166 p.state.load(Ordering::Relaxed),
4167 1,
4168 "the device is 3x faster once warm and must win"
4169 );
4170 }
4171
4172 #[test]
4176 fn the_warmup_is_spent_once_and_never_underflows() {
4177 let p = Probe::new();
4178 for _ in 0..8 {
4179 probe_record_into(&p, "matmat", None, true, ms(10.0));
4180 }
4181 assert_eq!(p.gpu_burn.load(Ordering::Relaxed), 0, "spent, not wrapped");
4182 assert_eq!(
4183 p.gpu_n.load(Ordering::Relaxed),
4184 7,
4185 "one sample burned, the rest counted"
4186 );
4187 }
4188
4189 #[test]
4195 fn a_class_whose_device_always_declines_settles_on_the_host() {
4196 let _probe_guard = probe_test_guard();
4197 let c = OpClass::MatmatWide;
4201 let p = &PROBES[c as usize];
4202 p.state.store(0, Ordering::Relaxed);
4203 p.declines.store(0, Ordering::Relaxed);
4204 for _ in 0..(PROBE_DECLINE_LIMIT - 1) {
4205 probe_note_decline(c);
4206 }
4207 assert_eq!(
4208 p.state.load(Ordering::Relaxed),
4209 0,
4210 "one short of the limit is still a question, not an answer"
4211 );
4212 probe_note_decline(c);
4213 assert_eq!(p.state.load(Ordering::Relaxed), 2, "settled on the host");
4214 assert!(matches!(probe_arm(c), ProbeArm::Cpu));
4215 p.state.store(0, Ordering::Relaxed);
4216 p.declines.store(0, Ordering::Relaxed);
4217 }
4218
4219 #[test]
4222 fn a_slow_device_still_loses_after_the_warmup() {
4223 let p = Probe::new();
4224 for _ in 0..4 {
4225 probe_record_into(&p, "matvec", None, true, ms(40.0));
4226 }
4227 for _ in 0..4 {
4228 probe_record_into(&p, "matvec", None, false, ms(2.0));
4229 }
4230 assert_eq!(p.state.load(Ordering::Relaxed), 2, "host wins on merit");
4231 }
4232}
4233
4234pub(crate) struct ImageStageGuard {
4237 #[cfg(target_os = "macos")]
4238 metal: Option<crate::gpu_metal::ImageStageGuard>,
4239 #[cfg(feature = "gpu")]
4240 wgpu: crate::gpu_wgpu::ImageStageGuard,
4241}
4242
4243pub(crate) fn image_stage_scope() -> ImageStageGuard {
4244 ImageStageGuard {
4245 #[cfg(target_os = "macos")]
4246 metal: if matches!(backend(), Backend::Metal) {
4247 Some(crate::gpu_metal::image_stage_scope())
4248 } else {
4249 None
4250 },
4251 #[cfg(feature = "gpu")]
4252 wgpu: crate::gpu_wgpu::image_stage_scope(),
4253 }
4254}
4255
4256impl ImageStageGuard {
4257 pub(crate) fn track_model(&mut self, uid: u64) {
4258 #[cfg(target_os = "macos")]
4259 if let Some(metal) = &mut self.metal {
4260 metal.track_model(uid);
4261 }
4262 #[cfg(feature = "gpu")]
4263 self.wgpu.track_model(uid);
4264 #[cfg(not(target_os = "macos"))]
4265 let _ = uid;
4266 }
4267}
4268
4269#[derive(Clone, Copy, Debug, PartialEq)]
4289pub struct Qi21Geom {
4290 pub hidden: usize,
4291 pub nh: usize,
4292 pub hd: usize,
4293 pub inter: usize,
4294 pub in_ch: usize,
4295 pub eps: f32,
4296}
4297
4298pub struct Qi21BlockRef<'a> {
4301 pub w: [usize; 7],
4302 pub norm_q: &'a [f32],
4303 pub norm_k: &'a [f32],
4304}
4305
4306pub struct Qi21PrefillArgs<'a> {
4308 pub model: &'a Arc<CmfModel>,
4309 pub geom: Qi21Geom,
4310 pub blocks: &'a [Qi21BlockRef<'a>],
4311 pub img_in: &'a [f32],
4313 pub proj_out: &'a [f32],
4314 pub key: u64,
4315 pub x: &'a [f32],
4317 pub lp: usize,
4318 pub rope_p: (&'a [f32], &'a [f32]),
4320 pub rope_t: (&'a [f32], &'a [f32]),
4321 pub vis: &'a [u32],
4323 pub mods0: &'a [f32],
4325 pub n: usize,
4327}
4328
4329#[allow(unused_variables)]
4331pub fn qi21_prefill(a: &Qi21PrefillArgs) -> bool {
4332 match backend() {
4333 #[cfg(target_os = "macos")]
4334 Backend::Metal => crate::gpu_metal::qi21::prefill(a),
4335 #[cfg(feature = "gpu")]
4336 Backend::Wgpu => crate::gpu_wgpu::qi21::prefill(a),
4337 #[allow(unreachable_patterns)]
4338 _ => false,
4339 }
4340}
4341
4342#[allow(unused_variables)]
4344pub fn qi21_step(key: u64, xtok: &[f32], mods: &[f32], fs: &[f32], out: &mut [f32]) -> bool {
4345 match backend() {
4346 #[cfg(target_os = "macos")]
4347 Backend::Metal => crate::gpu_metal::qi21::step(key, xtok, mods, fs, out),
4348 #[cfg(feature = "gpu")]
4349 Backend::Wgpu => crate::gpu_wgpu::qi21::step(key, xtok, mods, fs, out),
4350 #[allow(unreachable_patterns)]
4351 _ => false,
4352 }
4353}
4354
4355#[allow(unused_variables)]
4358pub fn qi21_release_key(key: u64) {
4359 #[cfg(target_os = "macos")]
4360 if matches!(backend(), Backend::Metal) {
4361 crate::gpu_metal::qi21::release_key(key);
4362 }
4363 #[cfg(feature = "gpu")]
4364 crate::gpu_wgpu::qi21::release_key(key);
4365}
4366
4367pub fn qi21_release() {
4370 #[cfg(target_os = "macos")]
4371 if matches!(backend(), Backend::Metal) {
4372 crate::gpu_metal::qi21::release();
4373 }
4374 #[cfg(feature = "gpu")]
4375 {
4376 crate::gpu_wgpu::qi21::release();
4377 crate::gpu_wgpu::qi21_vae::release();
4378 }
4379}
4380
4381#[derive(Clone, Copy)]
4384pub struct Qi21VaeConvRef<'a> {
4385 pub w: &'a [f32],
4386 pub b: &'a [f32],
4387 pub ci: usize,
4388 pub co: usize,
4389 pub k: usize,
4390}
4391
4392pub struct Qi21VaeResRef<'a> {
4395 pub g1: &'a [f32],
4396 pub c1: Qi21VaeConvRef<'a>,
4397 pub g2: &'a [f32],
4398 pub c2: Qi21VaeConvRef<'a>,
4399 pub shortcut: Option<Qi21VaeConvRef<'a>>,
4400}
4401
4402pub struct Qi21VaeUpRef<'a> {
4406 pub resnets: Vec<Qi21VaeResRef<'a>>,
4407 pub up: Option<(Qi21VaeConvRef<'a>, usize)>,
4408 pub in_dim: usize,
4409 pub out_dim: usize,
4410}
4411
4412pub struct Qi21VaeDecodeArgs<'a> {
4415 pub key: u64,
4417 pub post_quant: Qi21VaeConvRef<'a>,
4418 pub conv_in: Qi21VaeConvRef<'a>,
4419 pub mid_res: [Qi21VaeResRef<'a>; 2],
4420 pub attn_gamma: &'a [f32],
4423 pub attn_qkv: Qi21VaeConvRef<'a>,
4424 pub attn_proj: Qi21VaeConvRef<'a>,
4425 pub ups: Vec<Qi21VaeUpRef<'a>>,
4426 pub norm_out: &'a [f32],
4427 pub conv_out: Qi21VaeConvRef<'a>,
4428}
4429
4430#[allow(unused_variables)]
4434pub fn qi21_vae_decode(a: &Qi21VaeDecodeArgs, z: &[f32], h: usize, w: usize, out: &mut [f32]) -> bool {
4435 match backend() {
4436 #[cfg(feature = "gpu")]
4437 Backend::Wgpu => crate::gpu_wgpu::qi21_vae::decode(a, z, h, w, out),
4438 #[allow(unreachable_patterns)]
4439 _ => false,
4440 }
4441}
4442
4443#[derive(Clone, Copy)]
4451pub struct ZBlockRef<'a> {
4452 pub wq: usize,
4454 pub wk: usize,
4455 pub wv: usize,
4456 pub wo: usize,
4457 pub w1: usize,
4460 pub w3: usize,
4461 pub w2: usize,
4462 pub norm1: &'a [f32],
4464 pub norm2: &'a [f32],
4465 pub ffn_norm1: &'a [f32],
4467 pub ffn_norm2: &'a [f32],
4468 pub norm_q: &'a [f32],
4470 pub norm_k: &'a [f32],
4471}
4472
4473#[derive(Clone, Copy, Debug, PartialEq)]
4477pub struct ZGeom {
4478 pub hidden: usize,
4479 pub nh: usize,
4480 pub hd: usize,
4481 pub inter: usize,
4482 pub eps: f32,
4483 pub final_eps: f32,
4484 pub patch_dim: usize,
4485}
4486
4487pub struct ZPrepareArgs<'a> {
4491 pub model: &'a Arc<CmfModel>,
4492 pub geom: ZGeom,
4493 pub key: u64,
4496 pub n_img: usize,
4499 pub n_img_p: usize,
4500 pub n_cap_p: usize,
4501 pub grid: (usize, usize),
4504 pub cap: &'a [f32],
4506 pub rope_img: (&'a [f32], &'a [f32]),
4509 pub rope_joint: (&'a [f32], &'a [f32]),
4511 pub x_emb_w: &'a [f32],
4514 pub x_emb_b: &'a [f32],
4515 pub x_pad: &'a [f32],
4516 pub final_w: &'a [f32],
4518 pub final_b: &'a [f32],
4519 pub noise_refiner: &'a [ZBlockRef<'a>],
4521 pub layers: &'a [ZBlockRef<'a>],
4522 pub mods_all: Option<&'a [f32]>,
4528 pub final_scale_all: Option<&'a [f32]>,
4529 pub neg: Option<ZNegArgs<'a>>,
4535}
4536
4537pub struct ZNegArgs<'a> {
4541 pub cap: &'a [f32],
4543 pub n_cap_p: usize,
4544 pub rope_img: (&'a [f32], &'a [f32]),
4546 pub rope_joint: (&'a [f32], &'a [f32]),
4548}
4549
4550pub struct ZStepArgs<'a> {
4552 pub key: u64,
4555 pub step: usize,
4558 pub x_tok: &'a [f32],
4562 pub mods: &'a [f32],
4566 pub final_scale: &'a [f32],
4568 pub out: &'a mut [f32],
4571 pub out_neg: Option<&'a mut [f32]>,
4574}
4575
4576#[allow(unused_variables)]
4580pub fn zimage_prepare(a: &ZPrepareArgs) -> bool {
4581 match backend() {
4582 #[cfg(target_os = "macos")]
4583 Backend::Metal => crate::gpu_metal::zimage::prepare(a),
4584 #[cfg(feature = "gpu")]
4585 Backend::Wgpu => crate::gpu_wgpu::zimage::prepare(a),
4586 #[allow(unreachable_patterns)]
4587 _ => false,
4588 }
4589}
4590
4591#[allow(unused_variables)]
4595pub fn zimage_step(a: &mut ZStepArgs) -> bool {
4596 match backend() {
4597 #[cfg(target_os = "macos")]
4598 Backend::Metal => crate::gpu_metal::zimage::step(a),
4599 #[cfg(feature = "gpu")]
4600 Backend::Wgpu => crate::gpu_wgpu::zimage::step(a),
4601 #[allow(unreachable_patterns)]
4602 _ => false,
4603 }
4604}
4605
4606#[allow(unused_variables)]
4611pub fn zimage_preload(
4612 model: &Arc<CmfModel>,
4613 geom: &ZGeom,
4614 noise_refiner: &[ZBlockRef],
4615 layers: &[ZBlockRef],
4616 context_refiner: &[ZBlockRef],
4617) -> bool {
4618 match backend() {
4619 #[cfg(target_os = "macos")]
4620 Backend::Metal => {
4621 crate::gpu_metal::zimage::preload(model, geom, noise_refiner, layers, context_refiner)
4622 }
4623 #[cfg(feature = "gpu")]
4624 Backend::Wgpu => crate::gpu_wgpu::zimage::preload(model, geom, noise_refiner, layers, context_refiner),
4625 #[allow(unreachable_patterns)]
4626 _ => false,
4627 }
4628}
4629
4630pub fn zimage_flush_pipelines() {
4634 #[cfg(feature = "gpu")]
4635 if matches!(backend(), Backend::Wgpu) {
4636 crate::gpu_wgpu::pipeline_cache_flush();
4637 }
4638}
4639
4640pub fn zimage_warmup() -> bool {
4645 match backend() {
4646 #[cfg(target_os = "macos")]
4647 Backend::Metal => crate::gpu_metal::zimage::warmup(),
4648 #[cfg(feature = "gpu")]
4649 Backend::Wgpu => crate::gpu_wgpu::zimage::warmup(),
4650 #[allow(unreachable_patterns)]
4651 _ => false,
4652 }
4653}
4654
4655#[allow(unused_variables)]
4659pub fn vae_prewarm(a: &crate::vae::VaeChainArgs) -> bool {
4660 match backend() {
4661 #[cfg(target_os = "macos")]
4662 Backend::Metal => crate::gpu_metal::zimage::vae_prewarm(a),
4663 #[cfg(feature = "gpu")]
4664 Backend::Wgpu => crate::gpu_wgpu::zimage::vae_prewarm(a),
4665 #[allow(unreachable_patterns)]
4666 _ => false,
4667 }
4668}
4669
4670pub fn zimage_release_dit() {
4673 #[cfg(target_os = "macos")]
4674 crate::gpu_metal::zimage::release_dit();
4675 #[cfg(feature = "gpu")]
4676 crate::gpu_wgpu::zimage::release_dit();
4677}
4678
4679pub fn zimage_release() {
4684 #[cfg(target_os = "macos")]
4685 crate::gpu_metal::zimage::release();
4686 #[cfg(feature = "gpu")]
4687 crate::gpu_wgpu::zimage::release();
4688}
4689
4690#[allow(unused_variables)]
4696pub fn zimage_refine_caption(
4697 model: &Arc<CmfModel>,
4698 geom: &ZGeom,
4699 blocks: &[ZBlockRef],
4700 rope_cap: (&[f32], &[f32]),
4701 cap: &mut [f32],
4702) -> bool {
4703 match backend() {
4704 #[cfg(target_os = "macos")]
4705 Backend::Metal => {
4706 crate::gpu_metal::zimage::refine_caption(model, geom, blocks, rope_cap, cap)
4707 }
4708 #[cfg(feature = "gpu")]
4709 Backend::Wgpu => {
4710 crate::gpu_wgpu::zimage::refine_caption(model, geom, blocks, rope_cap, cap)
4711 }
4712 #[allow(unreachable_patterns)]
4713 _ => false,
4714 }
4715}
4716
4717#[allow(unused_variables)]
4723pub fn vae_decode_chain(
4724 a: &crate::vae::VaeChainArgs,
4725 z: &[f32],
4726 h: usize,
4727 w: usize,
4728 out: &mut [f32],
4729) -> bool {
4730 match backend() {
4731 #[cfg(target_os = "macos")]
4732 Backend::Metal => crate::gpu_metal::zimage::vae_decode_chain(a, z, h, w, out),
4733 #[cfg(feature = "gpu")]
4734 Backend::Wgpu => crate::gpu_wgpu::zimage::vae_decode_chain(a, z, h, w, out),
4735 #[allow(unreachable_patterns)]
4736 _ => false,
4737 }
4738}