use algocline_nn::arch::{LoraConfig, TinyLlamaModel};
use mlua::prelude::*;
use super::nn_card::{
wrap_gpt2_lora_bridge, wrap_tinyllama_lora_bridge, Gpt2Handle, LlamaHandle, NnHandle,
TinyLlamaHandle,
};
pub(super) fn register_nn_wrap(lua: &Lua, alc_table: &LuaTable) -> LuaResult<()> {
let nn_table: LuaTable = alc_table.get("nn")?;
let wrap_lora = lua.create_function(
|_lua, (base, opts): (LuaValue, LuaTable)| -> LuaResult<NnHandle> {
wrap_lora_impl(&base, opts)
},
)?;
nn_table.set("wrap_lora", wrap_lora)?;
Ok(())
}
fn wrap_lora_impl(base: &LuaValue, opts: LuaTable) -> LuaResult<NnHandle> {
let ud = match base {
LuaValue::UserData(u) => u,
_ => {
return Err(LuaError::external(format!(
"alc.nn.wrap_lora: expected NnHandle, got {}",
base.type_name()
)));
}
};
let handle: NnHandle = if let Ok(nn) = ud.borrow::<NnHandle>() {
(*nn).clone()
} else if let Ok(g) = ud.borrow::<Gpt2Handle>() {
NnHandle::Gpt2(g.clone())
} else if let Ok(t) = ud.borrow::<TinyLlamaHandle>() {
NnHandle::TinyLlama(t.clone())
} else if let Ok(l) = ud.borrow::<LlamaHandle>() {
NnHandle::Llama(l.clone())
} else {
return Err(LuaError::external(
"alc.nn.wrap_lora: expected NnHandle, got unknown userdata \
(Gpt2Handle / TinyLlamaHandle / LlamaHandle also accepted)",
));
};
if handle.is_lora_wrapped() {
return Err(LuaError::external(
"alc.nn.wrap_lora: handle is already LoRA-wrapped; drop the wrap or \
start from a base handle",
));
}
if let NnHandle::Llama(_) = handle {
return Err(LuaError::external(format!(
"alc.nn.wrap_lora: architecture {} is not LoRA-wrappable \
(only gpt2 / tinyllama families are supported)",
handle.arch()
)));
}
let arch = handle.arch();
let rank = extract_rank(&opts)?;
let alpha = extract_alpha(&opts)?;
let dropout = extract_dropout(&opts)?;
let target_modules = resolve_and_validate_targets(&opts, arch)?;
let mut cfg = LoraConfig::with_targets(rank, alpha, target_modules);
cfg.dropout = dropout;
match handle {
NnHandle::Gpt2(base_gpt2) => {
let wrapped = wrap_gpt2_lora_bridge(&base_gpt2, &cfg)?;
Ok(NnHandle::Gpt2(wrapped))
}
NnHandle::TinyLlama(base_tll) => {
let wrapped = wrap_tinyllama_lora_bridge(&base_tll, &cfg)?;
Ok(NnHandle::TinyLlama(wrapped))
}
NnHandle::Llama(_) => {
unreachable!("Llama variant guarded above")
}
}
}
fn extract_rank(opts: &LuaTable) -> LuaResult<usize> {
let raw: Option<i64> = opts.get("rank")?;
let n = raw.filter(|v| *v > 0).ok_or_else(|| {
LuaError::external("alc.nn.wrap_lora: opts.rank must be a positive integer")
})?;
Ok(n as usize)
}
fn extract_alpha(opts: &LuaTable) -> LuaResult<f32> {
let raw: Option<f64> = opts.get("alpha")?;
let v = raw.filter(|v| *v > 0.0).ok_or_else(|| {
LuaError::external("alc.nn.wrap_lora: opts.alpha must be a positive number")
})?;
Ok(v as f32)
}
fn extract_dropout(opts: &LuaTable) -> LuaResult<f32> {
let raw: Option<f64> = opts.get("dropout")?;
let v = raw.unwrap_or(0.0);
if !(0.0..1.0).contains(&v) {
return Err(LuaError::external(
"alc.nn.wrap_lora: opts.dropout must be in [0.0, 1.0)",
));
}
Ok(v as f32)
}
fn resolve_and_validate_targets(opts: &LuaTable, arch: &str) -> LuaResult<Vec<String>> {
let known = canonical_targets_for(arch).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.wrap_lora: architecture {arch} is not LoRA-wrappable \
(only gpt2 / tinyllama families are supported)"
))
})?;
let raw: LuaValue = opts.get("target_modules")?;
match raw {
LuaValue::Nil => Ok(known),
LuaValue::Table(tbl) => {
let entries: Vec<String> = tbl
.sequence_values::<String>()
.collect::<LuaResult<Vec<_>>>()?;
if entries.is_empty() {
return Err(LuaError::external(
"alc.nn.wrap_lora: opts.target_modules must be non-empty \
(or nil for the per-arch default)",
));
}
for entry in &entries {
if !known.iter().any(|k| k == entry) {
let known_list = known.join(", ");
return Err(LuaError::external(format!(
"alc.nn.wrap_lora: unknown target module {entry:?} for arch {arch} \
(known: [{known_list}])"
)));
}
}
Ok(entries)
}
other => Err(LuaError::external(format!(
"alc.nn.wrap_lora: opts.target_modules must be an array of strings \
(or nil for the per-arch default); got {}",
other.type_name()
))),
}
}
fn canonical_targets_for(arch: &str) -> Option<Vec<String>> {
match arch {
"gpt2" => Some(LoraConfig::default_targets()),
"tinyllama" => Some(TinyLlamaModel::default_lora_targets()),
_ => None,
}
}
#[cfg(test)]
mod wrap_lora_bridge_tests {
use super::super::nn_card::{build_gpt2_handle, build_tinyllama_handle};
use super::*;
use mlua::Lua;
use serde_json::json;
fn opts_table(lua: &Lua, v: serde_json::Value) -> LuaTable {
use mlua::LuaSerdeExt;
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_base_scaffold() -> (tempfile::TempDir, Gpt2Handle, Lua) {
let tmp = tempfile::TempDir::new().unwrap();
let nn_dir = tmp.path().join("nn");
let lua = Lua::new();
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let base =
build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).expect("build gpt2 tiny base");
(tmp, base, lua)
}
fn setup_tinyllama_base_scaffold() -> (tempfile::TempDir, TinyLlamaHandle, Lua) {
let tmp = tempfile::TempDir::new().unwrap();
let nn_dir = tmp.path().join("nn");
let lua = Lua::new();
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let base = build_tinyllama_handle("tinyllama-tiny", Some(&base_opts), &nn_dir)
.expect("build tinyllama-tiny base");
(tmp, base, lua)
}
fn snapshot_varmap(vm: &candle_nn::VarMap) -> Vec<Vec<f32>> {
vm.all_vars()
.iter()
.map(|v| v.as_tensor().flatten_all().unwrap().to_vec1().unwrap())
.collect()
}
#[test]
fn wrap_lora_gpt2_happy_path_returns_wrapped_handle() {
let (_tmp, base, lua) = setup_gpt2_base_scaffold();
let base_vm = base.varmap().expect("gpt2 tiny base carries a VarMap");
let before = snapshot_varmap(&base_vm);
let base_var_count = base_vm.all_vars().len();
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, json!({ "rank": 4, "alpha": 8.0 }));
let wrapped = wrap_lora_impl(&LuaValue::UserData(base_ud), opts).expect("wrap_lora");
assert!(
wrapped.is_lora_wrapped(),
"wrap_lora must set has_lora=true"
);
assert_eq!(wrapped.arch(), "gpt2");
let after = snapshot_varmap(&base_vm);
assert_eq!(
base_vm.all_vars().len(),
base_var_count,
"base VarMap var count changed during wrap_lora"
);
assert_eq!(before.len(), after.len(), "base VarMap length changed");
for (i, (b, a)) in before.iter().zip(after.iter()).enumerate() {
assert_eq!(
b, a,
"base VarMap tensor #{i} changed during wrap_lora (must stay frozen)"
);
}
}
#[test]
fn wrap_lora_tinyllama_happy_path_returns_wrapped_handle() {
let (_tmp, base, lua) = setup_tinyllama_base_scaffold();
let base_vm = base.varmap().expect("tinyllama tiny base carries a VarMap");
let before = snapshot_varmap(&base_vm);
let base_var_count = base_vm.all_vars().len();
let base_ud = lua.create_userdata(NnHandle::TinyLlama(base)).unwrap();
let opts = opts_table(&lua, json!({ "rank": 4, "alpha": 8.0 }));
let wrapped = wrap_lora_impl(&LuaValue::UserData(base_ud), opts).expect("wrap_lora");
assert!(
wrapped.is_lora_wrapped(),
"wrap_lora must set has_lora=true"
);
assert_eq!(wrapped.arch(), "tinyllama");
let after = snapshot_varmap(&base_vm);
assert_eq!(
base_vm.all_vars().len(),
base_var_count,
"base VarMap var count changed during wrap_lora"
);
assert_eq!(before.len(), after.len(), "base VarMap length changed");
for (i, (b, a)) in before.iter().zip(after.iter()).enumerate() {
assert_eq!(
b, a,
"base VarMap tensor #{i} changed during wrap_lora (must stay frozen)"
);
}
}
fn expect_err(result: LuaResult<NnHandle>) -> String {
match result {
Ok(_) => panic!("expected wrap_lora_impl to fail; got Ok(NnHandle)"),
Err(e) => e.to_string(),
}
}
#[test]
fn wrap_lora_refuses_zero_rank() {
let (_tmp, base, lua) = setup_gpt2_base_scaffold();
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, json!({ "rank": 0, "alpha": 8.0 }));
let msg = expect_err(wrap_lora_impl(&LuaValue::UserData(base_ud), opts));
assert!(
msg.contains("alc.nn.wrap_lora:") && msg.contains("opts.rank"),
"expected rank error, got: {msg}"
);
}
#[test]
fn wrap_lora_refuses_missing_alpha() {
let (_tmp, base, lua) = setup_gpt2_base_scaffold();
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, json!({ "rank": 4 }));
let msg = expect_err(wrap_lora_impl(&LuaValue::UserData(base_ud), opts));
assert!(
msg.contains("alc.nn.wrap_lora:") && msg.contains("opts.alpha"),
"expected alpha error, got: {msg}"
);
}
#[test]
fn wrap_lora_refuses_empty_target_modules_array() {
let (_tmp, base, lua) = setup_gpt2_base_scaffold();
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(
&lua,
json!({ "rank": 4, "alpha": 8.0, "target_modules": [] }),
);
let msg = expect_err(wrap_lora_impl(&LuaValue::UserData(base_ud), opts));
assert!(
msg.contains("alc.nn.wrap_lora:") && msg.contains("opts.target_modules"),
"expected target_modules error, got: {msg}"
);
}
#[test]
fn wrap_lora_refuses_unknown_target_module_per_arch() {
let (_tmp, base, lua) = setup_tinyllama_base_scaffold();
let base_ud = lua.create_userdata(NnHandle::TinyLlama(base)).unwrap();
let opts = opts_table(
&lua,
json!({
"rank": 4,
"alpha": 8.0,
"target_modules": ["q_proj", "k_proj", "v_proj", "o_proj", "up"],
}),
);
let msg = expect_err(wrap_lora_impl(&LuaValue::UserData(base_ud), opts));
assert!(
msg.contains("alc.nn.wrap_lora:")
&& msg.contains("unknown target module")
&& msg.contains("tinyllama"),
"expected per-arch unknown-target error, got: {msg}"
);
}
#[test]
fn wrap_lora_refuses_dropout_out_of_range() {
let (_tmp, base, lua) = setup_gpt2_base_scaffold();
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, json!({ "rank": 4, "alpha": 8.0, "dropout": 1.0 }));
let msg = expect_err(wrap_lora_impl(&LuaValue::UserData(base_ud), opts));
assert!(
msg.contains("alc.nn.wrap_lora:") && msg.contains("opts.dropout"),
"expected dropout error, got: {msg}"
);
}
}