#![cfg(all(feature = "metal", target_os = "macos"))]
#[test]
fn embedding_roundtrip() {
let path = std::path::PathBuf::from(std::env::var("HOME").expect("HOME not set"))
.join(".leap/models/LFM2-VL-450M-Q4_0/LFM2-VL-450M-Q4_0.gguf");
if !path.exists() {
eprintln!("skipping — model not found");
return;
}
let gguf = cera::gguf::GgufFile::open(&path).unwrap();
let model = cera::model::load_model_metal(gguf, Some(&path), 1024).unwrap();
let cfg = model.config().clone();
let mut state_a = cera::kv_cache::InferenceState::from_config(&cfg).unwrap();
let logits_a_t0 = model.forward(&[1], 0, &mut state_a);
let logits_a_t1 = model.forward(&[5242], 1, &mut state_a);
let mut state_b = cera::kv_cache::InferenceState::from_config(&cfg).unwrap();
let emb = model.forward_embedding(&[1], 0, &mut state_b);
let logits_b = model.forward_from_embedding(&emb, 1, &mut state_b);
assert_eq!(logits_a_t0.len(), cfg.vocab_size);
assert_eq!(logits_a_t1.len(), cfg.vocab_size);
assert_eq!(logits_b.len(), cfg.vocab_size);
assert!(
logits_b.iter().all(|x| x.is_finite()),
"logits contain NaN/Inf"
);
eprintln!("roundtrip OK: {} logits, all finite", logits_b.len());
}