#![allow(dead_code)]
#[cfg(feature = "remote")]
pub mod download;
pub fn hidden_states_with_lora(
model: &dyn cera::model::Model,
tokens: &[u32],
lora: Option<std::sync::Arc<cera::lora::LoraAdapterWeights>>,
) -> Vec<f32> {
let mut state =
cera::kv_cache::InferenceState::for_prefill(model.config(), tokens.len()).unwrap();
state.lora = lora;
model.hidden_states(tokens, &mut state)
}
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
pub fn metal_context() -> Option<cera::backend::metal::MetalContext> {
match cera::backend::metal::MetalContext::new() {
Ok(ctx) => Some(ctx),
Err(e) => {
assert!(
std::env::var("CERA_REQUIRE_METAL").as_deref() != Ok("1"),
"CERA_REQUIRE_METAL=1 but no Metal device is available ({e})"
);
eprintln!("skipping: no Metal device ({e})");
None
}
}
}
pub fn dense_model_or_skip() -> Option<std::path::PathBuf> {
use std::path::PathBuf;
if let Ok(p) = std::env::var("CERA_DENSE_MODEL") {
let p = PathBuf::from(p);
if p.exists() {
return Some(p);
}
panic!(
"CERA_DENSE_MODEL={} does not exist (unset it to use the defaults)",
p.display()
);
}
let mut tried: Vec<PathBuf> = Vec::new();
if let Ok(home) = std::env::var("HOME") {
let name = "Llama-3.2-1B-Instruct-Q8_0";
let leap = PathBuf::from(home)
.join(".leap/models")
.join(name)
.join(format!("{name}.gguf"));
if leap.exists() {
return Some(leap);
}
tried.push(leap);
}
let repo = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("../models/Llama-3.2-1B-Q4_0.gguf");
if repo.exists() {
return Some(repo);
}
tried.push(repo);
let tried = tried
.iter()
.map(|p| p.display().to_string())
.collect::<Vec<_>>()
.join(", ");
assert!(
std::env::var("CERA_REQUIRE_DENSE_MODEL").as_deref() != Ok("1"),
"CERA_REQUIRE_DENSE_MODEL=1 but no dense model found (CERA_DENSE_MODEL unset; tried {tried})"
);
eprintln!("skipping: no dense model (CERA_DENSE_MODEL unset; tried {tried})");
None
}
#[cfg(all(any(target_arch = "aarch64", target_arch = "x86_64"), not(has_blas)))]
pub fn dense_gemm_head_fixture(path: &std::path::Path) -> (Box<dyn cera::model::Model>, String) {
let gguf = cera::gguf::GgufFile::open(path).unwrap();
let model =
cera::model::load_model(cera::gguf::GgufFile::open(path).unwrap(), None, 8192).unwrap();
assert!(
model.supports_all_logits(),
"fixture is not a dense model with all-position logits, so it does not \
exercise the batched LM-head projection at all"
);
let hs = model.config().hidden_size;
let (reached, detail) = lm_head_gate(&gguf, hs);
assert!(
reached,
"{detail} — the batched LM-head projection would decline on this \
fixture, so anything measured against it is the per-row fallback"
);
(model, detail)
}
#[cfg(all(any(target_arch = "aarch64", target_arch = "x86_64"), not(has_blas)))]
fn lm_head_gate(gguf: &cera::gguf::GgufFile, hidden_size: usize) -> (bool, String) {
let Some(head) = gguf
.tensors
.get("output.weight")
.or_else(|| gguf.tensors.get("token_embd.weight"))
else {
return (
false,
"model has neither output.weight nor token_embd.weight".into(),
);
};
if !cera::model::transformer::batched_gemm_supports(head.dtype, hidden_size) {
return (
false,
format!(
"LM head `{}` is {:?} (hs={hidden_size}) — no batched GEMM kernel here",
head.name, head.dtype
),
);
}
(true, format!("LM head `{}` ({:?})", head.name, head.dtype))
}
#[derive(Clone, Default)]
pub struct WarnCapture(pub std::sync::Arc<std::sync::Mutex<Vec<String>>>);
impl WarnCapture {
pub fn messages(&self) -> Vec<String> {
match self.0.lock() {
Ok(g) => g.clone(),
Err(p) => p.into_inner().clone(),
}
}
}
impl<S: tracing::Subscriber> tracing_subscriber::Layer<S> for WarnCapture {
fn on_event(
&self,
event: &tracing::Event<'_>,
_ctx: tracing_subscriber::layer::Context<'_, S>,
) {
if *event.metadata().level() > tracing::Level::WARN {
return;
}
struct Visit(Option<String>);
impl tracing::field::Visit for Visit {
fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn std::fmt::Debug) {
if field.name() == "message" {
self.0 = Some(format!("{value:?}"));
}
}
}
let mut v = Visit(None);
event.record(&mut v);
let msg =
v.0.unwrap_or_else(|| format!("<no message> target={}", event.metadata().target()));
match self.0.lock() {
Ok(mut g) => g.push(msg),
Err(p) => p.into_inner().push(msg),
}
}
}
fn push_str(out: &mut Vec<u8>, s: &str) {
out.extend_from_slice(&(s.len() as u64).to_le_bytes());
out.extend_from_slice(s.as_bytes());
}
pub fn write_lora_gguf(tensors: &[(String, Vec<usize>, Vec<f32>)], alpha: f32) -> Vec<u8> {
let mut out = Vec::new();
out.extend_from_slice(b"GGUF");
out.extend_from_slice(&3u32.to_le_bytes());
out.extend_from_slice(&(tensors.len() as u64).to_le_bytes());
out.extend_from_slice(&1u64.to_le_bytes());
push_str(&mut out, "adapter.lora.alpha");
out.extend_from_slice(&6u32.to_le_bytes()); out.extend_from_slice(&alpha.to_le_bytes());
let mut offset = 0u64;
for (name, ne, data) in tensors {
push_str(&mut out, name);
out.extend_from_slice(&(ne.len() as u32).to_le_bytes());
ne.iter()
.for_each(|&d| out.extend_from_slice(&(d as u64).to_le_bytes()));
out.extend_from_slice(&0u32.to_le_bytes()); out.extend_from_slice(&offset.to_le_bytes());
offset += (data.len() * 4) as u64;
}
while !out.len().is_multiple_of(32) {
out.push(0);
}
for (_, _, data) in tensors {
out.extend(data.iter().flat_map(|x| x.to_le_bytes()));
}
out
}