use std::path::PathBuf;
use std::sync::{Arc, Mutex};
use algocline_nn::arch::adapter::{LlamaAdapter, LlamaAdapterConfig};
use algocline_nn::arch::{Gpt2Config, Gpt2Model, LoraConfig, TinyLlamaConfig, TinyLlamaModel};
use algocline_nn::card::{
validate_architecture, NnCandleBranch, NnCardMeta, NnLineage, NnLoraBranch,
};
use algocline_nn::merged::{export_merged, MergeError, MergedProvenance};
use algocline_nn::tokenizer::HfTokenizer;
use algocline_nn::train::{
run_distill, run_full_ft, run_lora_ft, Batch, CrossEntropyLoss, Dataset, DatasetOpts,
DistillLossKind, DistillSpec, FullFtConfig, JsonlDataset, ParquetDataset, ScheduleKind,
TokenizedDataset, TrainError, TrainingLease,
};
use candle_core::{DType, Device};
use candle_nn::VarMap;
use mlua::prelude::*;
use mlua::LuaSerdeExt;
use serde_json::{json, Value as Json};
use crate::card::{FileCardStore, SamplesQuery};
const NN_PKG: &str = "alc_nn";
pub(super) fn register_nn_card(
lua: &Lua,
alc_table: &LuaTable,
card_store: Arc<FileCardStore>,
nn_dir: PathBuf,
) -> LuaResult<()> {
let nn_table: LuaTable = alc_table.get("nn")?;
register_preset_ns(lua, &nn_table, nn_dir.clone())?;
register_data_ns(lua, &nn_table, Arc::clone(&card_store), nn_dir.clone())?;
register_trainer_ns(lua, &nn_table, nn_dir.clone())?;
let card_ns = lua.create_table()?;
let save_store = Arc::clone(&card_store);
let save = lua.create_function(
move |lua, (vars, name, meta): (LuaTable, String, LuaTable)| -> LuaResult<String> {
save_impl(lua, save_store.as_ref(), vars, &name, meta)
},
)?;
card_ns.set("save", save)?;
let load_vars_store = Arc::clone(&card_store);
let load_vars = lua.create_function(move |lua, card_id: String| -> LuaResult<LuaTable> {
load_impl(lua, load_vars_store.as_ref(), &card_id)
})?;
card_ns.set("load_vars", load_vars.clone())?;
card_ns.set("load", load_vars)?;
let load_handle_store = Arc::clone(&card_store);
let load_handle_nn_dir = nn_dir.clone();
let load_handle = lua.create_function(move |_lua, card_id: String| -> LuaResult<NnHandle> {
load_handle_impl(load_handle_store.as_ref(), &card_id, &load_handle_nn_dir)
})?;
card_ns.set("load_handle", load_handle)?;
let load_gpt2_store = Arc::clone(&card_store);
let load_gpt2 = lua.create_function(
move |_lua, (card_id, base_handle): (String, LuaAnyUserData)| -> LuaResult<Gpt2Handle> {
load_gpt2_impl(load_gpt2_store.as_ref(), &card_id, &base_handle)
},
)?;
card_ns.set("load_gpt2", load_gpt2)?;
let load_wrap_store = Arc::clone(&card_store);
let load_wrap = lua.create_function(
move |_lua, (card_id, base_handle): (String, LuaAnyUserData)| -> LuaResult<NnHandle> {
load_wrap_impl(load_wrap_store.as_ref(), &card_id, &base_handle)
},
)?;
card_ns.set("load_wrap", load_wrap)?;
let register_store = Arc::clone(&card_store);
let register = lua.create_function(
move |lua, (card_id, model_name): (String, String)| -> LuaResult<()> {
register_impl(lua, register_store.as_ref(), &card_id, &model_name)
},
)?;
card_ns.set("register", register)?;
let merge_lora_store = Arc::clone(&card_store);
let merge_lora_nn_dir = nn_dir.clone();
let merge_lora = lua.create_function(
move |_lua, (base_handle, opts): (LuaAnyUserData, LuaTable)| -> LuaResult<String> {
merge_lora_impl(
merge_lora_store.as_ref(),
&merge_lora_nn_dir,
&base_handle,
opts,
)
},
)?;
card_ns.set("merge_lora", merge_lora)?;
nn_table.set("card", card_ns)?;
let _ = nn_dir;
Ok(())
}
fn save_impl(
lua: &Lua,
store: &FileCardStore,
vars: LuaTable,
name: &str,
meta: LuaTable,
) -> LuaResult<String> {
let meta_json: Json = lua.from_value(LuaValue::Table(meta))?;
let card_id = generate_card_id(name);
let nn_save: LuaFunction = alc_nn_fn(lua, "save")?;
nn_save.call::<()>((vars, card_id.clone()))?;
let payload = build_create_payload(&card_id, name, &meta_json)?;
let (returned_id, _path) = store
.create(payload)
.map_err(|e| LuaError::external(format!("alc.nn.card.save: {e}")))?;
if returned_id != card_id {
return Err(LuaError::external(format!(
"alc.nn.card.save: card_id mismatch (expected {card_id}, got {returned_id})"
)));
}
Ok(card_id)
}
fn load_impl(lua: &Lua, store: &FileCardStore, card_id: &str) -> LuaResult<LuaTable> {
let card = store
.get(card_id)
.map_err(|e| LuaError::external(format!("alc.nn.card.load: {e}")))?
.ok_or_else(|| {
LuaError::external(format!("alc.nn.card.load: card '{card_id}' not found"))
})?;
let bundle_ref = card
.get("metadata")
.and_then(|m| m.get("nn"))
.and_then(|n| n.get("candle"))
.and_then(|c| c.get("bundle_ref"))
.and_then(|b| b.as_str())
.ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load: card '{card_id}' missing metadata.nn.candle.bundle_ref"
))
})?;
let expected = format!("nn/{card_id}");
if bundle_ref != expected {
return Err(LuaError::external(format!(
"alc.nn.card.load: bundle_ref '{bundle_ref}' does not match card_id \
'{card_id}' (expected '{expected}')"
)));
}
let nn_load: LuaFunction = alc_nn_fn(lua, "load")?;
let vars: LuaTable = nn_load.call(card_id.to_string())?;
Ok(vars)
}
fn load_handle_impl(
store: &FileCardStore,
card_id: &str,
nn_dir: &std::path::Path,
) -> LuaResult<NnHandle> {
let card = store
.get(card_id)
.map_err(|e| LuaError::external(format!("alc.nn.card.load_handle: {e}")))?
.ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_handle: card '{card_id}' not found"
))
})?;
let meta_json = card
.get("metadata")
.and_then(|m| m.get("nn"))
.cloned()
.ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_handle: card '{card_id}' missing metadata.nn"
))
})?;
let meta: NnCardMeta = serde_json::from_value(meta_json).map_err(|e| {
LuaError::external(format!(
"alc.nn.card.load_handle: card '{card_id}' invalid metadata.nn: {e}"
))
})?;
match meta.training_path.as_str() {
"full_ft" | "merged" | "distillation" => {}
"lora" => {
return Err(LuaError::external(format!(
"alc.nn.card.load_handle: card '{card_id}' has training_path=\"lora\"; \
LoRA cards need a base handle — call `alc.nn.card.load_wrap(card_id, base)` \
instead"
)));
}
other => {
return Err(LuaError::external(format!(
"alc.nn.card.load_handle: card '{card_id}' has unknown training_path \
{other:?} (expected one of full_ft / lora / merged / distillation)"
)));
}
}
let bundle_ref = meta
.candle
.as_ref()
.map(|c| c.bundle_ref.as_str())
.ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_handle: card '{card_id}' missing metadata.nn.candle"
))
})?;
let expected = format!("nn/{card_id}");
if bundle_ref != expected {
return Err(LuaError::external(format!(
"alc.nn.card.load_handle: bundle_ref '{bundle_ref}' does not match card_id \
'{card_id}' (expected '{expected}')"
)));
}
let ops = resolve_arch_ops(&meta.architecture).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_handle: card '{card_id}' architecture {:?} \
has no bridge dispatch (expected one of {})",
meta.architecture,
registered_arch_names().join(" / ")
))
})?;
let build = ops.build_from_safetensors.ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_handle: card '{card_id}' architecture {:?} \
does not support self-contained card load (adapter-style archs \
need a different entry point — Layer 4b §8 carry)",
meta.architecture
))
})?;
let path = nn_dir.join(format!("{card_id}.safetensors"));
if !path.exists() {
return Err(LuaError::external(format!(
"alc.nn.card.load_handle: bundle missing at {path:?} for card '{card_id}'"
)));
}
build(&meta, &path)
}
fn load_gpt2_impl(
store: &FileCardStore,
card_id: &str,
base_handle: &LuaAnyUserData,
) -> LuaResult<Gpt2Handle> {
let card = store
.get(card_id)
.map_err(|e| LuaError::external(format!("alc.nn.card.load_gpt2: {e}")))?
.ok_or_else(|| {
LuaError::external(format!("alc.nn.card.load_gpt2: card '{card_id}' not found"))
})?;
let meta = extract_nn_card_meta("alc.nn.card.load_gpt2", card_id, &card)?;
precheck_lora_card_meta("alc.nn.card.load_gpt2", card_id, &meta)?;
let base = base_handle
.borrow::<Gpt2Handle>()
.map_err(|e| LuaError::external(format!("alc.nn.card.load_gpt2: base handle: {e}")))?;
wrap_gpt2_lora_from_meta("alc.nn.card.load_gpt2", card_id, &meta, &base)
}
fn precheck_lora_card_meta(ctx: &str, card_id: &str, meta: &NnCardMeta) -> LuaResult<()> {
let candle = meta.candle.as_ref().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' missing metadata.nn.candle"
))
})?;
let lora_branch = candle.lora.as_ref().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' has no metadata.nn.candle.lora block \
(use alc.nn.card.load / load_handle for weight-only reload of a non-LoRA card)"
))
})?;
let delta_path_str = lora_branch.delta_path.as_ref().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' metadata.nn.candle.lora is missing delta_path \
(pre-ST-d cards do not record it; re-save via alc.nn.trainer.lora + \
alc.nn.card.save to populate)"
))
})?;
let delta_path = std::path::Path::new(delta_path_str);
if !delta_path.exists() {
return Err(LuaError::external(format!(
"{ctx}: delta safetensors missing at {delta_path:?} \
(expected the file produced by run_lora_ft; ckpt_dir may have been cleaned)"
)));
}
Ok(())
}
fn extract_nn_card_meta(ctx: &str, card_id: &str, card: &Json) -> LuaResult<NnCardMeta> {
let meta_json = card
.get("metadata")
.and_then(|m| m.get("nn"))
.cloned()
.ok_or_else(|| {
LuaError::external(format!("{ctx}: card '{card_id}' missing metadata.nn"))
})?;
serde_json::from_value(meta_json).map_err(|e| {
LuaError::external(format!("{ctx}: card '{card_id}' invalid metadata.nn: {e}"))
})
}
fn wrap_gpt2_lora_from_meta(
ctx: &str,
card_id: &str,
meta: &NnCardMeta,
base: &Gpt2Handle,
) -> LuaResult<Gpt2Handle> {
let candle = meta.candle.as_ref().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' missing metadata.nn.candle"
))
})?;
let lora_branch = candle.lora.as_ref().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' has no metadata.nn.candle.lora block \
(use alc.nn.card.load / load_handle for weight-only reload of a non-LoRA card)"
))
})?;
let delta_path_str = lora_branch.delta_path.clone().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' metadata.nn.candle.lora is missing delta_path \
(pre-ST-d cards do not record it; re-save via alc.nn.trainer.lora + \
alc.nn.card.save to populate)"
))
})?;
let delta_path = PathBuf::from(&delta_path_str);
if !delta_path.exists() {
return Err(LuaError::external(format!(
"{ctx}: delta safetensors missing at {delta_path:?} \
(expected the file produced by run_lora_ft; ckpt_dir may have been cleaned)"
)));
}
let card_arch = &meta.architecture;
let base_variant = &base.variant;
let base_cfg_id = if base_variant.starts_with("gpt2-") {
base_variant.clone()
} else {
format!("gpt2-{base_variant}")
};
if card_arch != &base_cfg_id && card_arch != base_variant {
return Err(LuaError::external(format!(
"{ctx}: architecture mismatch — card '{card_id}' was trained on \
'{card_arch}' but base handle is '{base_variant}'. Rebuild the base with \
`alc.nn.preset.gpt2('{card_arch}', ...)` (or the neutral \
`alc.nn.preset('gpt2', '{card_arch}', ...)`) to match."
)));
}
let mut lora_cfg = LoraConfig::with_targets(
lora_branch.rank as usize,
lora_branch.alpha as f32,
lora_branch.target_modules.iter().cloned(),
);
lora_cfg.dropout = lora_branch.dropout;
let model_arc = base.model();
let variant = base.variant.clone();
let layers = base.layers;
let heads = base.heads;
let dim = base.dim;
let ctx_len = base.ctx;
let vocab = base.vocab;
let device = base.device.clone();
let dtype = base.dtype.clone();
let pretrained = base.pretrained;
let mut model = model_arc
.lock()
.map_err(|e| LuaError::external(format!("{ctx}: model lock: {e}")))?;
let mut lora_vm = model
.wrap_lora(&lora_cfg)
.map_err(|e| LuaError::external(format!("{ctx}: wrap_lora: {e}")))?;
drop(model);
lora_vm
.load(&delta_path)
.map_err(|e| LuaError::external(format!("{ctx}: load delta {delta_path:?}: {e}")))?;
Ok(Gpt2Handle {
inner: model_arc,
varmap: Some(Arc::new(lora_vm)),
variant,
layers,
heads,
dim,
ctx: ctx_len,
vocab,
device,
dtype,
pretrained,
has_lora: true,
})
}
fn wrap_tinyllama_lora_from_meta(
ctx: &str,
card_id: &str,
meta: &NnCardMeta,
base: &TinyLlamaHandle,
) -> LuaResult<TinyLlamaHandle> {
let candle = meta.candle.as_ref().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' missing metadata.nn.candle"
))
})?;
let lora_branch = candle.lora.as_ref().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' has no metadata.nn.candle.lora block \
(use alc.nn.card.load / load_handle for weight-only reload of a non-LoRA card)"
))
})?;
let delta_path_str = lora_branch.delta_path.clone().ok_or_else(|| {
LuaError::external(format!(
"{ctx}: card '{card_id}' metadata.nn.candle.lora is missing delta_path"
))
})?;
let delta_path = PathBuf::from(&delta_path_str);
if !delta_path.exists() {
return Err(LuaError::external(format!(
"{ctx}: delta safetensors missing at {delta_path:?}"
)));
}
let card_arch = &meta.architecture;
let base_variant = &base.variant;
let base_cfg_id = if base_variant.starts_with("tinyllama-") {
base_variant.clone()
} else {
format!("tinyllama-{base_variant}")
};
if card_arch != &base_cfg_id && card_arch != base_variant {
return Err(LuaError::external(format!(
"{ctx}: architecture mismatch — card '{card_id}' was trained on \
'{card_arch}' but base handle is '{base_variant}'. Rebuild the base with \
`alc.nn.preset.tinyllama('{card_arch}', ...)` (or the neutral \
`alc.nn.preset('tinyllama', '{card_arch}', ...)`) to match."
)));
}
let mut lora_cfg = LoraConfig::with_targets(
lora_branch.rank as usize,
lora_branch.alpha as f32,
lora_branch.target_modules.iter().cloned(),
);
lora_cfg.dropout = lora_branch.dropout;
let model_arc = base.model();
let variant = base.variant.clone();
let layers = base.layers;
let heads = base.heads;
let kv_heads = base.kv_heads;
let dim = base.dim;
let ctx_len = base.ctx;
let vocab = base.vocab;
let device = base.device.clone();
let dtype = base.dtype.clone();
let pretrained = base.pretrained;
let mut model = model_arc
.lock()
.map_err(|e| LuaError::external(format!("{ctx}: model lock: {e}")))?;
let mut lora_vm = model
.wrap_lora(&lora_cfg)
.map_err(|e| LuaError::external(format!("{ctx}: wrap_lora: {e}")))?;
drop(model);
lora_vm
.load(&delta_path)
.map_err(|e| LuaError::external(format!("{ctx}: load delta {delta_path:?}: {e}")))?;
Ok(TinyLlamaHandle {
inner: model_arc,
varmap: Some(Arc::new(lora_vm)),
variant,
layers,
heads,
kv_heads,
dim,
ctx: ctx_len,
vocab,
device,
dtype,
pretrained,
has_lora: true,
})
}
pub(super) fn wrap_gpt2_lora_bridge(base: &Gpt2Handle, cfg: &LoraConfig) -> LuaResult<Gpt2Handle> {
let model_arc = base.model();
let variant = base.variant.clone();
let layers = base.layers;
let heads = base.heads;
let dim = base.dim;
let ctx_len = base.ctx;
let vocab = base.vocab;
let device = base.device.clone();
let dtype = base.dtype.clone();
let pretrained = base.pretrained;
let mut model = model_arc
.lock()
.map_err(|e| LuaError::external(format!("alc.nn.wrap_lora: model lock: {e}")))?;
let lora_vm = model
.wrap_lora(cfg)
.map_err(|e| LuaError::external(format!("alc.nn.wrap_lora: candle: {e}")))?;
drop(model);
Ok(Gpt2Handle {
inner: model_arc,
varmap: Some(Arc::new(lora_vm)),
variant,
layers,
heads,
dim,
ctx: ctx_len,
vocab,
device,
dtype,
pretrained,
has_lora: true,
})
}
pub(super) fn wrap_tinyllama_lora_bridge(
base: &TinyLlamaHandle,
cfg: &LoraConfig,
) -> LuaResult<TinyLlamaHandle> {
let model_arc = base.model();
let variant = base.variant.clone();
let layers = base.layers;
let heads = base.heads;
let kv_heads = base.kv_heads;
let dim = base.dim;
let ctx_len = base.ctx;
let vocab = base.vocab;
let device = base.device.clone();
let dtype = base.dtype.clone();
let pretrained = base.pretrained;
let mut model = model_arc
.lock()
.map_err(|e| LuaError::external(format!("alc.nn.wrap_lora: model lock: {e}")))?;
let lora_vm = model
.wrap_lora(cfg)
.map_err(|e| LuaError::external(format!("alc.nn.wrap_lora: candle: {e}")))?;
drop(model);
Ok(TinyLlamaHandle {
inner: model_arc,
varmap: Some(Arc::new(lora_vm)),
variant,
layers,
heads,
kv_heads,
dim,
ctx: ctx_len,
vocab,
device,
dtype,
pretrained,
has_lora: true,
})
}
fn register_impl(
lua: &Lua,
store: &FileCardStore,
card_id: &str,
model_name: &str,
) -> LuaResult<()> {
let exists = store
.get(card_id)
.map_err(|e| LuaError::external(format!("alc.nn.card.register: {e}")))?
.is_some();
if !exists {
return Err(LuaError::external(format!(
"alc.nn.card.register: card '{card_id}' not found"
)));
}
let placeholder_id = card_id.to_string();
let forward = lua.create_function(move |_, prompt: String| -> LuaResult<String> {
Ok(format!("[nn card {placeholder_id}]:{prompt}"))
})?;
let nn_register: LuaFunction = alc_nn_fn(lua, "register")?;
nn_register.call::<()>((model_name.to_string(), forward))?;
Ok(())
}
fn build_create_payload(card_id: &str, name: &str, user_meta: &Json) -> LuaResult<Json> {
let training_path = user_meta
.get("training_path")
.and_then(|v| v.as_str())
.ok_or_else(|| LuaError::external("alc.nn.card.save: meta.training_path is required"))?
.to_string();
let architecture = user_meta
.get("architecture")
.and_then(|v| v.as_str())
.ok_or_else(|| LuaError::external("alc.nn.card.save: meta.architecture is required"))?
.to_string();
validate_architecture(&architecture)
.map_err(|e| LuaError::external(format!("alc.nn.card.save: {e}")))?;
let task = user_meta
.get("task")
.and_then(|v| v.as_str())
.map(String::from);
let lineage = match user_meta.get("lineage").cloned() {
Some(v) => serde_json::from_value::<NnLineage>(v).map_err(|e| {
LuaError::external(format!("alc.nn.card.save: invalid meta.lineage: {e}"))
})?,
None => NnLineage::default(),
};
let hyperparams = normalise_object(user_meta.get("hyperparams").cloned());
let metrics = normalise_object(user_meta.get("metrics").cloned());
let candle_in = user_meta.get("candle");
let lora = match candle_in.and_then(|c| c.get("lora")) {
Some(v) if !v.is_null() => Some(
serde_json::from_value::<NnLoraBranch>(v.clone()).map_err(|e| {
LuaError::external(format!("alc.nn.card.save: invalid meta.candle.lora: {e}"))
})?,
),
_ => None,
};
let candle = NnCandleBranch {
bundle_ref: format!("nn/{card_id}"),
device: candle_in
.and_then(|c| c.get("device"))
.and_then(|v| v.as_str())
.map(String::from),
dtype: candle_in
.and_then(|c| c.get("dtype"))
.and_then(|v| v.as_str())
.map(String::from),
lora,
};
let nn_meta = NnCardMeta {
name: name.to_string(),
backend: "candle".into(),
task,
architecture,
training_path,
lineage,
hyperparams,
metrics,
candle: Some(candle),
};
let nn_meta_json = serde_json::to_value(&nn_meta)
.map_err(|e| LuaError::external(format!("alc.nn.card.save: serialize meta: {e}")))?;
Ok(json!({
"pkg": { "name": NN_PKG },
"card_id": card_id,
"metadata": {
"kind": "nn_model",
"nn": nn_meta_json,
}
}))
}
pub(super) fn build_create_payload_from_meta(card_id: &str, meta: &NnCardMeta) -> LuaResult<Json> {
validate_architecture(&meta.architecture)
.map_err(|e| LuaError::external(format!("alc.nn.card.merge_lora: {e}")))?;
let nn_meta_json = serde_json::to_value(meta)
.map_err(|e| LuaError::external(format!("alc.nn.card.merge_lora: serialize meta: {e}")))?;
Ok(json!({
"pkg": { "name": NN_PKG },
"card_id": card_id,
"metadata": {
"kind": "nn_model",
"nn": nn_meta_json,
}
}))
}
fn merge_error_to_lua(err: MergeError) -> LuaError {
let msg = match err {
MergeError::Provenance(inner) => format!("alc.nn.card.merge_lora: provenance: {inner}"),
MergeError::Merge(inner) => format!("alc.nn.card.merge_lora: merge: {inner}"),
MergeError::Io(inner) => format!("alc.nn.card.merge_lora: io: {inner}"),
MergeError::Serialize(inner) => format!("alc.nn.card.merge_lora: serialize: {inner}"),
};
LuaError::external(msg)
}
fn merge_lora_impl(
store: &FileCardStore,
nn_dir: &std::path::Path,
base_handle: &LuaAnyUserData,
opts: LuaTable,
) -> LuaResult<String> {
let name: Option<String> = opts.get("name")?;
let name = name.filter(|s| !s.is_empty()).ok_or_else(|| {
LuaError::external("alc.nn.card.merge_lora: opts.name must be a non-empty string")
})?;
let lora_card: Option<String> = opts.get("lora_card")?;
let lora_card = lora_card.filter(|s| !s.is_empty()).ok_or_else(|| {
LuaError::external("alc.nn.card.merge_lora: opts.lora_card must be a non-empty string")
})?;
let handle: NnHandle = if let Ok(nn) = base_handle.borrow::<NnHandle>() {
(*nn).clone()
} else if let Ok(g) = base_handle.borrow::<Gpt2Handle>() {
NnHandle::Gpt2(g.clone())
} else if let Ok(t) = base_handle.borrow::<TinyLlamaHandle>() {
NnHandle::TinyLlama(t.clone())
} else if let Ok(l) = base_handle.borrow::<LlamaHandle>() {
NnHandle::Llama(l.clone())
} else {
return Err(LuaError::external(
"alc.nn.card.merge_lora: base handle is not a recognised NnHandle / \
Gpt2Handle / TinyLlamaHandle / LlamaHandle",
));
};
if !handle.is_lora_wrapped() {
return Err(LuaError::external(format!(
"alc.nn.card.merge_lora: handle is not LoRA-wrapped (arch={:?}); \
use `alc.nn.card.load_wrap(lora_card_id, base_handle)` to obtain a \
wrapped handle before calling merge_lora",
handle.arch()
)));
}
let merged_card_id = generate_card_id(&name);
let arch = handle.arch_family_variant();
let bundle_ref = format!("nn/{merged_card_id}");
let provenance = MergedProvenance {
lora_card,
arch,
bundle_ref,
};
let out_path = nn_dir.join(format!("{merged_card_id}.safetensors"));
let (_bytes, meta) = match &handle {
NnHandle::Gpt2(gpt2) => {
let model_arc = gpt2.model();
let model_guard = model_arc.lock().map_err(|e| {
LuaError::external(format!("alc.nn.card.merge_lora: model lock: {e}"))
})?;
export_merged(&*model_guard, &provenance, &out_path).map_err(merge_error_to_lua)?
}
NnHandle::TinyLlama(tll) => {
let model_arc = tll.model();
let model_guard = model_arc.lock().map_err(|e| {
LuaError::external(format!("alc.nn.card.merge_lora: model lock: {e}"))
})?;
export_merged(&*model_guard, &provenance, &out_path).map_err(merge_error_to_lua)?
}
NnHandle::Llama(_) => {
return Err(LuaError::external(
"alc.nn.card.merge_lora: llama adapter path does not support LoRA merge",
));
}
};
let mut meta = meta;
meta.name = name.clone();
let payload = build_create_payload_from_meta(&merged_card_id, &meta)?;
let (returned_id, _path) = store
.create(payload)
.map_err(|e| LuaError::external(format!("alc.nn.card.merge_lora: card store: {e}")))?;
if returned_id != merged_card_id {
return Err(LuaError::external(format!(
"alc.nn.card.merge_lora: card_id mismatch (expected {merged_card_id}, got {returned_id})"
)));
}
Ok(merged_card_id)
}
fn normalise_object(v: Option<Json>) -> Json {
match v {
Some(Json::Object(m)) => Json::Object(m),
Some(Json::Array(a)) if a.is_empty() => Json::Object(serde_json::Map::new()),
Some(other) => other,
None => Json::Object(serde_json::Map::new()),
}
}
fn generate_card_id(name: &str) -> String {
let ts = compact_epoch_us();
let sanitized = sanitize_name(name);
format!("{sanitized}_{ts}")
}
pub(super) fn sanitize_name(name: &str) -> String {
let mut out = String::with_capacity(name.len());
for c in name.chars() {
if c.is_ascii_alphanumeric() || c == '-' || c == '_' {
out.push(c);
} else {
out.push('_');
}
}
if out.is_empty() {
"nn".into()
} else {
out
}
}
pub(super) fn compact_epoch_us() -> String {
let d = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default();
format!("{}{:06}", d.as_secs(), d.subsec_micros())
}
fn alc_nn_fn(lua: &Lua, key: &str) -> LuaResult<LuaFunction> {
let alc: LuaTable = lua
.globals()
.get("alc")
.map_err(|e| LuaError::external(format!("alc.nn.card: `alc` global missing: {e}")))?;
let nn: LuaTable = alc
.get("nn")
.map_err(|e| LuaError::external(format!("alc.nn.card: `alc.nn` missing: {e}")))?;
nn.get::<LuaFunction>(key)
.map_err(|e| LuaError::external(format!("alc.nn.card: `alc.nn.{key}` missing: {e}")))
}
#[derive(Clone)]
pub(super) struct Gpt2Handle {
inner: Arc<Mutex<Gpt2Model>>,
varmap: Option<Arc<VarMap>>,
variant: String,
layers: usize,
heads: usize,
dim: usize,
ctx: usize,
vocab: usize,
device: String,
dtype: String,
pretrained: bool,
pub(super) has_lora: bool,
}
impl mlua::UserData for Gpt2Handle {
fn add_methods<M: mlua::UserDataMethods<Self>>(methods: &mut M) {
methods.add_method("variant", |_, this, ()| Ok(this.variant.clone()));
methods.add_method("layers", |_, this, ()| Ok(this.layers));
methods.add_method("heads", |_, this, ()| Ok(this.heads));
methods.add_method("dim", |_, this, ()| Ok(this.dim));
methods.add_method("ctx", |_, this, ()| Ok(this.ctx));
methods.add_method("vocab", |_, this, ()| Ok(this.vocab));
methods.add_method("device", |_, this, ()| Ok(this.device.clone()));
methods.add_method("dtype", |_, this, ()| Ok(this.dtype.clone()));
methods.add_method("pretrained", |_, this, ()| Ok(this.pretrained));
methods.add_method("forward_shape", |_, this, (batch, seq): (usize, usize)| {
Ok(vec![batch, seq, this.vocab])
});
}
}
impl Gpt2Handle {
pub(super) fn model(&self) -> Arc<Mutex<Gpt2Model>> {
Arc::clone(&self.inner)
}
pub(super) fn varmap(&self) -> Option<Arc<VarMap>> {
self.varmap.as_ref().map(Arc::clone)
}
#[cfg(test)]
pub(super) fn for_test_pretrained_like(mut self) -> Self {
self.varmap = None;
self.pretrained = true;
self
}
}
fn register_preset_ns(lua: &Lua, nn_table: &LuaTable, nn_dir: PathBuf) -> LuaResult<()> {
let preset = lua.create_table()?;
let gpt2_nn_dir = nn_dir.clone();
let gpt2 = lua.create_function(
move |_lua, (variant, opts): (String, Option<LuaTable>)| -> LuaResult<Gpt2Handle> {
build_gpt2_handle(&variant, opts.as_ref(), &gpt2_nn_dir)
},
)?;
preset.set("gpt2", gpt2)?;
let tinyllama_nn_dir = nn_dir.clone();
let tinyllama = lua.create_function(
move |_lua, (variant, opts): (String, Option<LuaTable>)| -> LuaResult<TinyLlamaHandle> {
build_tinyllama_handle(&variant, opts.as_ref(), &tinyllama_nn_dir)
},
)?;
preset.set("tinyllama", tinyllama)?;
let llama = lua.create_function(
move |_lua, (variant, opts): (String, Option<LuaTable>)| -> LuaResult<LlamaHandle> {
build_llama_handle(&variant, opts.as_ref())
},
)?;
preset.set("llama", llama)?;
let neutral_nn_dir = nn_dir.clone();
let preset_call = lua.create_function(
move |_lua,
(_self, arch, variant, opts): (LuaTable, String, String, Option<LuaTable>)|
-> LuaResult<NnHandle> {
build_neutral_preset(&arch, &variant, opts.as_ref(), &neutral_nn_dir)
},
)?;
let preset_meta = lua.create_table()?;
preset_meta.set("__call", preset_call)?;
preset.set_metatable(Some(preset_meta))?;
nn_table.set("preset", preset)?;
Ok(())
}
fn build_neutral_preset(
arch: &str,
variant: &str,
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
) -> LuaResult<NnHandle> {
let ops = resolve_arch_ops(arch).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.preset: arch '{arch}' not registered \
(expected one of {}); qwen2 / phi / gemma are declared in \
SUPPORTED_ARCHITECTURE_FAMILIES but do not yet have a \
bridge preset entry",
registered_arch_names().join(" / ")
))
})?;
(ops.build_preset)(variant, opts, nn_dir)
}
struct ArchOps {
build_preset: fn(&str, Option<&LuaTable>, &std::path::Path) -> LuaResult<NnHandle>,
build_from_safetensors: Option<fn(&NnCardMeta, &std::path::Path) -> LuaResult<NnHandle>>,
#[allow(dead_code)]
build_from_wrap: Option<fn(&NnCardMeta, &NnHandle) -> LuaResult<NnHandle>>,
}
const ARCH_OPS: &[(&str, ArchOps)] = &[
(
"gpt2",
ArchOps {
build_preset: preset_gpt2_neutral,
build_from_safetensors: Some(gpt2_from_safetensors),
build_from_wrap: Some(wrap_gpt2_from_card),
},
),
(
"tinyllama",
ArchOps {
build_preset: preset_tinyllama_neutral,
build_from_safetensors: Some(tinyllama_from_safetensors),
build_from_wrap: Some(wrap_tinyllama_from_card),
},
),
(
"llama",
ArchOps {
build_preset: preset_llama_neutral,
build_from_safetensors: None,
build_from_wrap: None,
},
),
];
fn resolve_arch_ops(arch: &str) -> Option<&'static ArchOps> {
for (family, ops) in ARCH_OPS {
if arch == *family {
return Some(ops);
}
if let Some(rest) = arch.strip_prefix(*family) {
if rest.starts_with('-') {
return Some(ops);
}
}
}
None
}
fn registered_arch_names() -> Vec<&'static str> {
ARCH_OPS.iter().map(|(name, _)| *name).collect()
}
fn preset_gpt2_neutral(
variant: &str,
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
) -> LuaResult<NnHandle> {
build_gpt2_handle(variant, opts, nn_dir).map(NnHandle::Gpt2)
}
fn preset_tinyllama_neutral(
variant: &str,
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
) -> LuaResult<NnHandle> {
build_tinyllama_handle(variant, opts, nn_dir).map(NnHandle::TinyLlama)
}
fn preset_llama_neutral(
variant: &str,
opts: Option<&LuaTable>,
_nn_dir: &std::path::Path,
) -> LuaResult<NnHandle> {
build_llama_handle(variant, opts).map(NnHandle::Llama)
}
fn gpt2_from_safetensors(meta: &NnCardMeta, path: &std::path::Path) -> LuaResult<NnHandle> {
let mut cfg = Gpt2Config::from_variant(&meta.architecture).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load: unknown gpt2 variant {:?} on card {:?}",
meta.architecture, meta.name
))
})?;
apply_candle_branch_device_dtype("alc.nn.card.load", meta, &mut cfg.device, &mut cfg.dtype)?;
guard_device_dtype_matrix("alc.nn.card.load", &cfg.device, cfg.dtype)?;
let model = Gpt2Model::from_safetensors_file(&cfg, path)
.map_err(|e| LuaError::external(format!("alc.nn.card.load: {e}")))?;
let (device_str, dtype_str) = candle_branch_device_dtype_strings(meta, &cfg.device, cfg.dtype);
Ok(NnHandle::Gpt2(Gpt2Handle {
inner: Arc::new(Mutex::new(model)),
varmap: None,
variant: meta.architecture.clone(),
layers: cfg.layers,
heads: cfg.heads,
dim: cfg.dim,
ctx: cfg.ctx,
vocab: cfg.vocab,
device: device_str,
dtype: dtype_str,
pretrained: true,
has_lora: false,
}))
}
fn tinyllama_from_safetensors(meta: &NnCardMeta, path: &std::path::Path) -> LuaResult<NnHandle> {
let mut cfg = TinyLlamaConfig::from_variant(&meta.architecture).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load: unknown tinyllama variant {:?} on card {:?}",
meta.architecture, meta.name
))
})?;
apply_candle_branch_device_dtype("alc.nn.card.load", meta, &mut cfg.device, &mut cfg.dtype)?;
guard_device_dtype_matrix("alc.nn.card.load", &cfg.device, cfg.dtype)?;
let model = TinyLlamaModel::from_safetensors_file(&cfg, path)
.map_err(|e| LuaError::external(format!("alc.nn.card.load: {e}")))?;
let (device_str, dtype_str) = candle_branch_device_dtype_strings(meta, &cfg.device, cfg.dtype);
Ok(NnHandle::TinyLlama(TinyLlamaHandle {
inner: Arc::new(Mutex::new(model)),
varmap: None,
variant: meta.architecture.clone(),
layers: cfg.layers,
heads: cfg.heads,
kv_heads: cfg.kv_heads,
dim: cfg.dim,
ctx: cfg.ctx,
vocab: cfg.vocab,
device: device_str,
dtype: dtype_str,
pretrained: true,
has_lora: false,
}))
}
fn apply_candle_branch_device_dtype(
ctx: &str,
meta: &NnCardMeta,
device: &mut Device,
dtype: &mut DType,
) -> LuaResult<()> {
if let Some(candle) = &meta.candle {
if let Some(device_str) = &candle.device {
*device = parse_device_for(ctx, device_str)?;
}
if let Some(dtype_str) = &candle.dtype {
*dtype = parse_dtype_for(ctx, dtype_str)?;
}
}
Ok(())
}
fn candle_branch_device_dtype_strings(
meta: &NnCardMeta,
effective_device: &Device,
effective_dtype: DType,
) -> (String, String) {
let device_str = meta
.candle
.as_ref()
.and_then(|c| c.device.clone())
.unwrap_or_else(|| device_display(effective_device));
let dtype_str = meta
.candle
.as_ref()
.and_then(|c| c.dtype.clone())
.unwrap_or_else(|| dtype_display(effective_dtype));
(device_str, dtype_str)
}
fn device_display(d: &Device) -> String {
match d {
Device::Cpu => "cpu".into(),
Device::Cuda(_) => "cuda".into(),
Device::Metal(_) => "metal".into(),
}
}
fn wrap_gpt2_from_card(meta: &NnCardMeta, base: &NnHandle) -> LuaResult<NnHandle> {
let gpt2 = base.as_gpt2().ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_wrap: gpt2 card requires a gpt2 base handle; got '{}'",
base.arch()
))
})?;
let card_id = meta.name.as_str();
let wrapped = wrap_gpt2_lora_from_meta("alc.nn.card.load_wrap", card_id, meta, gpt2)?;
Ok(NnHandle::Gpt2(wrapped))
}
fn wrap_tinyllama_from_card(meta: &NnCardMeta, base: &NnHandle) -> LuaResult<NnHandle> {
let tll = base.as_tinyllama().ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_wrap: tinyllama card requires a tinyllama base handle; got '{}'",
base.arch()
))
})?;
let card_id = meta.name.as_str();
let wrapped = wrap_tinyllama_lora_from_meta("alc.nn.card.load_wrap", card_id, meta, tll)?;
Ok(NnHandle::TinyLlama(wrapped))
}
pub(super) fn load_wrap_impl(
store: &FileCardStore,
card_id: &str,
base_handle: &LuaAnyUserData,
) -> LuaResult<NnHandle> {
let card = store
.get(card_id)
.map_err(|e| LuaError::external(format!("alc.nn.card.load_wrap: {e}")))?
.ok_or_else(|| {
LuaError::external(format!("alc.nn.card.load_wrap: card '{card_id}' not found"))
})?;
let mut meta = extract_nn_card_meta("alc.nn.card.load_wrap", card_id, &card)?;
meta.name = card_id.to_string();
match meta.training_path.as_str() {
"lora" => {}
"full_ft" | "merged" | "distillation" => {
return Err(LuaError::external(format!(
"alc.nn.card.load_wrap: card '{card_id}' has training_path=\"{}\"; \
self-contained cards do not need a base handle — call \
`alc.nn.card.load_handle(card_id)` instead",
meta.training_path
)));
}
other => {
return Err(LuaError::external(format!(
"alc.nn.card.load_wrap: card '{card_id}' has unknown training_path \
{other:?} (expected one of full_ft / lora / merged / distillation)"
)));
}
}
precheck_lora_card_meta("alc.nn.card.load_wrap", card_id, &meta)?;
let ops = resolve_arch_ops(&meta.architecture).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_wrap: card '{card_id}' architecture {:?} \
has no bridge dispatch (expected one of {})",
meta.architecture,
registered_arch_names().join(" / ")
))
})?;
let wrap = ops.build_from_wrap.ok_or_else(|| {
LuaError::external(format!(
"alc.nn.card.load_wrap: card '{card_id}' architecture {:?} \
does not support LoRA wrap load",
meta.architecture
))
})?;
let base_nn: NnHandle = if let Ok(nn) = base_handle.borrow::<NnHandle>() {
(*nn).clone()
} else if let Ok(g) = base_handle.borrow::<Gpt2Handle>() {
NnHandle::Gpt2(g.clone())
} else if let Ok(t) = base_handle.borrow::<TinyLlamaHandle>() {
NnHandle::TinyLlama(t.clone())
} else if let Ok(l) = base_handle.borrow::<LlamaHandle>() {
NnHandle::Llama(l.clone())
} else {
return Err(LuaError::external(
"alc.nn.card.load_wrap: base handle is not a recognised NnHandle / \
Gpt2Handle / TinyLlamaHandle / LlamaHandle",
));
};
wrap(&meta, &base_nn)
}
fn dtype_display(d: DType) -> String {
match d {
DType::F32 => "f32".into(),
DType::F16 => "f16".into(),
DType::BF16 => "bf16".into(),
DType::U8 => "u8".into(),
DType::U32 => "u32".into(),
DType::I64 => "i64".into(),
DType::F64 => "f64".into(),
_ => format!("{d:?}"),
}
}
#[derive(Clone)]
pub(super) struct LlamaHandle {
inner: Arc<LlamaAdapter>,
variant: String,
layers: usize,
heads: usize,
kv_heads: usize,
dim: usize,
ctx: usize,
vocab: usize,
device: String,
dtype: String,
}
impl mlua::UserData for LlamaHandle {
fn add_methods<M: mlua::UserDataMethods<Self>>(methods: &mut M) {
methods.add_method("variant", |_, this, ()| Ok(this.variant.clone()));
methods.add_method("layers", |_, this, ()| Ok(this.layers));
methods.add_method("heads", |_, this, ()| Ok(this.heads));
methods.add_method("kv_heads", |_, this, ()| Ok(this.kv_heads));
methods.add_method("dim", |_, this, ()| Ok(this.dim));
methods.add_method("ctx", |_, this, ()| Ok(this.ctx));
methods.add_method("vocab", |_, this, ()| Ok(this.vocab));
methods.add_method("device", |_, this, ()| Ok(this.device.clone()));
methods.add_method("dtype", |_, this, ()| Ok(this.dtype.clone()));
methods.add_method("forward_shape", |_, this, (batch, _seq): (usize, usize)| {
Ok(vec![batch, this.vocab])
});
}
}
impl LlamaHandle {
#[allow(dead_code)]
pub(super) fn adapter(&self) -> Arc<LlamaAdapter> {
Arc::clone(&self.inner)
}
}
#[derive(Clone)]
pub(super) struct TinyLlamaHandle {
inner: Arc<Mutex<TinyLlamaModel>>,
varmap: Option<Arc<VarMap>>,
variant: String,
layers: usize,
heads: usize,
kv_heads: usize,
dim: usize,
ctx: usize,
vocab: usize,
device: String,
dtype: String,
pretrained: bool,
pub(super) has_lora: bool,
}
impl mlua::UserData for TinyLlamaHandle {
fn add_methods<M: mlua::UserDataMethods<Self>>(methods: &mut M) {
methods.add_method("variant", |_, this, ()| Ok(this.variant.clone()));
methods.add_method("layers", |_, this, ()| Ok(this.layers));
methods.add_method("heads", |_, this, ()| Ok(this.heads));
methods.add_method("kv_heads", |_, this, ()| Ok(this.kv_heads));
methods.add_method("dim", |_, this, ()| Ok(this.dim));
methods.add_method("ctx", |_, this, ()| Ok(this.ctx));
methods.add_method("vocab", |_, this, ()| Ok(this.vocab));
methods.add_method("device", |_, this, ()| Ok(this.device.clone()));
methods.add_method("dtype", |_, this, ()| Ok(this.dtype.clone()));
methods.add_method("pretrained", |_, this, ()| Ok(this.pretrained));
methods.add_method("forward_shape", |_, this, (batch, seq): (usize, usize)| {
Ok(vec![batch, seq, this.vocab])
});
}
}
impl TinyLlamaHandle {
#[allow(dead_code)]
pub(super) fn model(&self) -> Arc<Mutex<TinyLlamaModel>> {
Arc::clone(&self.inner)
}
#[allow(dead_code)]
pub(super) fn varmap(&self) -> Option<Arc<VarMap>> {
self.varmap.as_ref().map(Arc::clone)
}
#[cfg(test)]
pub(super) fn for_test_pretrained_like(mut self) -> Self {
self.varmap = None;
self.pretrained = true;
self
}
}
fn build_llama_handle(variant: &str, opts: Option<&LuaTable>) -> LuaResult<LlamaHandle> {
let flash_attn = opts
.and_then(|t| t.get::<Option<bool>>("flash_attn").ok().flatten())
.unwrap_or(false);
let mut cfg = LlamaAdapterConfig::from_variant(variant, flash_attn).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.preset.llama: unknown variant '{variant}' \
(expected 'tiny' / '7b-v1' / '7b-v2', or one of their 'llama-*' aliases)"
))
})?;
let device_str = opts
.and_then(|t| t.get::<Option<String>>("device").ok().flatten())
.unwrap_or_else(|| "cpu".to_string());
let dtype_str = opts
.and_then(|t| t.get::<Option<String>>("dtype").ok().flatten())
.unwrap_or_else(|| default_dtype_for_device(&device_str).to_string());
let use_kv_cache = opts
.and_then(|t| t.get::<Option<bool>>("use_kv_cache").ok().flatten())
.unwrap_or(true);
cfg.device = parse_llama_device(&device_str)?;
cfg.dtype = parse_llama_dtype(&dtype_str)?;
cfg.use_kv_cache = use_kv_cache;
guard_device_dtype_matrix("alc.nn.preset.llama", &cfg.device, cfg.dtype)?;
let weights_paths = extract_weights_paths(opts)?;
let (adapter, cfg_snapshot) = if let Some(paths) = weights_paths {
let cfg_snapshot = cfg.clone();
let adapter = LlamaAdapter::from_safetensors_files(&paths, cfg)
.map_err(|e| LuaError::external(format!("alc.nn.preset.llama: {e}")))?;
(adapter, cfg_snapshot)
} else {
let cfg_snapshot = cfg.clone();
let vm = VarMap::new();
let vb = candle_nn::VarBuilder::from_varmap(&vm, cfg.dtype, &cfg.device);
let adapter = LlamaAdapter::load(vb, cfg)
.map_err(|e| LuaError::external(format!("alc.nn.preset.llama: {e}")))?;
drop(vm);
(adapter, cfg_snapshot)
};
Ok(LlamaHandle {
inner: Arc::new(adapter),
variant: variant.to_string(),
layers: cfg_snapshot.config.num_hidden_layers,
heads: cfg_snapshot.config.num_attention_heads,
kv_heads: cfg_snapshot.config.num_key_value_heads,
dim: cfg_snapshot.config.hidden_size,
ctx: cfg_snapshot.config.max_position_embeddings,
vocab: cfg_snapshot.config.vocab_size,
device: device_str,
dtype: dtype_str,
})
}
fn extract_weights_paths(opts: Option<&LuaTable>) -> LuaResult<Option<Vec<PathBuf>>> {
let Some(opts) = opts else {
return Ok(None);
};
let Some(raw) = opts.get::<Option<mlua::Value>>("weights").ok().flatten() else {
return Ok(None);
};
match raw {
mlua::Value::String(s) => Ok(Some(vec![PathBuf::from(s.to_str()?.to_string())])),
mlua::Value::Table(tbl) => {
let mut out = Vec::new();
for pair in tbl.sequence_values::<String>() {
out.push(PathBuf::from(pair?));
}
if out.is_empty() {
return Err(LuaError::external(
"alc.nn.preset.llama: opts.weights is empty; provide at least one path",
));
}
Ok(Some(out))
}
_ => Err(LuaError::external(
"alc.nn.preset.llama: opts.weights must be a string or an array of strings",
)),
}
}
fn parse_llama_device(s: &str) -> LuaResult<Device> {
parse_device_for("alc.nn.preset.llama", s)
}
fn parse_llama_dtype(s: &str) -> LuaResult<DType> {
parse_dtype_for("alc.nn.preset.llama", s)
}
pub(super) fn build_gpt2_handle(
variant: &str,
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
) -> LuaResult<Gpt2Handle> {
let mut cfg = Gpt2Config::from_variant(variant).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.preset.gpt2: unknown variant '{variant}' (expected 'medium' or 'large')"
))
})?;
let device_str = opts
.and_then(|t| t.get::<Option<String>>("device").ok().flatten())
.unwrap_or_else(|| "cpu".to_string());
let dtype_str = opts
.and_then(|t| t.get::<Option<String>>("dtype").ok().flatten())
.unwrap_or_else(|| default_dtype_for_device(&device_str).to_string());
let pretrained = opts
.and_then(|t| t.get::<Option<bool>>("pretrained").ok().flatten())
.unwrap_or(true);
cfg.device = parse_device(&device_str)?;
cfg.dtype = parse_dtype(&dtype_str)?;
guard_device_dtype_matrix("alc.nn.preset.gpt2", &cfg.device, cfg.dtype)?;
let (model, varmap) = if pretrained {
let cache_dir = nn_dir.to_path_buf();
let m = Gpt2Model::from_pretrained(variant, &cfg, &cache_dir)
.map_err(|e| LuaError::external(format!("alc.nn.preset.gpt2: {e}")))?;
(m, None)
} else {
let vm = VarMap::new();
let vs = candle_nn::VarBuilder::from_varmap(&vm, cfg.dtype, &cfg.device);
let m = Gpt2Model::new(&cfg, vs)
.map_err(|e| LuaError::external(format!("alc.nn.preset.gpt2: {e}")))?;
(m, Some(Arc::new(vm)))
};
Ok(Gpt2Handle {
inner: Arc::new(Mutex::new(model)),
varmap,
variant: variant.to_string(),
layers: cfg.layers,
heads: cfg.heads,
dim: cfg.dim,
ctx: cfg.ctx,
vocab: cfg.vocab,
device: device_str,
dtype: dtype_str,
pretrained,
has_lora: false,
})
}
fn parse_device(s: &str) -> LuaResult<Device> {
parse_device_for("alc.nn.preset.gpt2", s)
}
pub(super) fn build_tinyllama_handle(
variant: &str,
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
) -> LuaResult<TinyLlamaHandle> {
let mut cfg = TinyLlamaConfig::from_variant(variant).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.preset.tinyllama: unknown variant '{variant}' \
(expected 'tinyllama-1.1b' / '1.1b' / 'tinyllama-tiny' / 'tiny')"
))
})?;
let device_str = opts
.and_then(|t| t.get::<Option<String>>("device").ok().flatten())
.unwrap_or_else(|| "cpu".to_string());
let dtype_str = opts
.and_then(|t| t.get::<Option<String>>("dtype").ok().flatten())
.unwrap_or_else(|| default_dtype_for_device(&device_str).to_string());
let pretrained = opts
.and_then(|t| t.get::<Option<bool>>("pretrained").ok().flatten())
.unwrap_or(true);
cfg.device = parse_device_for("alc.nn.preset.tinyllama", &device_str)?;
cfg.dtype = parse_dtype_for("alc.nn.preset.tinyllama", &dtype_str)?;
guard_device_dtype_matrix("alc.nn.preset.tinyllama", &cfg.device, cfg.dtype)?;
let (model, varmap) = if pretrained {
let cache_dir = nn_dir.to_path_buf();
let m = TinyLlamaModel::from_pretrained(variant, &cfg, &cache_dir)
.map_err(|e| LuaError::external(format!("alc.nn.preset.tinyllama: {e}")))?;
(m, None)
} else {
let vm = VarMap::new();
let vs = candle_nn::VarBuilder::from_varmap(&vm, cfg.dtype, &cfg.device);
let m = TinyLlamaModel::new(&cfg, vs)
.map_err(|e| LuaError::external(format!("alc.nn.preset.tinyllama: {e}")))?;
(m, Some(Arc::new(vm)))
};
Ok(TinyLlamaHandle {
inner: Arc::new(Mutex::new(model)),
varmap,
variant: variant.to_string(),
layers: cfg.layers,
heads: cfg.heads,
kv_heads: cfg.kv_heads,
dim: cfg.dim,
ctx: cfg.ctx,
vocab: cfg.vocab,
device: device_str,
dtype: dtype_str,
pretrained,
has_lora: false,
})
}
fn default_dtype_for_device(device: &str) -> &'static str {
if device.starts_with("cuda") {
"bf16"
} else if device.starts_with("metal") {
"f16"
} else {
"f32"
}
}
fn parse_dtype(s: &str) -> LuaResult<DType> {
parse_dtype_for("alc.nn.preset.gpt2", s)
}
fn parse_device_for(preset: &str, s: &str) -> LuaResult<Device> {
if s == "cpu" {
return Ok(Device::Cpu);
}
if let Some(rest) = s.strip_prefix("cuda:") {
let ord: usize = rest.parse().map_err(|e| {
LuaError::external(format!("{preset}: invalid cuda ordinal '{rest}': {e}"))
})?;
return Device::new_cuda(ord)
.map_err(|e| LuaError::external(format!("{preset}: cuda:{ord} unavailable: {e}")));
}
if s == "cuda" {
return Device::new_cuda(0)
.map_err(|e| LuaError::external(format!("{preset}: cuda unavailable: {e}")));
}
if let Some(rest) = s.strip_prefix("metal:") {
let ord: usize = rest.parse().map_err(|e| {
LuaError::external(format!("{preset}: invalid metal ordinal '{rest}': {e}"))
})?;
return Device::new_metal(ord)
.map_err(|e| LuaError::external(format!("{preset}: metal:{ord} unavailable: {e}")));
}
if s == "metal" {
return Device::new_metal(0)
.map_err(|e| LuaError::external(format!("{preset}: metal unavailable: {e}")));
}
Err(LuaError::external(format!(
"{preset}: unknown device '{s}' (expected 'cpu', 'cuda', 'cuda:N', 'metal', or 'metal:N')"
)))
}
fn parse_dtype_for(preset: &str, s: &str) -> LuaResult<DType> {
match s {
"f32" | "fp32" => Ok(DType::F32),
"bf16" => Ok(DType::BF16),
"f16" | "fp16" => Ok(DType::F16),
other => Err(LuaError::external(format!(
"{preset}: unknown dtype '{other}' (expected 'f32', 'bf16', or 'f16')"
))),
}
}
fn guard_device_dtype_matrix(preset: &str, device: &Device, dtype: DType) -> LuaResult<()> {
if dtype != DType::BF16 {
return Ok(());
}
match device {
Device::Cpu => Err(LuaError::external(format!(
"{preset}: bf16 dtype requires a CUDA device (use dtype='f32' on CPU, \
or dtype='f16' on Metal)"
))),
Device::Metal(_) => Err(LuaError::external(format!(
"{preset}: bf16 dtype is not supported on Metal (use dtype='f16' or 'f32' on Metal, \
or move to CUDA for bf16)"
))),
_ => Ok(()),
}
}
pub(super) struct DatasetHandle {
inner: Mutex<Box<dyn Dataset + Send>>,
source: String,
batch_size: usize,
ctx_len: usize,
}
impl DatasetHandle {
pub(super) fn inner_lock(
&self,
) -> LuaResult<std::sync::MutexGuard<'_, Box<dyn Dataset + Send>>> {
self.inner.lock().map_err(|e| {
LuaError::external(format!(
"alc.nn.trainer.run_lora_ft: dataset lock poisoned: {e}"
))
})
}
#[cfg(test)]
pub(super) fn for_test(
inner: Box<dyn Dataset + Send>,
source: String,
batch_size: usize,
ctx_len: usize,
) -> Self {
Self {
inner: Mutex::new(inner),
source,
batch_size,
ctx_len,
}
}
}
impl mlua::UserData for DatasetHandle {
fn add_methods<M: mlua::UserDataMethods<Self>>(methods: &mut M) {
methods.add_method("source", |_, this, ()| Ok(this.source.clone()));
methods.add_method("batch_size", |_, this, ()| Ok(this.batch_size));
methods.add_method("ctx_len", |_, this, ()| Ok(this.ctx_len));
methods.add_method("len_hint", |_, this, ()| {
let ds = this.inner.lock().map_err(|e| {
LuaError::external(format!("alc.nn.data: dataset lock poisoned: {e}"))
})?;
Ok(ds.len_hint())
});
methods.add_method_mut(
"next_batch",
|lua, this, ()| -> LuaResult<Option<LuaTable>> {
let mut ds = this.inner.lock().map_err(|e| {
LuaError::external(format!("alc.nn.data: dataset lock poisoned: {e}"))
})?;
match ds
.next_batch()
.map_err(|e| LuaError::external(format!("alc.nn.data.next_batch: {e}")))?
{
Some(batch) => Ok(Some(batch_to_lua(lua, batch)?)),
None => Ok(None),
}
},
);
}
}
fn batch_to_lua(lua: &Lua, batch: Batch) -> LuaResult<LuaTable> {
let out = lua.create_table()?;
let rows = lua.create_table()?;
for (i, row) in batch.input_ids.into_iter().enumerate() {
let arr = lua.create_table()?;
for (j, id) in row.into_iter().enumerate() {
arr.set(j + 1, id)?;
}
rows.set(i + 1, arr)?;
}
out.set("input_ids", rows)?;
out.set("is_last", batch.is_last)?;
Ok(out)
}
fn register_data_ns(
lua: &Lua,
nn_table: &LuaTable,
card_store: Arc<FileCardStore>,
nn_dir: PathBuf,
) -> LuaResult<()> {
let data = lua.create_table()?;
let jsonl_tok_dir = nn_dir.join("tokenizers");
let jsonl = lua.create_function(
move |_lua, (path, opts): (String, Option<LuaTable>)| -> LuaResult<DatasetHandle> {
let dopts = extract_dataset_opts(opts.as_ref())?;
let tokenizer_name = opts
.as_ref()
.and_then(|t| t.get::<Option<String>>("tokenizer").ok().flatten())
.unwrap_or_else(|| "gpt2".to_string());
let tok = HfTokenizer::load_cached(&tokenizer_name, &jsonl_tok_dir)
.map_err(|e| LuaError::external(format!("alc.nn.data.jsonl: {e}")))?;
let ds = JsonlDataset::new(std::path::Path::new(&path), dopts.clone(), tok)
.map_err(|e| LuaError::external(format!("alc.nn.data.jsonl: {e}")))?;
Ok(DatasetHandle {
inner: Mutex::new(Box::new(ds)),
source: format!("jsonl:{path}"),
batch_size: dopts.batch_size,
ctx_len: dopts.ctx_len,
})
},
)?;
data.set("jsonl", jsonl)?;
let parquet = lua.create_function(
move |_lua, (path, opts): (String, Option<LuaTable>)| -> LuaResult<DatasetHandle> {
let dopts = extract_dataset_opts(opts.as_ref())?;
let ds = ParquetDataset::new(std::path::Path::new(&path), dopts.clone());
Ok(DatasetHandle {
inner: Mutex::new(Box::new(ds)),
source: format!("parquet:{path}"),
batch_size: dopts.batch_size,
ctx_len: dopts.ctx_len,
})
},
)?;
data.set("parquet", parquet)?;
let from_card_store = Arc::clone(&card_store);
let from_card_tok_dir = nn_dir.join("tokenizers");
let from_card = lua.create_function(
move |_lua, (card_id, opts): (String, Option<LuaTable>)| -> LuaResult<DatasetHandle> {
let dopts = extract_dataset_opts(opts.as_ref())?;
let tokenizer_name = opts
.as_ref()
.and_then(|t| t.get::<Option<String>>("tokenizer").ok().flatten())
.unwrap_or_else(|| "gpt2".to_string());
let tok = HfTokenizer::load_cached(&tokenizer_name, &from_card_tok_dir)
.map_err(|e| LuaError::external(format!("alc.nn.data.from_card: {e}")))?;
let samples = from_card_store
.read_samples(&card_id, SamplesQuery::default())
.map_err(|e| LuaError::external(format!("alc.nn.data.from_card: {e}")))?;
let mut rows: Vec<Vec<u32>> = Vec::with_capacity(samples.len());
for (idx, sample) in samples.into_iter().enumerate() {
let prompt = sample
.get("prompt")
.and_then(|v| v.as_str())
.unwrap_or_default();
let response = sample
.get("response")
.and_then(|v| v.as_str())
.unwrap_or_default();
let text = if !prompt.is_empty() && !response.is_empty() {
format!("{prompt}\n{response}")
} else if !prompt.is_empty() {
prompt.to_string()
} else if !response.is_empty() {
response.to_string()
} else {
continue;
};
let ids = tok.encode(&text).map_err(|e| {
LuaError::external(format!("alc.nn.data.from_card: sample {idx}: {e}"))
})?;
if !ids.is_empty() {
rows.push(ids);
}
}
let ds = TokenizedDataset::new(rows, dopts.clone());
Ok(DatasetHandle {
inner: Mutex::new(Box::new(ds)),
source: format!("card:{card_id}"),
batch_size: dopts.batch_size,
ctx_len: dopts.ctx_len,
})
},
)?;
data.set("from_card", from_card)?;
let synthetic = lua.create_function(
move |_lua, (rows_tbl, opts): (LuaTable, Option<LuaTable>)| -> LuaResult<DatasetHandle> {
let dopts = extract_dataset_opts(opts.as_ref())?;
let row_count = rows_tbl.raw_len();
if row_count == 0 {
return Err(LuaError::external(
"alc.nn.data.synthetic: rows must be a non-empty array of token id \
sequences (each row itself an array of u32)"
.to_string(),
));
}
let mut rows: Vec<Vec<u32>> = Vec::with_capacity(row_count);
for i in 1..=row_count {
let row: LuaTable = rows_tbl.get(i).map_err(|e| {
LuaError::external(format!("alc.nn.data.synthetic: row {i} not a table: {e}"))
})?;
let len = row.raw_len();
if len == 0 {
return Err(LuaError::external(format!(
"alc.nn.data.synthetic: row {i} is empty (need at least 1 token)"
)));
}
let mut ids: Vec<u32> = Vec::with_capacity(len);
for j in 1..=len {
let id: u32 = row.get(j).map_err(|e| {
LuaError::external(format!(
"alc.nn.data.synthetic: row {i} token {j} not a u32 integer: {e}"
))
})?;
ids.push(id);
}
rows.push(ids);
}
let ds = TokenizedDataset::new(rows, dopts.clone());
Ok(DatasetHandle {
inner: Mutex::new(Box::new(ds)),
source: format!("synthetic:{row_count}rows"),
batch_size: dopts.batch_size,
ctx_len: dopts.ctx_len,
})
},
)?;
data.set("synthetic", synthetic)?;
nn_table.set("data", data)?;
Ok(())
}
fn extract_dataset_opts(opts: Option<&LuaTable>) -> LuaResult<DatasetOpts> {
let mut d = DatasetOpts::default();
if let Some(t) = opts {
if let Some(v) = t.get::<Option<usize>>("batch_size")? {
d.batch_size = v;
}
if let Some(v) = t.get::<Option<usize>>("ctx_len")? {
d.ctx_len = v;
}
if let Some(v) = t.get::<Option<bool>>("shuffle")? {
d.shuffle = v;
}
if let Some(v) = t.get::<Option<u32>>("pad_id")? {
d.pad_id = v;
}
if let Some(v) = t.get::<Option<String>>("text_field")? {
d.text_field = v;
}
}
if d.batch_size == 0 {
return Err(LuaError::external(
"alc.nn.data: batch_size must be >= 1".to_string(),
));
}
if d.ctx_len == 0 {
return Err(LuaError::external(
"alc.nn.data: ctx_len must be >= 1".to_string(),
));
}
Ok(d)
}
fn register_trainer_ns(lua: &Lua, nn_table: &LuaTable, nn_dir: PathBuf) -> LuaResult<()> {
let trainer = lua.create_table()?;
let lease = Arc::new(TrainingLease::new());
let full_ft_lease = Arc::clone(&lease);
let full_ft_dir = nn_dir.clone();
let full_ft = lua.create_function(
move |lua,
(handle, dataset, opts): (LuaAnyUserData, LuaAnyUserData, Option<LuaTable>)|
-> LuaResult<LuaTable> {
full_ft_impl(
lua,
&handle,
&dataset,
opts.as_ref(),
&full_ft_dir,
Arc::clone(&full_ft_lease),
)
},
)?;
trainer.set("full_ft", full_ft)?;
let lora_lease = Arc::clone(&lease);
let lora_dir = nn_dir.clone();
let lora = lua.create_function(
move |lua,
(handle, dataset, opts): (LuaAnyUserData, LuaAnyUserData, Option<LuaTable>)|
-> LuaResult<LuaTable> {
lora_impl(
lua,
&handle,
&dataset,
opts.as_ref(),
&lora_dir,
Arc::clone(&lora_lease),
)
},
)?;
trainer.set("lora", lora)?;
let distill_lease = Arc::clone(&lease);
let distill_dir = nn_dir;
let distill = lua.create_function(
move |lua,
(handle, dataset, opts): (LuaAnyUserData, LuaAnyUserData, Option<LuaTable>)|
-> LuaResult<LuaTable> {
distill_impl(
lua,
&handle,
&dataset,
opts.as_ref(),
&distill_dir,
Arc::clone(&distill_lease),
)
},
)?;
trainer.set("distill", distill)?;
nn_table.set("trainer", trainer)?;
Ok(())
}
fn full_ft_impl(
lua: &Lua,
handle: &LuaAnyUserData,
dataset: &LuaAnyUserData,
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
lease: Arc<TrainingLease>,
) -> LuaResult<LuaTable> {
let cfg = extract_full_ft_opts(opts)?;
let (ckpt_dir, ckpt_prefix) = resolve_ckpt_dest(opts, nn_dir, "full_ft")?;
let card_id = ckpt_prefix.clone();
let gpt2 = handle.borrow::<Gpt2Handle>()?;
let model_arc = gpt2.model();
let vm_arc = gpt2.varmap().ok_or_else(|| {
LuaError::external(
"alc.nn.trainer.full_ft: handle was built with pretrained=true; \
full-fine-tune requires a from-scratch handle (pretrained=false)"
.to_string(),
)
})?;
drop(gpt2);
let ds_guard = dataset.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_guard
.inner
.lock()
.map_err(|e| LuaError::external(format!("alc.nn.trainer.full_ft: dataset lock: {e}")))?;
let loss_fn = CrossEntropyLoss::new();
let model = model_arc
.lock()
.map_err(|e| LuaError::external(format!("alc.nn.trainer.full_ft: model lock: {e}")))?;
let result = run_full_ft(
&*model,
&vm_arc,
ds_lock.as_mut(),
&cfg,
&loss_fn,
&ckpt_dir,
&ckpt_prefix,
lease,
);
drop(model);
drop(ds_lock);
drop(ds_guard);
let ckpt = result.map_err(train_err_to_lua)?;
checkpoint_to_lua(lua, &ckpt, &card_id, None)
}
fn lora_impl(
lua: &Lua,
handle: &LuaAnyUserData,
dataset: &LuaAnyUserData,
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
lease: Arc<TrainingLease>,
) -> LuaResult<LuaTable> {
let train_cfg = extract_full_ft_opts(opts)?;
let lora_cfg = extract_lora_cfg(opts)?;
let (ckpt_dir, card_id) = resolve_ckpt_dest(opts, nn_dir, "lora")?;
let gpt2 = handle.borrow::<Gpt2Handle>()?;
let model_arc = gpt2.model();
let base_bundle_ref = match opts {
Some(t) => t
.get::<Option<String>>("base_bundle_ref")?
.unwrap_or_else(|| format!("nn/{}", gpt2.variant)),
None => format!("nn/{}", gpt2.variant),
};
drop(gpt2);
let ds_guard = dataset.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_guard
.inner
.lock()
.map_err(|e| LuaError::external(format!("alc.nn.trainer.lora: dataset lock: {e}")))?;
let loss_fn = CrossEntropyLoss::new();
let mut model = model_arc
.lock()
.map_err(|e| LuaError::external(format!("alc.nn.trainer.lora: model lock: {e}")))?;
let result = run_lora_ft(
&mut *model,
ds_lock.as_mut(),
&lora_cfg,
&train_cfg,
&loss_fn,
&ckpt_dir,
&card_id,
lease,
);
drop(model);
drop(ds_lock);
drop(ds_guard);
let ckpt = result.map_err(train_err_to_lua)?;
let lora_tbl = lua.create_table()?;
lora_tbl.set("rank", lora_cfg.rank as u32)?;
lora_tbl.set("alpha", lora_cfg.alpha as u32)?;
lora_tbl.set("base_bundle_ref", base_bundle_ref)?;
let target_modules_tbl = lua.create_table()?;
for (i, m) in lora_cfg.target_modules.iter().enumerate() {
target_modules_tbl.set(i + 1, m.clone())?;
}
lora_tbl.set("target_modules", target_modules_tbl)?;
lora_tbl.set("dropout", lora_cfg.dropout)?;
let delta_path = ckpt_dir.join("nn").join(&ckpt.bundle_ref);
lora_tbl.set("delta_path", delta_path.to_string_lossy().to_string())?;
checkpoint_to_lua(lua, &ckpt, &card_id, Some(lora_tbl))
}
fn distill_impl(
lua: &Lua,
handle: &LuaAnyUserData,
dataset: &LuaAnyUserData,
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
lease: Arc<TrainingLease>,
) -> LuaResult<LuaTable> {
let hyperparams = extract_full_ft_opts(opts)?;
let loss_kind = extract_distill_loss_kind(opts)?;
let spec = DistillSpec {
hyperparams,
loss_kind,
};
let (ckpt_dir, ckpt_prefix) = resolve_ckpt_dest(opts, nn_dir, "distill")?;
let card_id = ckpt_prefix.clone();
let gpt2 = handle.borrow::<Gpt2Handle>()?;
let model_arc = gpt2.model();
let vm_arc = gpt2.varmap().ok_or_else(|| {
LuaError::external(
"alc.nn.trainer.distill: handle was built with pretrained=true; \
distillation requires a from-scratch student handle (pretrained=false)"
.to_string(),
)
})?;
drop(gpt2);
let ds_guard = dataset.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_guard
.inner
.lock()
.map_err(|e| LuaError::external(format!("alc.nn.trainer.distill: dataset lock: {e}")))?;
let model = model_arc
.lock()
.map_err(|e| LuaError::external(format!("alc.nn.trainer.distill: model lock: {e}")))?;
let result = run_distill(
&*model,
&vm_arc,
ds_lock.as_mut(),
&spec,
&ckpt_dir,
&ckpt_prefix,
lease,
);
drop(model);
drop(ds_lock);
drop(ds_guard);
let ckpt = result.map_err(train_err_to_lua)?;
checkpoint_to_lua(lua, &ckpt, &card_id, None)
}
pub(super) fn extract_full_ft_opts(opts: Option<&LuaTable>) -> LuaResult<FullFtConfig> {
let mut cfg = FullFtConfig::default();
let Some(t) = opts else {
return Ok(cfg);
};
if let Some(v) = t.get::<Option<f64>>("lr")? {
cfg.lr = v;
}
if let Some(v) = t.get::<Option<usize>>("batch_size")? {
cfg.batch_size = v;
}
if let Some(v) = t.get::<Option<usize>>("grad_accum")? {
cfg.grad_accum = v;
}
if let Some(v) = t.get::<Option<usize>>("steps")? {
cfg.steps = v;
}
if let Some(v) = t.get::<Option<usize>>("warmup")? {
cfg.warmup = v;
}
if let Some(v) = t.get::<Option<String>>("schedule")? {
cfg.schedule = parse_schedule(&v)?;
}
if let Some(v) = t.get::<Option<f64>>("weight_decay")? {
cfg.weight_decay = v;
}
if let Some(v) = t.get::<Option<usize>>("ckpt_every")? {
cfg.ckpt_every = v;
}
if let Some(v) = t.get::<Option<usize>>("ckpt_keep")? {
cfg.ckpt_keep = v;
}
if cfg.batch_size == 0 {
return Err(LuaError::external(
"alc.nn.trainer: batch_size must be >= 1".to_string(),
));
}
Ok(cfg)
}
fn parse_schedule(s: &str) -> LuaResult<ScheduleKind> {
match s {
"cosine" | "cosine_with_warmup" => Ok(ScheduleKind::CosineWithWarmup),
"constant" => Ok(ScheduleKind::Constant),
other => Err(LuaError::external(format!(
"alc.nn.trainer: unknown schedule '{other}' \
(expected 'cosine' or 'constant')"
))),
}
}
fn extract_lora_cfg(opts: Option<&LuaTable>) -> LuaResult<LoraConfig> {
let t = opts.ok_or_else(|| {
LuaError::external(
"alc.nn.trainer.lora: opts table is required (need at least rank and alpha)"
.to_string(),
)
})?;
let rank = t.get::<Option<usize>>("rank")?.ok_or_else(|| {
LuaError::external("alc.nn.trainer.lora: opts.rank is required".to_string())
})?;
let alpha_raw = t.get::<Option<f32>>("alpha")?.ok_or_else(|| {
LuaError::external("alc.nn.trainer.lora: opts.alpha is required".to_string())
})?;
if rank == 0 {
return Err(LuaError::external(
"alc.nn.trainer.lora: rank must be >= 1".to_string(),
));
}
if !alpha_raw.is_finite() || alpha_raw <= 0.0 {
return Err(LuaError::external(
"alc.nn.trainer.lora: alpha must be a positive finite number".to_string(),
));
}
let dropout = t.get::<Option<f32>>("dropout")?.unwrap_or(0.0);
let target_modules = match t.get::<Option<LuaTable>>("target_modules")? {
Some(list) => {
let mut names = Vec::new();
for pair in list.pairs::<LuaValue, String>() {
let (_k, name) = pair?;
names.push(name);
}
if names.is_empty() {
return Err(LuaError::external(
"alc.nn.trainer.lora: opts.target_modules must not be empty".to_string(),
));
}
names
}
None => LoraConfig::default_targets(),
};
Ok(LoraConfig {
rank,
alpha: alpha_raw,
target_modules,
dropout,
})
}
fn extract_distill_loss_kind(opts: Option<&LuaTable>) -> LuaResult<DistillLossKind> {
let raw = match opts {
Some(t) => t
.get::<Option<String>>("loss_kind")?
.unwrap_or_else(|| "ce".to_string()),
None => "ce".to_string(),
};
match raw.as_str() {
"ce" => Ok(DistillLossKind::Ce),
other => Err(LuaError::external(format!(
"alc.nn.trainer.distill: unknown loss_kind '{other}' (expected 'ce')"
))),
}
}
fn resolve_ckpt_dest(
opts: Option<&LuaTable>,
nn_dir: &std::path::Path,
stage: &str,
) -> LuaResult<(PathBuf, String)> {
let ckpt_dir = match opts {
Some(t) => t
.get::<Option<String>>("ckpt_dir")?
.map(PathBuf::from)
.unwrap_or_else(|| nn_dir.join("ckpt")),
None => nn_dir.join("ckpt"),
};
std::fs::create_dir_all(&ckpt_dir).map_err(|e| {
LuaError::external(format!(
"alc.nn.trainer.{stage}: mkdir {:?}: {e}",
ckpt_dir.display()
))
})?;
let ckpt_prefix = match opts {
Some(t) => match t.get::<Option<String>>("card_id")? {
Some(id) => sanitize_name(&id),
None => generate_card_id(stage),
},
None => generate_card_id(stage),
};
Ok((ckpt_dir, ckpt_prefix))
}
fn train_err_to_lua(e: TrainError) -> LuaError {
LuaError::external(format!("alc.nn.trainer: {e}"))
}
fn checkpoint_to_lua(
lua: &Lua,
ckpt: &algocline_nn::train::Checkpoint,
card_id: &str,
lora_branch: Option<LuaTable>,
) -> LuaResult<LuaTable> {
let out = lua.create_table()?;
out.set("bundle_ref", ckpt.bundle_ref.clone())?;
out.set("card_id", card_id.to_string())?;
out.set("step", ckpt.step)?;
out.set("train_loss", ckpt.train_loss)?;
match ckpt.val_loss {
Some(v) => out.set("val_loss", v)?,
None => out.set("val_loss", LuaValue::Nil)?,
}
let metrics = lua.create_table()?;
for (k, v) in &ckpt.metrics {
metrics.set(k.as_str(), *v)?;
}
out.set("metrics", metrics)?;
if let Some(lora) = lora_branch {
out.set("lora", lora)?;
}
Ok(out)
}
#[allow(dead_code)]
#[derive(Clone)]
pub(super) enum NnHandle {
Gpt2(Gpt2Handle),
TinyLlama(TinyLlamaHandle),
Llama(LlamaHandle),
}
impl NnHandle {
#[allow(dead_code)]
pub(super) fn arch(&self) -> &'static str {
match self {
Self::Gpt2(_) => "gpt2",
Self::TinyLlama(_) => "tinyllama",
Self::Llama(_) => "llama",
}
}
#[allow(dead_code)]
pub(super) fn as_gpt2(&self) -> Option<&Gpt2Handle> {
match self {
Self::Gpt2(h) => Some(h),
_ => None,
}
}
#[allow(dead_code)]
pub(super) fn as_tinyllama(&self) -> Option<&TinyLlamaHandle> {
match self {
Self::TinyLlama(h) => Some(h),
_ => None,
}
}
#[allow(dead_code)]
pub(super) fn as_llama(&self) -> Option<&LlamaHandle> {
match self {
Self::Llama(h) => Some(h),
_ => None,
}
}
#[allow(dead_code)]
pub(super) fn arch_family_variant(&self) -> String {
let (family, variant) = match self {
Self::Gpt2(h) => ("gpt2", h.variant.as_str()),
Self::TinyLlama(h) => ("tinyllama", h.variant.as_str()),
Self::Llama(h) => ("llama", h.variant.as_str()),
};
let prefix = format!("{family}-");
if variant.starts_with(&prefix) {
variant.to_string()
} else {
format!("{prefix}{variant}")
}
}
#[allow(dead_code)]
pub(super) fn is_lora_wrapped(&self) -> bool {
match self {
Self::Gpt2(h) => h.has_lora,
Self::TinyLlama(h) => h.has_lora,
Self::Llama(_) => false,
}
}
}
impl mlua::UserData for NnHandle {
fn add_methods<M: mlua::UserDataMethods<Self>>(methods: &mut M) {
methods.add_method("arch", |_, this, ()| Ok(this.arch()));
methods.add_method("variant", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.variant.clone()),
NnHandle::TinyLlama(h) => Ok(h.variant.clone()),
NnHandle::Llama(h) => Ok(h.variant.clone()),
});
methods.add_method("layers", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.layers),
NnHandle::TinyLlama(h) => Ok(h.layers),
NnHandle::Llama(h) => Ok(h.layers),
});
methods.add_method("heads", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.heads),
NnHandle::TinyLlama(h) => Ok(h.heads),
NnHandle::Llama(h) => Ok(h.heads),
});
methods.add_method("kv_heads", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.heads),
NnHandle::TinyLlama(h) => Ok(h.kv_heads),
NnHandle::Llama(h) => Ok(h.kv_heads),
});
methods.add_method("dim", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.dim),
NnHandle::TinyLlama(h) => Ok(h.dim),
NnHandle::Llama(h) => Ok(h.dim),
});
methods.add_method("ctx", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.ctx),
NnHandle::TinyLlama(h) => Ok(h.ctx),
NnHandle::Llama(h) => Ok(h.ctx),
});
methods.add_method("vocab", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.vocab),
NnHandle::TinyLlama(h) => Ok(h.vocab),
NnHandle::Llama(h) => Ok(h.vocab),
});
methods.add_method("device", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.device.clone()),
NnHandle::TinyLlama(h) => Ok(h.device.clone()),
NnHandle::Llama(h) => Ok(h.device.clone()),
});
methods.add_method("dtype", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.dtype.clone()),
NnHandle::TinyLlama(h) => Ok(h.dtype.clone()),
NnHandle::Llama(h) => Ok(h.dtype.clone()),
});
methods.add_method("pretrained", |_, this, ()| match this {
NnHandle::Gpt2(h) => Ok(h.pretrained),
NnHandle::TinyLlama(h) => Ok(h.pretrained),
NnHandle::Llama(_) => Ok(true),
});
methods.add_method(
"forward_shape",
|_, this, (batch, seq): (usize, usize)| match this {
NnHandle::Gpt2(h) => Ok(vec![batch, seq, h.vocab]),
NnHandle::TinyLlama(h) => Ok(vec![batch, seq, h.vocab]),
NnHandle::Llama(h) => Ok(vec![batch, h.vocab]),
},
);
}
}
#[cfg(test)]
mod arch_ops_tests {
use super::*;
use algocline_nn::card::SUPPORTED_ARCHITECTURE_FAMILIES;
#[test]
fn arch_ops_entries_are_all_canonical_families() {
for (name, _) in ARCH_OPS {
assert!(
SUPPORTED_ARCHITECTURE_FAMILIES.contains(name),
"ARCH_OPS entry {name:?} is not in SUPPORTED_ARCHITECTURE_FAMILIES {SUPPORTED_ARCHITECTURE_FAMILIES:?}"
);
}
}
#[test]
fn resolve_arch_ops_matches_bare_and_variant_forms() {
assert!(resolve_arch_ops("gpt2").is_some());
assert!(resolve_arch_ops("gpt2-medium").is_some());
assert!(resolve_arch_ops("tinyllama").is_some());
assert!(resolve_arch_ops("tinyllama-1.1b").is_some());
assert!(resolve_arch_ops("llama").is_some());
assert!(resolve_arch_ops("llama-7b").is_some());
}
#[test]
fn resolve_arch_ops_rejects_prefix_only_matches() {
assert!(resolve_arch_ops("gpt2experimental").is_none());
assert!(resolve_arch_ops("tinyllamafork").is_none());
}
#[test]
fn resolve_arch_ops_rejects_unregistered_families() {
assert!(resolve_arch_ops("qwen2").is_none());
assert!(resolve_arch_ops("phi").is_none());
assert!(resolve_arch_ops("gemma").is_none());
}
#[test]
fn registered_arch_names_reports_gpt2_tinyllama_llama_today() {
let names = registered_arch_names();
assert_eq!(names, vec!["gpt2", "tinyllama", "llama"]);
}
}
#[cfg(test)]
mod nn_handle_helper_tests {
use super::*;
use mlua::Lua;
fn tempdir() -> tempfile::TempDir {
tempfile::tempdir().expect("tempdir")
}
#[test]
fn arch_family_variant_prepends_family_prefix_for_bare_variant() {
let dir = tempdir();
let lua = Lua::new();
let opts = lua.create_table().unwrap();
opts.set("pretrained", false).unwrap();
let gpt2 = build_gpt2_handle("tiny", Some(&opts), dir.path()).expect("build gpt2");
let handle = NnHandle::Gpt2(gpt2);
assert_eq!(handle.arch_family_variant(), "gpt2-tiny");
assert!(!handle.is_lora_wrapped());
}
#[test]
fn arch_family_variant_passes_through_prefixed_variant() {
let dir = tempdir();
let lua = Lua::new();
let opts = lua.create_table().unwrap();
opts.set("pretrained", false).unwrap();
let mut gpt2 = build_gpt2_handle("tiny", Some(&opts), dir.path()).expect("build gpt2");
gpt2.variant = "gpt2-tiny".to_string();
let handle = NnHandle::Gpt2(gpt2);
assert_eq!(handle.arch_family_variant(), "gpt2-tiny");
}
#[test]
fn arch_family_variant_prepends_tinyllama_prefix() {
let dir = tempdir();
let lua = Lua::new();
let opts = lua.create_table().unwrap();
opts.set("pretrained", false).unwrap();
let tll =
build_tinyllama_handle("tinyllama-tiny", Some(&opts), dir.path()).expect("build tll");
let handle = NnHandle::TinyLlama(tll);
assert_eq!(handle.arch_family_variant(), "tinyllama-tiny");
assert!(!handle.is_lora_wrapped());
}
#[test]
fn is_lora_wrapped_returns_false_for_base_handles() {
let dir = tempdir();
let lua = Lua::new();
let opts = lua.create_table().unwrap();
opts.set("pretrained", false).unwrap();
let gpt2 = build_gpt2_handle("tiny", Some(&opts), dir.path()).expect("build gpt2");
assert!(!NnHandle::Gpt2(gpt2).is_lora_wrapped());
let tll =
build_tinyllama_handle("tinyllama-tiny", Some(&opts), dir.path()).expect("build tll");
assert!(!NnHandle::TinyLlama(tll).is_lora_wrapped());
}
#[test]
fn is_lora_wrapped_returns_true_when_has_lora_flag_set() {
let dir = tempdir();
let lua = Lua::new();
let opts = lua.create_table().unwrap();
opts.set("pretrained", false).unwrap();
let mut gpt2 = build_gpt2_handle("tiny", Some(&opts), dir.path()).expect("build gpt2");
gpt2.has_lora = true;
assert!(NnHandle::Gpt2(gpt2).is_lora_wrapped());
let mut tll =
build_tinyllama_handle("tinyllama-tiny", Some(&opts), dir.path()).expect("build tll");
tll.has_lora = true;
assert!(NnHandle::TinyLlama(tll).is_lora_wrapped());
}
}
#[cfg(test)]
mod build_create_payload_from_meta_tests {
use super::*;
use algocline_nn::card::{NnCandleBranch, NnCardMeta, NnLineage};
use serde_json::json;
fn sample_merged_meta() -> NnCardMeta {
NnCardMeta {
name: "my-merged".into(),
backend: "candle".into(),
task: None,
architecture: "gpt2-medium".into(),
training_path: "merged".into(),
lineage: NnLineage {
parent: Some("cards/lora-src-001".into()),
..NnLineage::default()
},
hyperparams: json!({}),
metrics: json!({}),
candle: Some(NnCandleBranch {
bundle_ref: "nn/my-merged-1".into(),
device: None,
dtype: None,
lora: None,
}),
}
}
#[test]
fn envelope_shape_matches_build_create_payload() {
let meta = sample_merged_meta();
let payload = build_create_payload_from_meta("my-merged-1", &meta).expect("build payload");
assert_eq!(payload["pkg"]["name"], NN_PKG);
assert_eq!(payload["card_id"], "my-merged-1");
assert_eq!(payload["metadata"]["kind"], "nn_model");
let nn = &payload["metadata"]["nn"];
assert_eq!(nn["name"], "my-merged");
assert_eq!(nn["architecture"], "gpt2-medium");
assert_eq!(nn["training_path"], "merged");
assert_eq!(nn["lineage"]["parent"], "cards/lora-src-001");
assert_eq!(nn["candle"]["bundle_ref"], "nn/my-merged-1");
}
#[test]
fn refuses_unknown_architecture_family() {
let mut meta = sample_merged_meta();
meta.architecture = "nonexistent-arch".into();
let err = build_create_payload_from_meta("my-merged-1", &meta)
.expect_err("should reject unknown arch");
assert!(
err.to_string().contains("alc.nn.card.merge_lora"),
"expected merge_lora-prefixed error, got: {err}"
);
}
}
#[cfg(test)]
mod trainer_tests {
use super::*;
use mlua::Lua;
fn opts_from(lua: &Lua, pairs: &[(&str, LuaValue)]) -> LuaTable {
let t = lua.create_table().expect("create opts table");
for (k, v) in pairs {
t.set(*k, v.clone()).expect("set opt field");
}
t
}
#[test]
fn full_ft_opts_defaults_when_empty() {
let lua = Lua::new();
let cfg = extract_full_ft_opts(None).expect("None -> defaults");
let default = FullFtConfig::default();
assert_eq!(cfg.lr, default.lr);
assert_eq!(cfg.batch_size, default.batch_size);
assert_eq!(cfg.steps, default.steps);
let empty = lua.create_table().unwrap();
let cfg2 = extract_full_ft_opts(Some(&empty)).expect("empty table -> defaults");
assert_eq!(cfg2.lr, default.lr);
assert_eq!(cfg2.batch_size, default.batch_size);
}
#[test]
fn full_ft_opts_partial_merges_with_defaults() {
let lua = Lua::new();
let opts = opts_from(
&lua,
&[
("lr", LuaValue::Number(1e-3)),
("steps", LuaValue::Integer(42)),
],
);
let cfg = extract_full_ft_opts(Some(&opts)).expect("partial merge");
assert!((cfg.lr - 1e-3).abs() < 1e-12, "lr override");
assert_eq!(cfg.steps, 42, "steps override");
let d = FullFtConfig::default();
assert_eq!(cfg.batch_size, d.batch_size);
assert_eq!(cfg.warmup, d.warmup);
}
#[test]
fn full_ft_opts_reject_zero_batch_size() {
let lua = Lua::new();
let opts = opts_from(&lua, &[("batch_size", LuaValue::Integer(0))]);
let err = extract_full_ft_opts(Some(&opts)).expect_err("zero batch_size");
assert!(
err.to_string().contains("batch_size must be >= 1"),
"message: {err}"
);
}
#[test]
fn schedule_parser_accepts_known_and_rejects_unknown() {
assert!(matches!(
parse_schedule("cosine").unwrap(),
ScheduleKind::CosineWithWarmup
));
assert!(matches!(
parse_schedule("cosine_with_warmup").unwrap(),
ScheduleKind::CosineWithWarmup
));
assert!(matches!(
parse_schedule("constant").unwrap(),
ScheduleKind::Constant
));
let err = parse_schedule("linear").expect_err("unknown");
assert!(err.to_string().contains("linear"), "message: {err}");
}
#[test]
fn lora_cfg_requires_opts_table() {
let err = extract_lora_cfg(None).expect_err("None opts");
assert!(err.to_string().contains("opts table is required"));
}
#[test]
fn lora_cfg_requires_rank_and_alpha() {
let lua = Lua::new();
let no_rank = opts_from(&lua, &[("alpha", LuaValue::Number(16.0))]);
let err = extract_lora_cfg(Some(&no_rank)).expect_err("missing rank");
assert!(err.to_string().contains("opts.rank is required"));
let no_alpha = opts_from(&lua, &[("rank", LuaValue::Integer(8))]);
let err = extract_lora_cfg(Some(&no_alpha)).expect_err("missing alpha");
assert!(err.to_string().contains("opts.alpha is required"));
}
#[test]
fn lora_cfg_rejects_zero_rank_and_nonpositive_alpha() {
let lua = Lua::new();
let zero_rank = opts_from(
&lua,
&[
("rank", LuaValue::Integer(0)),
("alpha", LuaValue::Number(1.0)),
],
);
let err = extract_lora_cfg(Some(&zero_rank)).expect_err("zero rank");
assert!(err.to_string().contains("rank must be >= 1"));
let neg_alpha = opts_from(
&lua,
&[
("rank", LuaValue::Integer(4)),
("alpha", LuaValue::Number(-1.0)),
],
);
let err = extract_lora_cfg(Some(&neg_alpha)).expect_err("negative alpha");
assert!(err.to_string().contains("positive finite number"));
}
#[test]
fn lora_cfg_defaults_target_modules_when_omitted() {
let lua = Lua::new();
let opts = opts_from(
&lua,
&[
("rank", LuaValue::Integer(8)),
("alpha", LuaValue::Number(16.0)),
],
);
let cfg = extract_lora_cfg(Some(&opts)).expect("defaults");
let defaults = LoraConfig::default_targets();
assert_eq!(cfg.target_modules, defaults);
assert_eq!(cfg.rank, 8);
assert!((cfg.alpha - 16.0).abs() < 1e-6);
assert_eq!(cfg.dropout, 0.0);
}
#[test]
fn lora_cfg_rejects_empty_target_modules_list() {
let lua = Lua::new();
let opts = lua.create_table().unwrap();
opts.set("rank", LuaValue::Integer(4)).unwrap();
opts.set("alpha", LuaValue::Number(8.0)).unwrap();
let empty = lua.create_table().unwrap();
opts.set("target_modules", empty).unwrap();
let err = extract_lora_cfg(Some(&opts)).expect_err("empty targets");
assert!(err.to_string().contains("must not be empty"));
}
#[test]
fn lora_cfg_reads_custom_targets_dropout() {
let lua = Lua::new();
let opts = lua.create_table().unwrap();
opts.set("rank", LuaValue::Integer(16)).unwrap();
opts.set("alpha", LuaValue::Number(32.0)).unwrap();
opts.set("dropout", LuaValue::Number(0.05)).unwrap();
let targets = lua.create_table().unwrap();
targets.set(1, "q_proj").unwrap();
targets.set(2, "v_proj").unwrap();
opts.set("target_modules", targets).unwrap();
let cfg = extract_lora_cfg(Some(&opts)).expect("custom");
assert_eq!(
cfg.target_modules,
vec!["q_proj".to_string(), "v_proj".into()]
);
assert!((cfg.dropout - 0.05).abs() < 1e-6);
}
#[test]
fn distill_loss_kind_rejects_wrong_type_input() {
let lua = Lua::new();
let opts = opts_from(&lua, &[("loss_kind", LuaValue::Boolean(true))]);
let err = extract_distill_loss_kind(Some(&opts)).expect_err("wrong-type loss_kind");
let msg = err.to_string();
assert!(!msg.is_empty(), "type-mismatch error should have a message");
}
#[test]
fn resolve_ckpt_dest_rejects_wrong_type_ckpt_dir() {
let lua = Lua::new();
let opts = opts_from(&lua, &[("ckpt_dir", LuaValue::Boolean(false))]);
let tmp = std::env::temp_dir();
let err = resolve_ckpt_dest(Some(&opts), &tmp, "full_ft").expect_err("wrong-type ckpt_dir");
assert!(!err.to_string().is_empty());
}
#[test]
fn distill_loss_kind_defaults_to_ce_and_rejects_unknown() {
assert!(matches!(
extract_distill_loss_kind(None).unwrap(),
DistillLossKind::Ce
));
let lua = Lua::new();
let opts = opts_from(
&lua,
&[(
"loss_kind",
LuaValue::String(lua.create_string("ce").unwrap()),
)],
);
assert!(matches!(
extract_distill_loss_kind(Some(&opts)).unwrap(),
DistillLossKind::Ce
));
let bad = opts_from(
&lua,
&[(
"loss_kind",
LuaValue::String(lua.create_string("kl_soft").unwrap()),
)],
);
let err = extract_distill_loss_kind(Some(&bad)).expect_err("unknown loss");
assert!(err.to_string().contains("kl_soft"));
}
#[test]
fn build_create_payload_populates_lora_branch_when_meta_provides_it() {
let user_meta = json!({
"training_path": "lora",
"architecture": "gpt2-medium",
"candle": {
"device": "cuda:0",
"dtype": "bf16",
"lora": {
"rank": 8,
"alpha": 16,
"base_bundle_ref": "nn/base-gpt2-medium"
}
}
});
let payload =
build_create_payload("card-abc", "my-model", &user_meta).expect("payload with lora");
let lora = payload
.pointer("/metadata/nn/candle/lora")
.expect("lora sub-object");
assert_eq!(lora.get("rank"), Some(&json!(8)));
assert_eq!(lora.get("alpha"), Some(&json!(16)));
assert_eq!(
lora.get("base_bundle_ref"),
Some(&json!("nn/base-gpt2-medium"))
);
}
#[test]
fn build_create_payload_omits_lora_when_meta_absent_or_null() {
let no_candle = json!({
"training_path": "full_ft",
"architecture": "gpt2-medium",
});
let p = build_create_payload("c1", "m", &no_candle).unwrap();
let candle = p.pointer("/metadata/nn/candle").expect("candle present");
assert!(
candle.get("lora").is_none() || candle.get("lora") == Some(&Json::Null),
"lora must be absent: {candle}"
);
let candle_only = json!({
"training_path": "full_ft",
"architecture": "gpt2-medium",
"candle": { "device": "cpu" }
});
let p2 = build_create_payload("c2", "m", &candle_only).unwrap();
let candle2 = p2.pointer("/metadata/nn/candle").expect("candle present");
assert!(candle2.get("lora").is_none() || candle2.get("lora") == Some(&Json::Null));
let explicit_null = json!({
"training_path": "full_ft",
"architecture": "gpt2-medium",
"candle": { "lora": Json::Null }
});
let p3 = build_create_payload("c3", "m", &explicit_null).unwrap();
let candle3 = p3.pointer("/metadata/nn/candle").expect("candle present");
assert!(candle3.get("lora").is_none() || candle3.get("lora") == Some(&Json::Null));
}
#[test]
fn build_create_payload_preserves_lora_target_modules_dropout_delta_path() {
let user_meta = json!({
"training_path": "lora",
"architecture": "gpt2-medium",
"candle": {
"lora": {
"rank": 8,
"alpha": 16,
"base_bundle_ref": "nn/base-gpt2-medium",
"target_modules": ["q_proj", "v_proj"],
"dropout": 0.0625,
"delta_path": "/tmp/ckpt/nn/lora-run.safetensors"
}
}
});
let payload = build_create_payload("card-lora", "m", &user_meta)
.expect("payload with full lora branch");
let lora = payload
.pointer("/metadata/nn/candle/lora")
.expect("lora sub-object");
assert_eq!(
lora.get("target_modules"),
Some(&json!(["q_proj", "v_proj"]))
);
assert_eq!(lora.get("dropout"), Some(&json!(0.0625)));
assert_eq!(
lora.get("delta_path"),
Some(&json!("/tmp/ckpt/nn/lora-run.safetensors"))
);
}
#[test]
fn build_create_payload_defaults_lora_target_modules_when_meta_omits_them() {
let user_meta = json!({
"training_path": "lora",
"architecture": "gpt2-medium",
"candle": {
"lora": {
"rank": 8,
"alpha": 16,
"base_bundle_ref": "nn/base-gpt2-medium"
}
}
});
let payload =
build_create_payload("card-legacy", "m", &user_meta).expect("payload with legacy lora");
let lora = payload
.pointer("/metadata/nn/candle/lora")
.expect("lora sub-object");
let targets = lora.get("target_modules").expect("target_modules");
let arr = targets.as_array().expect("array");
assert_eq!(arr.len(), 6);
assert_eq!(lora.get("dropout"), Some(&json!(0.0)));
assert!(
lora.get("delta_path").is_none() || lora.get("delta_path") == Some(&Json::Null),
"delta_path must be omitted when absent: {lora}"
);
}
#[test]
fn load_gpt2_impl_errors_when_card_missing() {
let tmp = tempfile::TempDir::new().expect("tempdir");
let store = FileCardStore::new(tmp.path().to_path_buf());
let lua = Lua::new();
let placeholder = lua.create_any_userdata(0u32).expect("placeholder userdata");
let msg = match load_gpt2_impl(&store, "does-not-exist", &placeholder) {
Ok(_) => panic!("card missing must error"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("does-not-exist"),
"error must name the missing card: {msg}"
);
}
fn write_test_card(store: &FileCardStore, nn_meta: serde_json::Value) -> String {
let payload = json!({
"pkg": { "name": "alc_nn" },
"metadata": {
"kind": "nn_model",
"nn": nn_meta,
}
});
let (card_id, _path) = store.create(payload).expect("create test card");
card_id
}
#[test]
fn load_gpt2_impl_errors_when_metadata_nn_candle_missing() {
let tmp = tempfile::TempDir::new().expect("tempdir");
let store = FileCardStore::new(tmp.path().to_path_buf());
let card_id = write_test_card(
&store,
json!({
"name": "no-candle",
"backend": "endpoint",
"architecture": "gpt2-medium",
"training_path": "lora",
}),
);
let lua = Lua::new();
let placeholder = lua.create_any_userdata(0u32).expect("placeholder userdata");
let msg = match load_gpt2_impl(&store, &card_id, &placeholder) {
Ok(_) => panic!("missing candle must error"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("missing metadata.nn.candle"),
"error must name the missing candle branch: {msg}"
);
}
#[test]
fn load_gpt2_impl_errors_when_lora_branch_lacks_delta_path() {
let tmp = tempfile::TempDir::new().expect("tempdir");
let store = FileCardStore::new(tmp.path().to_path_buf());
let card_id = write_test_card(
&store,
json!({
"name": "legacy-lora",
"backend": "candle",
"architecture": "gpt2-medium",
"training_path": "lora",
"candle": {
"bundle_ref": "nn/placeholder",
"lora": {
"rank": 8,
"alpha": 16,
"base_bundle_ref": "nn/base-gpt2-medium",
}
}
}),
);
let lua = Lua::new();
let placeholder = lua.create_any_userdata(0u32).expect("placeholder userdata");
let msg = match load_gpt2_impl(&store, &card_id, &placeholder) {
Ok(_) => panic!("missing delta_path must error"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("delta_path"),
"error must name the missing delta_path field: {msg}"
);
}
#[test]
fn load_gpt2_impl_errors_when_delta_path_file_missing() {
let tmp = tempfile::TempDir::new().expect("tempdir");
let store = FileCardStore::new(tmp.path().to_path_buf());
let card_id = write_test_card(
&store,
json!({
"name": "missing-delta-file",
"backend": "candle",
"architecture": "gpt2-medium",
"training_path": "lora",
"candle": {
"bundle_ref": "nn/placeholder",
"lora": {
"rank": 8,
"alpha": 16,
"base_bundle_ref": "nn/base-gpt2-medium",
"target_modules": ["q_proj", "v_proj"],
"dropout": 0.0,
"delta_path": "/nonexistent/path/to/lora-run.safetensors"
}
}
}),
);
let lua = Lua::new();
let placeholder = lua.create_any_userdata(0u32).expect("placeholder userdata");
let msg = match load_gpt2_impl(&store, &card_id, &placeholder) {
Ok(_) => panic!("missing delta file must error"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("delta safetensors missing")
|| msg.contains("/nonexistent/path/to/lora-run.safetensors"),
"error must name the missing delta file: {msg}"
);
}
#[test]
fn build_create_payload_reports_invalid_lora_shape() {
let bad = json!({
"training_path": "lora",
"architecture": "gpt2-medium",
"candle": { "lora": { "rank": 8 } }
});
let err = build_create_payload("cx", "m", &bad).expect_err("invalid lora");
let msg = err.to_string();
assert!(msg.contains("invalid meta.candle.lora"), "message: {msg}");
}
}
#[cfg(test)]
mod load_dispatch_tests {
use super::*;
use algocline_nn::arch::{Gpt2Config, Gpt2Model, LoraConfig};
use candle_nn::{VarBuilder, VarMap};
use mlua::Lua;
use serde_json::json;
fn write_test_card(store: &FileCardStore, nn_meta: serde_json::Value) -> String {
let payload = json!({
"pkg": { "name": "alc_nn" },
"metadata": { "kind": "nn_model", "nn": nn_meta }
});
let (card_id, _path) = store.create(payload).expect("create test card");
card_id
}
#[test]
fn neutral_preset_gpt2_returns_gpt2_nn_handle() {
let tmp = tempfile::TempDir::new().unwrap();
let lua = Lua::new();
let opts_val = lua.to_value(&json!({ "pretrained": false })).unwrap();
let opts_tbl = match opts_val {
LuaValue::Table(t) => t,
_ => unreachable!(),
};
let h = build_neutral_preset("gpt2", "tiny", Some(&opts_tbl), tmp.path())
.expect("neutral preset gpt2 tiny");
assert_eq!(h.arch(), "gpt2");
assert!(h.as_gpt2().is_some());
assert!(h.as_tinyllama().is_none());
}
#[test]
fn neutral_preset_tinyllama_returns_tinyllama_nn_handle() {
let tmp = tempfile::TempDir::new().unwrap();
let lua = Lua::new();
let opts_val = lua.to_value(&json!({ "pretrained": false })).unwrap();
let opts_tbl = match opts_val {
LuaValue::Table(t) => t,
_ => unreachable!(),
};
let h = build_neutral_preset("tinyllama", "tinyllama-tiny", Some(&opts_tbl), tmp.path())
.expect("neutral preset tinyllama-tiny");
assert_eq!(h.arch(), "tinyllama");
assert!(h.as_tinyllama().is_some());
}
#[test]
fn neutral_preset_rejects_unregistered_arch() {
let tmp = tempfile::TempDir::new().unwrap();
let msg = match build_neutral_preset("qwen2", "1.5b", None, tmp.path()) {
Ok(_) => panic!("qwen2 must be rejected"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("qwen2") && msg.contains("not registered"),
"message: {msg}"
);
}
#[test]
fn load_handle_refuses_lora_card_with_directional_error() {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let nn_dir = tmp.path().join("nn");
let card_id = write_test_card(
&store,
json!({
"name": "gpt2-lora",
"backend": "candle",
"architecture": "gpt2-tiny",
"training_path": "lora",
"candle": {
"bundle_ref": "nn/placeholder",
"lora": {
"rank": 4, "alpha": 8, "base_bundle_ref": "nn/base",
"target_modules": ["q_proj"], "dropout": 0.0,
"delta_path": "/nonexistent/lora.safetensors"
}
}
}),
);
let msg = match load_handle_impl(&store, &card_id, &nn_dir) {
Ok(_) => panic!("lora card must be refused by load_handle"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("training_path=\"lora\"") && msg.contains("load_wrap"),
"directional error must point at load_wrap: {msg}"
);
}
#[test]
fn load_handle_refuses_unknown_training_path() {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let nn_dir = tmp.path().join("nn");
let card_id = write_test_card(
&store,
json!({
"name": "gpt2-bogus",
"backend": "candle",
"architecture": "gpt2-tiny",
"training_path": "quantized_awq",
"candle": { "bundle_ref": "nn/placeholder" }
}),
);
let msg = match load_handle_impl(&store, &card_id, &nn_dir) {
Ok(_) => panic!("unknown training_path must error"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("unknown training_path") || msg.contains("quantized_awq"),
"message: {msg}"
);
}
#[test]
fn load_wrap_refuses_merged_card_with_directional_error() {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let card_id = write_test_card(
&store,
json!({
"name": "gpt2-merged",
"backend": "candle",
"architecture": "gpt2-tiny",
"training_path": "merged",
"candle": { "bundle_ref": "nn/placeholder" }
}),
);
let lua = Lua::new();
let placeholder = lua.create_any_userdata(0u32).unwrap();
let msg = match load_wrap_impl(&store, &card_id, &placeholder) {
Ok(_) => panic!("merged card must be refused by load_wrap"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("training_path=\"merged\"") && msg.contains("load_handle"),
"directional error must point at load_handle: {msg}"
);
}
#[test]
fn load_wrap_refuses_full_ft_card_with_directional_error() {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let card_id = write_test_card(
&store,
json!({
"name": "gpt2-fullft",
"backend": "candle",
"architecture": "gpt2-tiny",
"training_path": "full_ft",
"candle": { "bundle_ref": "nn/placeholder" }
}),
);
let lua = Lua::new();
let placeholder = lua.create_any_userdata(0u32).unwrap();
let msg = match load_wrap_impl(&store, &card_id, &placeholder) {
Ok(_) => panic!("full_ft card must be refused by load_wrap"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("training_path=\"full_ft\"") && msg.contains("load_handle"),
"message: {msg}"
);
}
#[test]
fn load_wrap_rejects_arch_mismatched_base_handle() {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let cfg = Gpt2Config::from_variant("tiny").unwrap();
let vm = VarMap::new();
let vs = VarBuilder::from_varmap(&vm, cfg.dtype, &cfg.device);
let mut model = Gpt2Model::new(&cfg, vs).unwrap();
let lora_cfg = LoraConfig::new(4, 8.0);
let lora_vm = model.wrap_lora(&lora_cfg).unwrap();
let delta_path = tmp.path().join("gpt2-lora.safetensors");
lora_vm.save(&delta_path).unwrap();
let card_id = write_test_card(
&store,
json!({
"name": "gpt2-lora",
"backend": "candle",
"architecture": "gpt2-tiny",
"training_path": "lora",
"candle": {
"bundle_ref": "nn/placeholder",
"lora": {
"rank": 4, "alpha": 8, "base_bundle_ref": "nn/base-gpt2",
"target_modules": ["q_proj", "k_proj", "v_proj"],
"dropout": 0.0,
"delta_path": delta_path.to_str().unwrap()
}
}
}),
);
let tll_nn_dir = tmp.path().join("nn");
let lua = Lua::new();
let opts_val = lua.to_value(&json!({ "pretrained": false })).unwrap();
let opts_tbl = match opts_val {
LuaValue::Table(t) => t,
_ => unreachable!(),
};
let tll_handle =
build_tinyllama_handle("tinyllama-tiny", Some(&opts_tbl), &tll_nn_dir).unwrap();
let tll_ud = lua.create_userdata(tll_handle).unwrap();
let msg = match load_wrap_impl(&store, &card_id, &tll_ud) {
Ok(_) => panic!("arch mismatch must be refused"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("gpt2 card requires a gpt2 base handle") && msg.contains("tinyllama"),
"arch mismatch message must name both sides: {msg}"
);
}
#[test]
fn legacy_load_gpt2_shim_delegates_to_shared_core() {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let card_id = write_test_card(
&store,
json!({
"name": "legacy-lora",
"backend": "candle",
"architecture": "gpt2-medium",
"training_path": "lora",
"candle": {
"bundle_ref": "nn/placeholder",
"lora": {
"rank": 8, "alpha": 16,
"base_bundle_ref": "nn/base-gpt2-medium"
}
}
}),
);
let lua = Lua::new();
let placeholder = lua.create_any_userdata(0u32).unwrap();
let msg = match load_gpt2_impl(&store, &card_id, &placeholder) {
Ok(_) => panic!("missing delta_path must error via shim"),
Err(e) => e.to_string(),
};
assert!(msg.contains("delta_path"), "message: {msg}");
}
}
#[cfg(test)]
mod merge_lora_bridge_tests {
use super::*;
use algocline_nn::arch::{Gpt2Config, Gpt2Model, LoraConfig, TinyLlamaConfig, TinyLlamaModel};
use candle_nn::{VarBuilder, VarMap};
use mlua::Lua;
use serde_json::json;
fn write_test_card(store: &FileCardStore, nn_meta: serde_json::Value) -> String {
let payload = json!({
"pkg": { "name": "alc_nn" },
"metadata": { "kind": "nn_model", "nn": nn_meta }
});
let (card_id, _path) = store.create(payload).expect("create test card");
card_id
}
fn opts_table(lua: &Lua, v: serde_json::Value) -> LuaTable {
let val = lua.to_value(&v).expect("to_value");
match val {
LuaValue::Table(t) => t,
_ => unreachable!("json object must serialise to Lua table"),
}
}
fn setup_gpt2_lora_scaffold() -> (tempfile::TempDir, FileCardStore, PathBuf, String, Lua) {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let nn_dir = tmp.path().join("nn");
let cfg = Gpt2Config::from_variant("tiny").unwrap();
let vm = VarMap::new();
let vs = VarBuilder::from_varmap(&vm, cfg.dtype, &cfg.device);
let mut model = Gpt2Model::new(&cfg, vs).unwrap();
let lora_cfg = LoraConfig::new(4, 8.0);
let lora_vm = model.wrap_lora(&lora_cfg).unwrap();
let delta_path = tmp.path().join("gpt2-lora-delta.safetensors");
lora_vm.save(&delta_path).unwrap();
let lora_card_id = write_test_card(
&store,
json!({
"name": "gpt2-lora-src",
"backend": "candle",
"architecture": "gpt2-tiny",
"training_path": "lora",
"candle": {
"bundle_ref": "nn/placeholder",
"lora": {
"rank": 4, "alpha": 8,
"base_bundle_ref": "nn/base-gpt2-tiny",
"target_modules": ["q_proj", "k_proj", "v_proj", "o_proj", "up", "down"],
"dropout": 0.0,
"delta_path": delta_path.to_str().unwrap()
}
}
}),
);
let lua = Lua::new();
(tmp, store, nn_dir, lora_card_id, lua)
}
fn setup_tinyllama_lora_scaffold() -> (tempfile::TempDir, FileCardStore, PathBuf, String, Lua) {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let nn_dir = tmp.path().join("nn");
let cfg = TinyLlamaConfig::from_variant("tinyllama-tiny").unwrap();
let vm = VarMap::new();
let vs = VarBuilder::from_varmap(&vm, cfg.dtype, &cfg.device);
let mut model = TinyLlamaModel::new(&cfg, vs).unwrap();
let lora_cfg = LoraConfig::with_targets(4, 8.0, TinyLlamaModel::default_lora_targets());
let lora_vm = model.wrap_lora(&lora_cfg).unwrap();
let delta_path = tmp.path().join("tinyllama-lora-delta.safetensors");
lora_vm.save(&delta_path).unwrap();
let lora_card_id = write_test_card(
&store,
json!({
"name": "tinyllama-lora-src",
"backend": "candle",
"architecture": "tinyllama-tiny",
"training_path": "lora",
"candle": {
"bundle_ref": "nn/placeholder",
"lora": {
"rank": 4, "alpha": 8,
"base_bundle_ref": "nn/base-tinyllama-tiny",
"target_modules": TinyLlamaModel::default_lora_targets(),
"dropout": 0.0,
"delta_path": delta_path.to_str().unwrap()
}
}
}),
);
let lua = Lua::new();
(tmp, store, nn_dir, lora_card_id, lua)
}
#[test]
fn merge_lora_gpt2_happy_path_produces_merged_card_and_bundle() {
let (_tmp, store, nn_dir, lora_card_id, lua) = setup_gpt2_lora_scaffold();
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let gpt2_base = build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).unwrap();
let gpt2_ud = lua.create_userdata(gpt2_base).unwrap();
let wrapped_nn = load_wrap_impl(&store, &lora_card_id, &gpt2_ud).unwrap();
assert!(wrapped_nn.is_lora_wrapped(), "wrap must set has_lora=true");
let wrapped_ud = lua.create_userdata(wrapped_nn).unwrap();
let merge_opts = opts_table(
&lua,
json!({
"name": "my-merged-gpt2",
"lora_card": lora_card_id.clone(),
}),
);
let merged_card_id =
merge_lora_impl(&store, &nn_dir, &wrapped_ud, merge_opts).expect("merge_lora");
let bundle_path = nn_dir.join(format!("{merged_card_id}.safetensors"));
assert!(
bundle_path.exists(),
"merged safetensors must exist at {bundle_path:?}"
);
let card = store.get(&merged_card_id).unwrap().unwrap();
let nn = card.get("metadata").and_then(|m| m.get("nn")).unwrap();
assert_eq!(nn.get("training_path").unwrap().as_str().unwrap(), "merged");
assert_eq!(
nn.get("architecture").unwrap().as_str().unwrap(),
"gpt2-tiny"
);
assert_eq!(
nn.get("lineage")
.and_then(|l| l.get("parent"))
.and_then(|p| p.as_str())
.unwrap(),
lora_card_id
);
assert_eq!(
nn.get("candle")
.and_then(|c| c.get("bundle_ref"))
.and_then(|b| b.as_str())
.unwrap(),
format!("nn/{merged_card_id}")
);
assert_eq!(nn.get("name").unwrap().as_str().unwrap(), "my-merged-gpt2");
let merged_handle = load_handle_impl(&store, &merged_card_id, &nn_dir).unwrap();
assert_eq!(merged_handle.arch(), "gpt2");
assert!(!merged_handle.is_lora_wrapped());
}
#[test]
fn merge_lora_tinyllama_happy_path_produces_merged_card_and_bundle() {
let (_tmp, store, nn_dir, lora_card_id, lua) = setup_tinyllama_lora_scaffold();
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let tll_base = build_tinyllama_handle("tinyllama-tiny", Some(&base_opts), &nn_dir).unwrap();
let tll_ud = lua.create_userdata(tll_base).unwrap();
let wrapped_nn = load_wrap_impl(&store, &lora_card_id, &tll_ud).unwrap();
assert!(wrapped_nn.is_lora_wrapped());
let wrapped_ud = lua.create_userdata(wrapped_nn).unwrap();
let merge_opts = opts_table(
&lua,
json!({
"name": "my-merged-tinyllama",
"lora_card": lora_card_id.clone(),
}),
);
let merged_card_id =
merge_lora_impl(&store, &nn_dir, &wrapped_ud, merge_opts).expect("merge_lora");
let bundle_path = nn_dir.join(format!("{merged_card_id}.safetensors"));
assert!(bundle_path.exists());
let card = store.get(&merged_card_id).unwrap().unwrap();
let nn = card.get("metadata").and_then(|m| m.get("nn")).unwrap();
assert_eq!(nn.get("training_path").unwrap().as_str().unwrap(), "merged");
assert_eq!(
nn.get("architecture").unwrap().as_str().unwrap(),
"tinyllama-tiny"
);
let merged_handle = load_handle_impl(&store, &merged_card_id, &nn_dir).unwrap();
assert_eq!(merged_handle.arch(), "tinyllama");
assert!(!merged_handle.is_lora_wrapped());
}
#[test]
fn merge_lora_refuses_unwrapped_base_handle_with_directional_error() {
let tmp = tempfile::TempDir::new().unwrap();
let store = FileCardStore::new(tmp.path().join("cards"));
let nn_dir = tmp.path().join("nn");
let lua = Lua::new();
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let gpt2_base = build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).unwrap();
let base_nn = NnHandle::Gpt2(gpt2_base);
let base_ud = lua.create_userdata(base_nn).unwrap();
let merge_opts = opts_table(
&lua,
json!({ "name": "should-fail", "lora_card": "cards/whatever" }),
);
let msg = match merge_lora_impl(&store, &nn_dir, &base_ud, merge_opts) {
Ok(id) => panic!("base handle must be refused; got merged card {id:?}"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("is not LoRA-wrapped") && msg.contains("load_wrap"),
"directional error must mention load_wrap: {msg}"
);
}
#[test]
fn merge_lora_refuses_missing_opts_name() {
let (_tmp, store, nn_dir, lora_card_id, lua) = setup_gpt2_lora_scaffold();
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let gpt2_base = build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).unwrap();
let gpt2_ud = lua.create_userdata(gpt2_base).unwrap();
let wrapped_nn = load_wrap_impl(&store, &lora_card_id, &gpt2_ud).unwrap();
let wrapped_ud = lua.create_userdata(wrapped_nn).unwrap();
let merge_opts = opts_table(&lua, json!({ "lora_card": lora_card_id.clone() }));
let msg = match merge_lora_impl(&store, &nn_dir, &wrapped_ud, merge_opts) {
Ok(_) => panic!("missing name must be refused"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("opts.name must be a non-empty string"),
"message: {msg}"
);
}
#[test]
fn merge_lora_refuses_missing_opts_lora_card() {
let (_tmp, store, nn_dir, lora_card_id, lua) = setup_gpt2_lora_scaffold();
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let gpt2_base = build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).unwrap();
let gpt2_ud = lua.create_userdata(gpt2_base).unwrap();
let wrapped_nn = load_wrap_impl(&store, &lora_card_id, &gpt2_ud).unwrap();
let wrapped_ud = lua.create_userdata(wrapped_nn).unwrap();
let merge_opts = opts_table(&lua, json!({ "name": "no-lora-card" }));
let msg = match merge_lora_impl(&store, &nn_dir, &wrapped_ud, merge_opts) {
Ok(_) => panic!("missing lora_card must be refused"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("opts.lora_card must be a non-empty string"),
"message: {msg}"
);
}
#[test]
fn merge_lora_refuses_empty_string_opts() {
let (_tmp, store, nn_dir, lora_card_id, lua) = setup_gpt2_lora_scaffold();
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let gpt2_base = build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).unwrap();
let gpt2_ud = lua.create_userdata(gpt2_base).unwrap();
let wrapped_nn = load_wrap_impl(&store, &lora_card_id, &gpt2_ud).unwrap();
let wrapped_ud = lua.create_userdata(wrapped_nn).unwrap();
let merge_opts = opts_table(
&lua,
json!({ "name": "", "lora_card": lora_card_id.clone() }),
);
let msg = match merge_lora_impl(&store, &nn_dir, &wrapped_ud, merge_opts) {
Ok(_) => panic!("empty name must be refused"),
Err(e) => e.to_string(),
};
assert!(
msg.contains("opts.name must be a non-empty string"),
"message: {msg}"
);
let merge_opts_2 = opts_table(&lua, json!({ "name": "ok", "lora_card": "" }));
let msg2 = match merge_lora_impl(&store, &nn_dir, &wrapped_ud, merge_opts_2) {
Ok(_) => panic!("empty lora_card must be refused"),
Err(e) => e.to_string(),
};
assert!(
msg2.contains("opts.lora_card must be a non-empty string"),
"message: {msg2}"
);
}
}