use std::path::PathBuf;
use std::sync::Arc;
use algocline_nn::arch::{LoraConfig, TinyLlamaModel};
use algocline_nn::card::{NnCandleBranch, NnCardMeta, NnLineage, NnLoraBranch};
use algocline_nn::train::{
run_distill, run_full_ft, run_lora_ft, CrossEntropyLoss, DistillLossKind, DistillSpec,
FullFtConfig, TrainError, TrainingLease,
};
use mlua::prelude::*;
use serde_json::json;
use crate::card::FileCardStore;
use super::nn_card::{
build_create_payload_from_meta, compact_epoch_us, extract_full_ft_opts, sanitize_name,
DatasetHandle, Gpt2Handle, LlamaHandle, NnHandle, TinyLlamaHandle,
};
pub(super) fn register_nn_trainer(
lua: &Lua,
alc_table: &LuaTable,
card_store: Arc<FileCardStore>,
nn_dir: PathBuf,
) -> LuaResult<()> {
let nn_table: LuaTable = alc_table.get("nn")?;
let trainer: LuaTable = nn_table.get("trainer")?;
let store = Arc::clone(&card_store);
let dir = nn_dir.clone();
let run_lora_ft = lua.create_function(
move |_lua, (base, dataset, opts): (LuaValue, LuaValue, LuaTable)| -> LuaResult<String> {
run_lora_ft_impl(&store, &dir, &base, &dataset, opts)
},
)?;
trainer.set("run_lora_ft", run_lora_ft)?;
let store_ff = Arc::clone(&card_store);
let dir_ff = nn_dir.clone();
let run_full_ft = lua.create_function(
move |_lua, (base, dataset, opts): (LuaValue, LuaValue, LuaTable)| -> LuaResult<String> {
run_full_ft_impl(&store_ff, &dir_ff, &base, &dataset, opts)
},
)?;
trainer.set("run_full_ft", run_full_ft)?;
let store_rd = Arc::clone(&card_store);
let dir_rd = nn_dir;
let run_distill = lua.create_function(
move |_lua,
(student, dataset, opts): (LuaValue, LuaValue, LuaTable)|
-> LuaResult<String> {
run_distill_impl(&store_rd, &dir_rd, &student, &dataset, opts)
},
)?;
trainer.set("run_distill", run_distill)?;
Ok(())
}
fn run_lora_ft_impl(
store: &FileCardStore,
nn_dir: &std::path::Path,
base: &LuaValue,
dataset: &LuaValue,
opts: LuaTable,
) -> LuaResult<String> {
let base_ud = match base {
LuaValue::UserData(u) => u,
_ => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_lora_ft: expected NnHandle, got {}",
base.type_name()
)));
}
};
let handle: NnHandle = if let Ok(nn) = base_ud.borrow::<NnHandle>() {
(*nn).clone()
} else if let Ok(g) = base_ud.borrow::<Gpt2Handle>() {
NnHandle::Gpt2(g.clone())
} else if let Ok(t) = base_ud.borrow::<TinyLlamaHandle>() {
NnHandle::TinyLlama(t.clone())
} else if let Ok(l) = base_ud.borrow::<LlamaHandle>() {
NnHandle::Llama(l.clone())
} else {
return Err(LuaError::external(
"alc.nn.trainer.run_lora_ft: expected NnHandle, got unknown userdata \
(Gpt2Handle / TinyLlamaHandle / LlamaHandle also accepted)",
));
};
if handle.is_lora_wrapped() {
return Err(LuaError::external(
"alc.nn.trainer.run_lora_ft: expected base (unwrapped) NnHandle; \
drop the wrap first",
));
}
if let NnHandle::Llama(_) = handle {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_lora_ft: architecture {} is not LoRA-wrappable \
(only gpt2 / tinyllama families are supported)",
handle.arch()
)));
}
let dataset_ud = match dataset {
LuaValue::UserData(u) => u,
_ => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_lora_ft: dataset must be an alc.nn.dataset \
(got {})",
dataset.type_name()
)));
}
};
if dataset_ud.borrow::<DatasetHandle>().is_err() {
return Err(LuaError::external(
"alc.nn.trainer.run_lora_ft: dataset must be an alc.nn.dataset \
(got unknown userdata)",
));
}
let arch = handle.arch();
let lora_cfg = extract_lora_cfg(&opts, arch)?;
let train_cfg = extract_train_cfg(&opts)?;
let name: Option<String> = opts.get("name")?;
let name_base = name
.as_deref()
.filter(|s| !s.is_empty())
.unwrap_or("run_lora_ft");
let lora_card_id = format!("{}_{}", sanitize_name(name_base), compact_epoch_us());
let base_bundle_ref = format!("nn/{}", handle.arch_family_variant());
let architecture = handle.arch_family_variant();
let lease = Arc::new(TrainingLease::new());
let ckpt = match &handle {
NnHandle::Gpt2(gpt2) => {
let model_arc = gpt2.model();
let ds_handle = dataset_ud.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_handle.inner_lock()?;
let loss_fn = CrossEntropyLoss::new();
let mut model = model_arc.lock().map_err(|e| {
LuaError::external(format!("alc.nn.trainer.run_lora_ft: model lock: {e}"))
})?;
let result = run_lora_ft(
&mut *model,
ds_lock.as_mut(),
&lora_cfg,
&train_cfg,
&loss_fn,
nn_dir,
&lora_card_id,
Arc::clone(&lease),
);
drop(model);
drop(ds_lock);
drop(ds_handle);
result.map_err(train_err_to_lua)?
}
NnHandle::TinyLlama(tll) => {
let model_arc = tll.model();
let ds_handle = dataset_ud.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_handle.inner_lock()?;
let loss_fn = CrossEntropyLoss::new();
let mut model = model_arc.lock().map_err(|e| {
LuaError::external(format!("alc.nn.trainer.run_lora_ft: model lock: {e}"))
})?;
let result = run_lora_ft(
&mut *model,
ds_lock.as_mut(),
&lora_cfg,
&train_cfg,
&loss_fn,
nn_dir,
&lora_card_id,
Arc::clone(&lease),
);
drop(model);
drop(ds_lock);
drop(ds_handle);
result.map_err(train_err_to_lua)?
}
NnHandle::Llama(_) => {
unreachable!("Llama variant guarded above")
}
};
let delta_path = nn_dir.join("nn").join(&ckpt.bundle_ref);
let delta_path_string = delta_path.to_string_lossy().to_string();
let lora_branch = NnLoraBranch {
rank: lora_cfg.rank as u32,
alpha: lora_cfg.alpha as u32,
base_bundle_ref: base_bundle_ref.clone(),
target_modules: lora_cfg.target_modules.clone(),
dropout: lora_cfg.dropout,
delta_path: Some(delta_path_string.clone()),
};
let candle = NnCandleBranch {
bundle_ref: format!("nn/{lora_card_id}"),
device: None,
dtype: None,
lora: Some(lora_branch),
};
let meta = NnCardMeta {
name: name_base.to_string(),
backend: "candle".into(),
task: None,
architecture,
training_path: "lora".into(),
lineage: NnLineage::default(),
hyperparams: json!({
"lr": train_cfg.lr,
"batch": train_cfg.batch_size,
"steps": train_cfg.steps,
"warmup": train_cfg.warmup,
}),
metrics: json!({
"train_loss": ckpt.train_loss,
"step": ckpt.step,
}),
candle: Some(candle),
};
let payload = build_create_payload_from_meta(&lora_card_id, &meta)?;
let (returned_id, _path) = store
.create(payload)
.map_err(|e| LuaError::external(format!("alc.nn.trainer.run_lora_ft: card store: {e}")))?;
if returned_id != lora_card_id {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_lora_ft: card_id mismatch (expected \
{lora_card_id}, got {returned_id})"
)));
}
Ok(lora_card_id)
}
fn extract_lora_cfg(opts: &LuaTable, arch: &str) -> LuaResult<LoraConfig> {
let rank: Option<i64> = opts.get("rank")?;
let rank = rank.filter(|v| *v > 0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_lora_ft: opts.rank must be a positive integer")
})? as usize;
let alpha: Option<f64> = opts.get("alpha")?;
let alpha = alpha.filter(|v| *v > 0.0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_lora_ft: opts.alpha must be a positive number")
})? as f32;
let dropout: Option<f64> = opts.get("dropout")?;
let dropout = dropout.unwrap_or(0.0);
if !(0.0..1.0).contains(&dropout) {
return Err(LuaError::external(
"alc.nn.trainer.run_lora_ft: opts.dropout must be in [0.0, 1.0)",
));
}
let dropout = dropout as f32;
let known = canonical_targets_for(arch).ok_or_else(|| {
LuaError::external(format!(
"alc.nn.trainer.run_lora_ft: architecture {arch} is not LoRA-wrappable \
(only gpt2 / tinyllama families are supported)"
))
})?;
let raw: LuaValue = opts.get("target_modules")?;
let target_modules: Vec<String> = match raw {
LuaValue::Nil => known.clone(),
LuaValue::Table(tbl) => {
let entries: Vec<String> = tbl
.sequence_values::<String>()
.collect::<LuaResult<Vec<_>>>()?;
if entries.is_empty() {
return Err(LuaError::external(
"alc.nn.trainer.run_lora_ft: 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.trainer.run_lora_ft: unknown target module {entry:?} \
for arch {arch} (known: [{known_list}])"
)));
}
}
entries
}
other => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_lora_ft: opts.target_modules must be an array of strings \
(or nil for the per-arch default); got {}",
other.type_name()
)));
}
};
let mut cfg = LoraConfig::with_targets(rank, alpha, target_modules);
cfg.dropout = dropout;
Ok(cfg)
}
fn extract_train_cfg(opts: &LuaTable) -> LuaResult<FullFtConfig> {
let lr: Option<f64> = opts.get("lr")?;
let lr = lr.filter(|v| v.is_finite() && *v > 0.0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_lora_ft: opts.lr must be a positive number")
})?;
let batch: Option<i64> = opts.get("batch")?;
let batch = batch.filter(|v| *v > 0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_lora_ft: opts.batch must be a positive integer")
})? as usize;
let steps: Option<i64> = opts.get("steps")?;
let steps = steps.filter(|v| *v > 0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_lora_ft: opts.steps must be a positive integer")
})? as usize;
let warmup = match opts.get::<Option<i64>>("warmup")? {
Some(v) if v < 0 => {
return Err(LuaError::external(
"alc.nn.trainer.run_lora_ft: opts.warmup must be >= 0",
));
}
Some(v) => v as usize,
None => 0,
};
let schedule_canonical = match opts
.get::<Option<String>>("schedule")?
.as_deref()
.unwrap_or("CosineWithWarmup")
{
"CosineWithWarmup" => "cosine_with_warmup",
"Constant" => "constant",
other => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_lora_ft: opts.schedule must be one of \
\"CosineWithWarmup\" / \"Constant\" (got {other:?})"
)));
}
};
opts.set("lr", lr)?;
opts.set("batch_size", batch as i64)?;
opts.set("steps", steps as i64)?;
opts.set("warmup", warmup as i64)?;
opts.set("schedule", schedule_canonical.to_string())?;
extract_full_ft_opts(Some(opts))
}
fn canonical_targets_for(arch: &str) -> Option<Vec<String>> {
match arch {
"gpt2" => Some(LoraConfig::default_targets()),
"tinyllama" => Some(TinyLlamaModel::default_lora_targets()),
_ => None,
}
}
fn train_err_to_lua(e: TrainError) -> LuaError {
let msg = match e {
TrainError::ZeroSteps => "alc.nn.trainer.run_lora_ft: zero steps".to_string(),
TrainError::LeaseHeld => {
"alc.nn.trainer.run_lora_ft: training lease already active on this VM".to_string()
}
TrainError::DatasetExhausted { seen, requested } => format!(
"alc.nn.trainer.run_lora_ft: dataset exhausted after {seen} steps \
(requested {requested})"
),
TrainError::Ckpt(inner) => {
format!("alc.nn.trainer.run_lora_ft: checkpoint: {inner}")
}
TrainError::Candle(inner) => {
format!("alc.nn.trainer.run_lora_ft: candle: {inner}")
}
other => format!("alc.nn.trainer.run_lora_ft: {other}"),
};
LuaError::external(msg)
}
fn run_full_ft_impl(
store: &FileCardStore,
nn_dir: &std::path::Path,
base: &LuaValue,
dataset: &LuaValue,
opts: LuaTable,
) -> LuaResult<String> {
let base_ud = match base {
LuaValue::UserData(u) => u,
_ => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_full_ft: expected NnHandle, got {}",
base.type_name()
)));
}
};
let handle: NnHandle = if let Ok(nn) = base_ud.borrow::<NnHandle>() {
(*nn).clone()
} else if let Ok(g) = base_ud.borrow::<Gpt2Handle>() {
NnHandle::Gpt2(g.clone())
} else if let Ok(t) = base_ud.borrow::<TinyLlamaHandle>() {
NnHandle::TinyLlama(t.clone())
} else if let Ok(l) = base_ud.borrow::<LlamaHandle>() {
NnHandle::Llama(l.clone())
} else {
return Err(LuaError::external(
"alc.nn.trainer.run_full_ft: expected NnHandle, got unknown userdata \
(Gpt2Handle / TinyLlamaHandle / LlamaHandle also accepted)",
));
};
if handle.is_lora_wrapped() {
return Err(LuaError::external(
"alc.nn.trainer.run_full_ft: expected base (unwrapped) NnHandle; \
drop the wrap first (a LoRA-wrapped handle is a LoRA-training \
target, not a full-fine-tune target)",
));
}
if let NnHandle::Llama(_) = handle {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_full_ft: architecture {} is not trainable \
(only gpt2 / tinyllama families are supported)",
handle.arch()
)));
}
let dataset_ud = match dataset {
LuaValue::UserData(u) => u,
_ => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_full_ft: dataset must be an alc.nn.dataset \
(got {})",
dataset.type_name()
)));
}
};
if dataset_ud.borrow::<DatasetHandle>().is_err() {
return Err(LuaError::external(
"alc.nn.trainer.run_full_ft: dataset must be an alc.nn.dataset \
(got unknown userdata)",
));
}
let train_cfg = extract_train_cfg_ff(&opts)?;
let name: Option<String> = opts.get("name")?;
let name_base = name
.as_deref()
.filter(|s| !s.is_empty())
.unwrap_or("run_full_ft");
let card_id = format!("{}_{}", sanitize_name(name_base), compact_epoch_us());
let architecture = handle.arch_family_variant();
let lease = Arc::new(TrainingLease::new());
let ckpt = match &handle {
NnHandle::Gpt2(gpt2) => {
let vm_arc = gpt2.varmap().ok_or_else(|| {
LuaError::external(
"alc.nn.trainer.run_full_ft: handle was built with \
pretrained=true; full-fine-tune requires a from-scratch \
handle (pretrained=false)",
)
})?;
let model_arc = gpt2.model();
let ds_handle = dataset_ud.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_handle.inner_lock()?;
let loss_fn = CrossEntropyLoss::new();
let model = model_arc.lock().map_err(|e| {
LuaError::external(format!("alc.nn.trainer.run_full_ft: model lock: {e}"))
})?;
let result = run_full_ft(
&*model,
&vm_arc,
ds_lock.as_mut(),
&train_cfg,
&loss_fn,
nn_dir,
&card_id,
Arc::clone(&lease),
);
drop(model);
drop(ds_lock);
drop(ds_handle);
result.map_err(train_err_to_lua_ff)?
}
NnHandle::TinyLlama(tll) => {
let vm_arc = tll.varmap().ok_or_else(|| {
LuaError::external(
"alc.nn.trainer.run_full_ft: handle was built with \
pretrained=true; full-fine-tune requires a from-scratch \
handle (pretrained=false)",
)
})?;
let model_arc = tll.model();
let ds_handle = dataset_ud.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_handle.inner_lock()?;
let loss_fn = CrossEntropyLoss::new();
let model = model_arc.lock().map_err(|e| {
LuaError::external(format!("alc.nn.trainer.run_full_ft: model lock: {e}"))
})?;
let result = run_full_ft(
&*model,
&vm_arc,
ds_lock.as_mut(),
&train_cfg,
&loss_fn,
nn_dir,
&card_id,
Arc::clone(&lease),
);
drop(model);
drop(ds_lock);
drop(ds_handle);
result.map_err(train_err_to_lua_ff)?
}
NnHandle::Llama(_) => {
unreachable!("Llama variant guarded above")
}
};
let candle = NnCandleBranch {
bundle_ref: format!("nn/{card_id}"),
device: None,
dtype: None,
lora: None,
};
let meta = NnCardMeta {
name: name_base.to_string(),
backend: "candle".into(),
task: None,
architecture,
training_path: "full_ft".into(),
lineage: NnLineage::default(),
hyperparams: json!({
"lr": train_cfg.lr,
"batch": train_cfg.batch_size,
"steps": train_cfg.steps,
"warmup": train_cfg.warmup,
}),
metrics: json!({
"train_loss": ckpt.train_loss,
"step": ckpt.step,
}),
candle: Some(candle),
};
let payload = build_create_payload_from_meta(&card_id, &meta)?;
let (returned_id, _path) = store
.create(payload)
.map_err(|e| LuaError::external(format!("alc.nn.trainer.run_full_ft: card store: {e}")))?;
if returned_id != card_id {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_full_ft: card_id mismatch (expected \
{card_id}, got {returned_id})"
)));
}
let _ = ckpt.bundle_ref;
Ok(card_id)
}
fn extract_train_cfg_ff(opts: &LuaTable) -> LuaResult<FullFtConfig> {
let lr: Option<f64> = opts.get("lr")?;
let lr = lr.filter(|v| v.is_finite() && *v > 0.0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_full_ft: opts.lr must be a positive number")
})?;
let batch: Option<i64> = opts.get("batch")?;
let batch = batch.filter(|v| *v > 0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_full_ft: opts.batch must be a positive integer")
})? as usize;
let steps: Option<i64> = opts.get("steps")?;
let steps = steps.filter(|v| *v > 0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_full_ft: opts.steps must be a positive integer")
})? as usize;
let warmup = match opts.get::<Option<i64>>("warmup")? {
Some(v) if v < 0 => {
return Err(LuaError::external(
"alc.nn.trainer.run_full_ft: opts.warmup must be >= 0",
));
}
Some(v) => v as usize,
None => 0,
};
let schedule_canonical = match opts
.get::<Option<String>>("schedule")?
.as_deref()
.unwrap_or("CosineWithWarmup")
{
"CosineWithWarmup" => "cosine_with_warmup",
"Constant" => "constant",
other => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_full_ft: opts.schedule must be one of \
\"CosineWithWarmup\" / \"Constant\" (got {other:?})"
)));
}
};
opts.set("lr", lr)?;
opts.set("batch_size", batch as i64)?;
opts.set("steps", steps as i64)?;
opts.set("warmup", warmup as i64)?;
opts.set("schedule", schedule_canonical.to_string())?;
extract_full_ft_opts(Some(opts))
}
fn train_err_to_lua_ff(e: TrainError) -> LuaError {
let msg = match e {
TrainError::ZeroSteps => "alc.nn.trainer.run_full_ft: zero steps".to_string(),
TrainError::LeaseHeld => {
"alc.nn.trainer.run_full_ft: training lease already active on this VM".to_string()
}
TrainError::DatasetExhausted { seen, requested } => format!(
"alc.nn.trainer.run_full_ft: dataset exhausted after {seen} steps \
(requested {requested})"
),
TrainError::Ckpt(inner) => {
format!("alc.nn.trainer.run_full_ft: checkpoint: {inner}")
}
TrainError::Candle(inner) => {
format!("alc.nn.trainer.run_full_ft: candle: {inner}")
}
other => format!("alc.nn.trainer.run_full_ft: {other}"),
};
LuaError::external(msg)
}
fn run_distill_impl(
store: &FileCardStore,
nn_dir: &std::path::Path,
student: &LuaValue,
dataset: &LuaValue,
opts: LuaTable,
) -> LuaResult<String> {
let student_ud = match student {
LuaValue::UserData(u) => u,
_ => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_distill: expected NnHandle, got {}",
student.type_name()
)));
}
};
let handle: NnHandle = if let Ok(nn) = student_ud.borrow::<NnHandle>() {
(*nn).clone()
} else if let Ok(g) = student_ud.borrow::<Gpt2Handle>() {
NnHandle::Gpt2(g.clone())
} else if let Ok(t) = student_ud.borrow::<TinyLlamaHandle>() {
NnHandle::TinyLlama(t.clone())
} else if let Ok(l) = student_ud.borrow::<LlamaHandle>() {
NnHandle::Llama(l.clone())
} else {
return Err(LuaError::external(
"alc.nn.trainer.run_distill: expected NnHandle, got unknown userdata \
(Gpt2Handle / TinyLlamaHandle / LlamaHandle also accepted)",
));
};
if handle.is_lora_wrapped() {
return Err(LuaError::external(
"alc.nn.trainer.run_distill: expected base (unwrapped) NnHandle; \
drop the wrap first (a LoRA-wrapped handle is a LoRA-training \
target, not a distillation student)",
));
}
if let NnHandle::Llama(_) = handle {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_distill: architecture {} is not trainable \
(only gpt2 / tinyllama families are supported)",
handle.arch()
)));
}
let dataset_ud = match dataset {
LuaValue::UserData(u) => u,
_ => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_distill: dataset must be an alc.nn.dataset \
(got {})",
dataset.type_name()
)));
}
};
if dataset_ud.borrow::<DatasetHandle>().is_err() {
return Err(LuaError::external(
"alc.nn.trainer.run_distill: dataset must be an alc.nn.dataset \
(got unknown userdata)",
));
}
let train_cfg = extract_train_cfg_rd(&opts)?;
let loss_kind = extract_distill_loss_kind_rd(&opts)?;
let spec = DistillSpec {
hyperparams: train_cfg,
loss_kind,
};
let name: Option<String> = opts.get("name")?;
let name_base = name
.as_deref()
.filter(|s| !s.is_empty())
.unwrap_or("run_distill");
let card_id = format!("{}_{}", sanitize_name(name_base), compact_epoch_us());
let architecture = handle.arch_family_variant();
let lease = Arc::new(TrainingLease::new());
let ckpt = match &handle {
NnHandle::Gpt2(gpt2) => {
let vm_arc = gpt2.varmap().ok_or_else(|| {
LuaError::external(
"alc.nn.trainer.run_distill: handle was built with \
pretrained=true; distillation requires a from-scratch \
student handle (pretrained=false)",
)
})?;
let model_arc = gpt2.model();
let ds_handle = dataset_ud.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_handle.inner_lock()?;
let model = model_arc.lock().map_err(|e| {
LuaError::external(format!("alc.nn.trainer.run_distill: model lock: {e}"))
})?;
let result = run_distill(
&*model,
&vm_arc,
ds_lock.as_mut(),
&spec,
nn_dir,
&card_id,
Arc::clone(&lease),
);
drop(model);
drop(ds_lock);
drop(ds_handle);
result.map_err(train_err_to_lua_rd)?
}
NnHandle::TinyLlama(tll) => {
let vm_arc = tll.varmap().ok_or_else(|| {
LuaError::external(
"alc.nn.trainer.run_distill: handle was built with \
pretrained=true; distillation requires a from-scratch \
student handle (pretrained=false)",
)
})?;
let model_arc = tll.model();
let ds_handle = dataset_ud.borrow_mut::<DatasetHandle>()?;
let mut ds_lock = ds_handle.inner_lock()?;
let model = model_arc.lock().map_err(|e| {
LuaError::external(format!("alc.nn.trainer.run_distill: model lock: {e}"))
})?;
let result = run_distill(
&*model,
&vm_arc,
ds_lock.as_mut(),
&spec,
nn_dir,
&card_id,
Arc::clone(&lease),
);
drop(model);
drop(ds_lock);
drop(ds_handle);
result.map_err(train_err_to_lua_rd)?
}
NnHandle::Llama(_) => {
unreachable!("Llama variant guarded above")
}
};
let candle = NnCandleBranch {
bundle_ref: format!("nn/{card_id}"),
device: None,
dtype: None,
lora: None,
};
let loss_kind_str = match spec.loss_kind {
DistillLossKind::Ce => "ce",
};
let meta = NnCardMeta {
name: name_base.to_string(),
backend: "candle".into(),
task: None,
architecture,
training_path: "distillation".into(),
lineage: NnLineage::default(),
hyperparams: json!({
"lr": spec.hyperparams.lr,
"batch": spec.hyperparams.batch_size,
"steps": spec.hyperparams.steps,
"warmup": spec.hyperparams.warmup,
"loss_kind": loss_kind_str,
}),
metrics: json!({
"train_loss": ckpt.train_loss,
"step": ckpt.step,
}),
candle: Some(candle),
};
let payload = build_create_payload_from_meta(&card_id, &meta)?;
let (returned_id, _path) = store
.create(payload)
.map_err(|e| LuaError::external(format!("alc.nn.trainer.run_distill: card store: {e}")))?;
if returned_id != card_id {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_distill: card_id mismatch (expected \
{card_id}, got {returned_id})"
)));
}
let _ = ckpt.bundle_ref;
Ok(card_id)
}
fn extract_train_cfg_rd(opts: &LuaTable) -> LuaResult<FullFtConfig> {
let lr: Option<f64> = opts.get("lr")?;
let lr = lr.filter(|v| v.is_finite() && *v > 0.0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_distill: opts.lr must be a positive number")
})?;
let batch: Option<i64> = opts.get("batch")?;
let batch = batch.filter(|v| *v > 0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_distill: opts.batch must be a positive integer")
})? as usize;
let steps: Option<i64> = opts.get("steps")?;
let steps = steps.filter(|v| *v > 0).ok_or_else(|| {
LuaError::external("alc.nn.trainer.run_distill: opts.steps must be a positive integer")
})? as usize;
let warmup = match opts.get::<Option<i64>>("warmup")? {
Some(v) if v < 0 => {
return Err(LuaError::external(
"alc.nn.trainer.run_distill: opts.warmup must be >= 0",
));
}
Some(v) => v as usize,
None => 0,
};
let schedule_canonical = match opts
.get::<Option<String>>("schedule")?
.as_deref()
.unwrap_or("CosineWithWarmup")
{
"CosineWithWarmup" => "cosine_with_warmup",
"Constant" => "constant",
other => {
return Err(LuaError::external(format!(
"alc.nn.trainer.run_distill: opts.schedule must be one of \
\"CosineWithWarmup\" / \"Constant\" (got {other:?})"
)));
}
};
opts.set("lr", lr)?;
opts.set("batch_size", batch as i64)?;
opts.set("steps", steps as i64)?;
opts.set("warmup", warmup as i64)?;
opts.set("schedule", schedule_canonical.to_string())?;
extract_full_ft_opts(Some(opts))
}
fn extract_distill_loss_kind_rd(opts: &LuaTable) -> LuaResult<DistillLossKind> {
let raw = opts
.get::<Option<String>>("loss_kind")?
.unwrap_or_else(|| "ce".to_string());
match raw.as_str() {
"ce" => Ok(DistillLossKind::Ce),
other => Err(LuaError::external(format!(
"alc.nn.trainer.run_distill: unknown loss_kind '{other}' (expected 'ce')"
))),
}
}
fn train_err_to_lua_rd(e: TrainError) -> LuaError {
let msg = match e {
TrainError::ZeroSteps => "alc.nn.trainer.run_distill: zero steps".to_string(),
TrainError::LeaseHeld => {
"alc.nn.trainer.run_distill: training lease already active on this VM".to_string()
}
TrainError::DatasetExhausted { seen, requested } => format!(
"alc.nn.trainer.run_distill: dataset exhausted after {seen} steps \
(requested {requested})"
),
TrainError::Ckpt(inner) => {
format!("alc.nn.trainer.run_distill: checkpoint: {inner}")
}
TrainError::Candle(inner) => {
format!("alc.nn.trainer.run_distill: candle: {inner}")
}
other => format!("alc.nn.trainer.run_distill: {other}"),
};
LuaError::external(msg)
}
#[cfg(test)]
mod run_ft_bridge_tests {
use super::super::nn_card::{build_gpt2_handle, build_tinyllama_handle, load_wrap_impl};
use super::*;
use algocline_nn::train::{DatasetOpts, TokenizedDataset};
use candle_nn::VarMap;
use mlua::Lua;
use serde_json::json;
fn overfit_row() -> Vec<u32> {
vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16]
}
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 make_dataset_handle(lua: &Lua, row: Vec<u32>, n: usize) -> LuaAnyUserData {
let rows: Vec<Vec<u32>> = std::iter::repeat_with(|| row.clone()).take(n).collect();
let dopts = DatasetOpts {
batch_size: 1,
ctx_len: 16,
shuffle: false,
pad_id: 0,
text_field: "text".into(),
};
let ds = TokenizedDataset::new(rows, dopts);
let handle = DatasetHandle::for_test(Box::new(ds), "test-synthetic".into(), 1, 16);
lua.create_userdata(handle).expect("dataset userdata")
}
fn setup_gpt2_scaffold() -> (tempfile::TempDir, FileCardStore, PathBuf, Gpt2Handle, Lua) {
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 base =
build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).expect("build gpt2 tiny base");
(tmp, store, nn_dir, base, lua)
}
fn setup_tinyllama_scaffold() -> (
tempfile::TempDir,
FileCardStore,
PathBuf,
TinyLlamaHandle,
Lua,
) {
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 base = build_tinyllama_handle("tinyllama-tiny", Some(&base_opts), &nn_dir)
.expect("build tinyllama-tiny base");
(tmp, store, nn_dir, base, lua)
}
fn snapshot_varmap(vm: &VarMap) -> Vec<Vec<f32>> {
vm.all_vars()
.iter()
.map(|v| v.as_tensor().flatten_all().unwrap().to_vec1().unwrap())
.collect()
}
fn base_train_opts() -> serde_json::Value {
json!({
"rank": 4,
"alpha": 8.0,
"lr": 5e-3,
"batch": 1,
"steps": 3,
"warmup": 0,
"schedule": "CosineWithWarmup",
})
}
#[test]
fn run_lora_ft_gpt2_happy_path_writes_lora_card_and_delta() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, base_train_opts());
let card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_lora_ft");
let delta_path = nn_dir
.join("nn")
.join(format!("lora-{card_id}.safetensors"));
assert!(
delta_path.exists(),
"Δ safetensors must exist at {delta_path:?}"
);
let card = store.get(&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(), "lora");
assert_eq!(
nn.get("architecture").unwrap().as_str().unwrap(),
"gpt2-tiny"
);
assert_eq!(
nn.get("candle")
.and_then(|c| c.get("bundle_ref"))
.and_then(|b| b.as_str())
.unwrap(),
format!("nn/{card_id}")
);
let lora = nn.get("candle").and_then(|c| c.get("lora")).unwrap();
assert_eq!(lora.get("rank").unwrap().as_u64().unwrap(), 4);
assert_eq!(lora.get("alpha").unwrap().as_u64().unwrap(), 8);
assert_eq!(
lora.get("base_bundle_ref").unwrap().as_str().unwrap(),
"nn/gpt2-tiny"
);
assert_eq!(
lora.get("delta_path").unwrap().as_str().unwrap(),
delta_path.to_string_lossy()
);
}
#[test]
fn run_lora_ft_tinyllama_happy_path_writes_lora_card_and_delta() {
let (_tmp, store, nn_dir, base, lua) = setup_tinyllama_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::TinyLlama(base)).unwrap();
let opts = opts_table(&lua, base_train_opts());
let card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_lora_ft");
let delta_path = nn_dir
.join("nn")
.join(format!("lora-{card_id}.safetensors"));
assert!(delta_path.exists());
let card = store.get(&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(), "lora");
assert_eq!(
nn.get("architecture").unwrap().as_str().unwrap(),
"tinyllama-tiny"
);
let lora = nn.get("candle").and_then(|c| c.get("lora")).unwrap();
assert_eq!(
lora.get("base_bundle_ref").unwrap().as_str().unwrap(),
"nn/tinyllama-tiny"
);
}
fn expect_err(result: LuaResult<String>) -> String {
match result {
Ok(id) => panic!("expected run_lora_ft_impl to fail; got card_id {id:?}"),
Err(e) => e.to_string(),
}
}
#[test]
fn run_lora_ft_refuses_zero_steps() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 5);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let mut o = base_train_opts();
o["steps"] = json!(0);
let opts = opts_table(&lua, o);
let msg = expect_err(run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_lora_ft:") && msg.contains("opts.steps"),
"expected steps error, got: {msg}"
);
}
#[test]
fn run_lora_ft_refuses_unknown_schedule() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 5);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let mut o = base_train_opts();
o["schedule"] = json!("Adam");
let opts = opts_table(&lua, o);
let msg = expect_err(run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_lora_ft:")
&& msg.contains("opts.schedule")
&& msg.contains("Adam"),
"expected schedule error, got: {msg}"
);
}
#[test]
fn run_lora_ft_leaves_base_vars_bit_identical() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let base_vm = base.varmap().expect("from-scratch base carries VarMap");
let before = snapshot_varmap(&base_vm);
let base_var_count = base_vm.all_vars().len();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, base_train_opts());
let _card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_lora_ft");
let after = snapshot_varmap(&base_vm);
assert_eq!(
base_vm.all_vars().len(),
base_var_count,
"base VarMap var count changed"
);
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} drifted through run_lora_ft (must stay frozen)"
);
}
}
#[test]
fn run_lora_ft_produces_delta_of_expected_var_count() {
let (_tmp, store, nn_dir, base, lua) = setup_tinyllama_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::TinyLlama(base)).unwrap();
let opts = opts_table(&lua, base_train_opts());
let card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_lora_ft");
let delta_path = nn_dir
.join("nn")
.join(format!("lora-{card_id}.safetensors"));
let tensors = candle_core::safetensors::load(&delta_path, &candle_core::Device::Cpu)
.expect("delta safetensors must load");
assert_eq!(
tensors.len(),
28,
"delta safetensors must contain exactly 28 tensors (2 layers × 7 targets × 2), \
got {}: keys={:?}",
tensors.len(),
tensors.keys().collect::<Vec<_>>()
);
for name in tensors.keys() {
assert!(
name.ends_with(".lora_a.weight") || name.ends_with(".lora_b.weight"),
"unexpected non-LoRA key in delta bundle: {name}"
);
}
}
#[test]
fn run_lora_ft_wraps_have_lora_true() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, base_train_opts());
let card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_lora_ft");
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let fresh_base =
build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).expect("fresh gpt2 base");
let fresh_ud = lua.create_userdata(fresh_base).unwrap();
let wrapped = load_wrap_impl(&store, &card_id, &fresh_ud).expect("load_wrap");
assert!(
wrapped.is_lora_wrapped(),
"load_wrap of run_lora_ft output must set has_lora=true"
);
}
#[test]
fn run_lora_ft_card_roundtrips_through_load_wrap_gpt2() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, base_train_opts());
let card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_lora_ft");
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let fresh_base =
build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).expect("fresh gpt2 base");
let fresh_ud = lua.create_userdata(fresh_base).unwrap();
let wrapped = load_wrap_impl(&store, &card_id, &fresh_ud).expect("load_wrap");
assert_eq!(wrapped.arch(), "gpt2");
assert!(wrapped.is_lora_wrapped());
}
#[test]
fn run_lora_ft_card_roundtrips_through_load_wrap_tinyllama() {
let (_tmp, store, nn_dir, base, lua) = setup_tinyllama_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::TinyLlama(base)).unwrap();
let opts = opts_table(&lua, base_train_opts());
let card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_lora_ft");
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let fresh_base = build_tinyllama_handle("tinyllama-tiny", Some(&base_opts), &nn_dir)
.expect("fresh tinyllama base");
let fresh_ud = lua.create_userdata(fresh_base).unwrap();
let wrapped = load_wrap_impl(&store, &card_id, &fresh_ud).expect("load_wrap");
assert_eq!(wrapped.arch(), "tinyllama");
assert!(wrapped.is_lora_wrapped());
}
fn base_full_ft_opts() -> serde_json::Value {
json!({
"lr": 5e-3,
"batch": 1,
"steps": 3,
"warmup": 0,
"schedule": "CosineWithWarmup",
})
}
#[test]
fn run_full_ft_gpt2_happy_path_writes_full_ft_card() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, base_full_ft_opts());
let card_id = run_full_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_full_ft");
let ckpt_path = nn_dir.join(format!("{card_id}.safetensors"));
assert!(
ckpt_path.exists(),
"full-ft safetensors must exist at {ckpt_path:?}"
);
let card = store.get(&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(),
"full_ft"
);
assert_eq!(
nn.get("architecture").unwrap().as_str().unwrap(),
"gpt2-tiny"
);
let candle = nn.get("candle").unwrap();
assert_eq!(
candle.get("bundle_ref").unwrap().as_str().unwrap(),
format!("nn/{card_id}")
);
let lora = candle.get("lora");
assert!(
lora.is_none() || lora.unwrap().is_null(),
"full-ft Card must not carry a LoRA branch; got: {lora:?}"
);
}
#[test]
fn run_full_ft_tinyllama_happy_path_writes_full_ft_card() {
let (_tmp, store, nn_dir, base, lua) = setup_tinyllama_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::TinyLlama(base)).unwrap();
let opts = opts_table(&lua, base_full_ft_opts());
let card_id = run_full_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_full_ft");
let ckpt_path = nn_dir.join(format!("{card_id}.safetensors"));
assert!(ckpt_path.exists());
let card = store.get(&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(),
"full_ft"
);
assert_eq!(
nn.get("architecture").unwrap().as_str().unwrap(),
"tinyllama-tiny"
);
}
#[test]
fn run_full_ft_refuses_zero_steps() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 5);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let mut o = base_full_ft_opts();
o["steps"] = json!(0);
let opts = opts_table(&lua, o);
let msg = expect_err(run_full_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_full_ft:") && msg.contains("opts.steps"),
"expected steps error, got: {msg}"
);
}
#[test]
fn run_full_ft_refuses_lora_wrapped_handle() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let lora_opts = opts_table(&lua, base_train_opts());
let lora_card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
lora_opts,
)
.expect("run_lora_ft");
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let fresh_base =
build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).expect("fresh gpt2 base");
let fresh_ud = lua.create_userdata(fresh_base).unwrap();
let wrapped = load_wrap_impl(&store, &lora_card_id, &fresh_ud).expect("load_wrap");
assert!(wrapped.is_lora_wrapped());
let wrapped_ud = lua.create_userdata(wrapped).unwrap();
let ds_ud2 = make_dataset_handle(&lua, overfit_row(), 5);
let ff_opts = opts_table(&lua, base_full_ft_opts());
let msg = expect_err(run_full_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(wrapped_ud),
&LuaValue::UserData(ds_ud2),
ff_opts,
));
assert!(
msg.contains("alc.nn.trainer.run_full_ft:") && msg.contains("drop the wrap first"),
"expected wrapped-handle refusal, got: {msg}"
);
}
#[test]
fn run_full_ft_refuses_pretrained_handle() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let pretrained_like = base.for_test_pretrained_like();
let base_ud = lua
.create_userdata(NnHandle::Gpt2(pretrained_like))
.unwrap();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 5);
let opts = opts_table(&lua, base_full_ft_opts());
let msg = expect_err(run_full_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_full_ft:") && msg.contains("pretrained=true"),
"expected pretrained refusal, got: {msg}"
);
}
#[test]
fn run_distill_gpt2_happy_path_writes_distillation_card() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let opts = opts_table(&lua, base_full_ft_opts());
let card_id = run_distill_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_distill");
let ckpt_path = nn_dir.join(format!("{card_id}.safetensors"));
assert!(
ckpt_path.exists(),
"distill safetensors must exist at {ckpt_path:?}"
);
let card = store.get(&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(),
"distillation"
);
assert_eq!(
nn.get("architecture").unwrap().as_str().unwrap(),
"gpt2-tiny"
);
assert_eq!(
nn.get("hyperparams")
.and_then(|h| h.get("loss_kind"))
.and_then(|l| l.as_str())
.unwrap(),
"ce"
);
let candle = nn.get("candle").unwrap();
assert_eq!(
candle.get("bundle_ref").unwrap().as_str().unwrap(),
format!("nn/{card_id}")
);
let lora = candle.get("lora");
assert!(
lora.is_none() || lora.unwrap().is_null(),
"distillation Card must not carry a LoRA branch; got: {lora:?}"
);
}
#[test]
fn run_distill_tinyllama_happy_path_writes_distillation_card() {
let (_tmp, store, nn_dir, base, lua) = setup_tinyllama_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::TinyLlama(base)).unwrap();
let opts = opts_table(&lua, base_full_ft_opts());
let card_id = run_distill_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
)
.expect("run_distill");
let ckpt_path = nn_dir.join(format!("{card_id}.safetensors"));
assert!(ckpt_path.exists());
let card = store.get(&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(),
"distillation"
);
assert_eq!(
nn.get("architecture").unwrap().as_str().unwrap(),
"tinyllama-tiny"
);
}
#[test]
fn run_distill_refuses_unknown_loss_kind() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 5);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let mut o = base_full_ft_opts();
o["loss_kind"] = json!("kl");
let opts = opts_table(&lua, o);
let msg = expect_err(run_distill_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_distill:")
&& msg.contains("loss_kind")
&& msg.contains("kl"),
"expected loss_kind error, got: {msg}"
);
}
#[test]
fn run_distill_refuses_pretrained_handle() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let pretrained_like = base.for_test_pretrained_like();
let base_ud = lua
.create_userdata(NnHandle::Gpt2(pretrained_like))
.unwrap();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 5);
let opts = opts_table(&lua, base_full_ft_opts());
let msg = expect_err(run_distill_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_distill:") && msg.contains("pretrained=true"),
"expected pretrained refusal, got: {msg}"
);
}
#[test]
fn run_distill_refuses_lora_wrapped_handle() {
let (_tmp, store, nn_dir, base, lua) = setup_gpt2_scaffold();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 20);
let base_ud = lua.create_userdata(NnHandle::Gpt2(base)).unwrap();
let lora_opts = opts_table(&lua, base_train_opts());
let lora_card_id = run_lora_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
lora_opts,
)
.expect("run_lora_ft");
let base_opts = opts_table(&lua, json!({ "pretrained": false }));
let fresh_base =
build_gpt2_handle("tiny", Some(&base_opts), &nn_dir).expect("fresh gpt2 base");
let fresh_ud = lua.create_userdata(fresh_base).unwrap();
let wrapped = load_wrap_impl(&store, &lora_card_id, &fresh_ud).expect("load_wrap");
assert!(wrapped.is_lora_wrapped());
let wrapped_ud = lua.create_userdata(wrapped).unwrap();
let ds_ud2 = make_dataset_handle(&lua, overfit_row(), 5);
let opts = opts_table(&lua, base_full_ft_opts());
let msg = expect_err(run_distill_impl(
&store,
&nn_dir,
&LuaValue::UserData(wrapped_ud),
&LuaValue::UserData(ds_ud2),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_distill:") && msg.contains("drop the wrap first"),
"expected wrapped-handle refusal, got: {msg}"
);
}
#[test]
fn run_full_ft_refuses_pretrained_handle_tinyllama() {
let (_tmp, store, nn_dir, base, lua) = setup_tinyllama_scaffold();
let pretrained_like = base.for_test_pretrained_like();
let base_ud = lua
.create_userdata(NnHandle::TinyLlama(pretrained_like))
.unwrap();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 5);
let opts = opts_table(&lua, base_full_ft_opts());
let msg = expect_err(run_full_ft_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_full_ft:") && msg.contains("pretrained=true"),
"expected pretrained refusal, got: {msg}"
);
}
#[test]
fn run_distill_refuses_pretrained_handle_tinyllama() {
let (_tmp, store, nn_dir, base, lua) = setup_tinyllama_scaffold();
let pretrained_like = base.for_test_pretrained_like();
let base_ud = lua
.create_userdata(NnHandle::TinyLlama(pretrained_like))
.unwrap();
let ds_ud = make_dataset_handle(&lua, overfit_row(), 5);
let opts = opts_table(&lua, base_full_ft_opts());
let msg = expect_err(run_distill_impl(
&store,
&nn_dir,
&LuaValue::UserData(base_ud),
&LuaValue::UserData(ds_ud),
opts,
));
assert!(
msg.contains("alc.nn.trainer.run_distill:") && msg.contains("pretrained=true"),
"expected pretrained refusal, got: {msg}"
);
}
}