use std::fs;
use std::path::{Path, PathBuf};
use anyhow::{Context, Result, anyhow};
use clap::{Parser, ValueEnum};
use burn::module::Module;
use burn::record::{BinFileRecorder, FullPrecisionSettings, Recorder};
use burn::tensor::backend::Backend;
use burn::tensor::{Int, Tensor, TensorData};
use burn_dragon_hatchling::wgpu::init_runtime;
use burn_dragon_hatchling::{
BDH, BDHConfig, GenerationConfig, ModelOverrides, TrainingConfig, TrainingHyperparameters,
load_training_config,
};
use burn_wgpu::Wgpu;
#[cfg(feature = "cuda")]
use burn_cuda::Cuda;
fn main() {
if let Err(err) = run() {
eprintln!("error: {err:#}");
std::process::exit(1);
}
}
fn run() -> Result<()> {
let args = Args::parse();
let mut config_paths = vec![PathBuf::from("config/base.toml")];
config_paths.extend(args.config.clone());
let config = load_training_config(&config_paths)?;
match args.backend {
BackendArg::Wgpu => {
infer_backend::<Wgpu<f32>, _>(&config, &args, "wgpu", |device| init_runtime(device))
}
BackendArg::Cuda => {
#[cfg(feature = "cuda")]
{
infer_backend::<Cuda<f32>, _>(&config, &args, "cuda", |_| {})
}
#[cfg(not(feature = "cuda"))]
{
Err(anyhow!(
"cuda backend selected but this build lacks `cuda` feature; rebuild with `--features cuda`"
))
}
}
}
}
fn infer_backend<B, Init>(
config: &TrainingConfig,
args: &Args,
backend_name: &str,
init_backend: Init,
) -> Result<()>
where
B: Backend + 'static,
B::Device: Clone,
Init: Fn(&B::Device),
{
B::seed(1337);
let device = B::Device::default();
init_backend(&device);
let checkpoint_dir = args
.checkpoint
.clone()
.unwrap_or_else(|| PathBuf::from("runs").join(backend_name).join("checkpoint"));
let (checkpoint_base, epoch) =
resolve_checkpoint_base(&checkpoint_dir, args.epoch, backend_name)?;
let mut model = BDH::<B>::new(build_model_config(&config.model), &device);
let recorder = BinFileRecorder::<FullPrecisionSettings>::new();
let record = recorder
.load::<<BDH<B> as Module<B>>::Record>(checkpoint_base.clone(), &device)
.with_context(|| {
format!(
"failed to load checkpoint {}",
format_checkpoint(&checkpoint_base)
)
})?;
model = model.load_record(record);
let mut generation = config.generation.clone();
apply_generation_overrides(&mut generation, args);
let output = generate_text::<B>(&model, &device, &config.training, &generation)?;
eprintln!(
"Loaded epoch {epoch} from {} using {backend_name} backend.",
format_checkpoint(&checkpoint_base)
);
println!("{output}");
Ok(())
}
fn build_model_config(overrides: &ModelOverrides) -> BDHConfig {
let mut model_config = BDHConfig::default();
if let Some(n_layer) = overrides.n_layer {
model_config.n_layer = n_layer;
}
if let Some(n_embd) = overrides.n_embd {
model_config.n_embd = n_embd;
}
if let Some(n_head) = overrides.n_head {
model_config.n_head = n_head;
}
if let Some(multiplier) = overrides.mlp_internal_dim_multiplier {
model_config.mlp_internal_dim_multiplier = multiplier;
}
if let Some(dropout) = overrides.dropout {
model_config.dropout = dropout;
}
if let Some(enabled) = overrides.fused_kernels {
model_config.fused_kernels.enabled = enabled;
}
if let Some(block) = overrides.block_size {
model_config.fused_kernels.set_block_sizes(block, block);
}
if let Some(use_alibi) = overrides.use_alibi {
model_config.fused_kernels.set_use_alibi(use_alibi);
if !use_alibi {
model_config
.fused_kernels
.set_alibi_slopes(vec![0.0; model_config.n_head]);
}
}
model_config
}
fn apply_generation_overrides(generation: &mut GenerationConfig, args: &Args) {
if let Some(prompt) = &args.prompt {
generation.prompt = prompt.clone();
}
if let Some(max_tokens) = args.max_tokens {
generation.max_tokens = max_tokens;
}
if let Some(temperature) = args.temperature {
generation.temperature = temperature;
}
if let Some(top_k) = args.top_k {
generation.top_k = Some(top_k);
}
}
fn generate_text<B: Backend>(
model: &BDH<B>,
device: &B::Device,
training: &TrainingHyperparameters,
generation: &GenerationConfig,
) -> Result<String> {
let mut prompt_tokens: Vec<i64> = generation
.prompt
.as_bytes()
.iter()
.map(|b| *b as i64)
.collect();
if prompt_tokens.len() > training.block_size {
prompt_tokens = prompt_tokens[prompt_tokens.len() - training.block_size..].to_vec();
}
let prompt_len = prompt_tokens.len();
let prompt_tensor =
Tensor::<B, 2, Int>::from_data(TensorData::new(prompt_tokens, [1, prompt_len]), device);
let generated = model.generate(
prompt_tensor,
generation.max_tokens,
generation.temperature,
generation.top_k,
);
let tokens = generated
.into_data()
.convert::<i64>()
.into_vec::<i64>()
.map_err(|err| anyhow!("{err:?}"))?;
let bytes: Vec<u8> = tokens.iter().map(|&tok| tok as u8).collect();
Ok(String::from_utf8_lossy(&bytes).to_string())
}
fn resolve_checkpoint_base(
path: &Path,
epoch: Option<usize>,
backend_name: &str,
) -> Result<(PathBuf, usize)> {
if path.is_dir() {
let target_epoch = epoch.unwrap_or(find_latest_epoch(path)?);
let base = path.join(format!("model-{target_epoch}"));
ensure_checkpoint_exists(&base)?;
return Ok((base, target_epoch));
}
let mut base = if path.extension().is_some() {
let mut without_ext = path.to_path_buf();
without_ext.set_extension("");
without_ext
} else {
path.to_path_buf()
};
let detected_epoch = parse_epoch_from_stem(&base);
let target_epoch = match (epoch, detected_epoch) {
(Some(explicit), Some(detected)) if explicit != detected => {
let parent = base.parent().map(Path::to_path_buf).unwrap_or_default();
base = parent.join(format!("model-{explicit}"));
explicit
}
(Some(explicit), _) => {
if detected_epoch.is_none() {
let parent = base
.parent()
.map(Path::to_path_buf)
.unwrap_or_else(|| PathBuf::from("runs").join(backend_name).join("checkpoint"));
base = parent.join(format!("model-{explicit}"));
}
explicit
}
(None, Some(detected)) => detected,
(None, None) => {
return Err(anyhow!(
"unable to infer checkpoint epoch from {}; provide --epoch",
path.display()
));
}
};
ensure_checkpoint_exists(&base)?;
Ok((base, target_epoch))
}
fn ensure_checkpoint_exists(base: &Path) -> Result<()> {
let mut candidate = base.to_path_buf();
candidate.set_extension("bin");
if candidate.is_file() {
return Ok(());
}
Err(anyhow!("checkpoint file {}.bin not found", base.display()))
}
fn find_latest_epoch(dir: &Path) -> Result<usize> {
let mut max_epoch = None;
for entry in fs::read_dir(dir)
.with_context(|| format!("failed to read checkpoint directory {}", dir.display()))?
{
let entry = entry?;
if !entry.file_type()?.is_file() {
continue;
}
let mut base = entry.path();
base.set_extension("");
if let Some(epoch) = parse_epoch_from_stem(&base) {
let updated = max_epoch
.map(|current: usize| current.max(epoch))
.unwrap_or(epoch);
max_epoch = Some(updated);
}
}
max_epoch.ok_or_else(|| anyhow!("no model checkpoints found in {}", dir.display()))
}
fn parse_epoch_from_stem(path: &Path) -> Option<usize> {
let stem = path.file_name()?.to_string_lossy();
let stem = stem.strip_suffix(".bin").unwrap_or(&stem);
let epoch_part = stem.strip_prefix("model-")?;
epoch_part.parse().ok()
}
fn format_checkpoint(base: &Path) -> String {
let mut path = base.to_path_buf();
path.set_extension("bin");
path.display().to_string()
}
#[derive(Parser, Debug)]
#[command(
author,
version,
about = "Run inference with a trained Baby Dragon Hatchling model"
)]
struct Args {
#[arg(short = 'c', long = "config", value_name = "PATH")]
config: Vec<PathBuf>,
#[arg(long, value_enum, default_value_t = BackendArg::Cuda)]
backend: BackendArg,
#[arg(long, value_name = "PATH")]
checkpoint: Option<PathBuf>,
#[arg(long, value_name = "N")]
epoch: Option<usize>,
#[arg(long)]
prompt: Option<String>,
#[arg(long, value_name = "N")]
max_tokens: Option<usize>,
#[arg(long, value_name = "T")]
temperature: Option<f32>,
#[arg(long, value_name = "K")]
top_k: Option<usize>,
}
#[derive(Copy, Clone, Debug, ValueEnum)]
enum BackendArg {
Wgpu,
Cuda,
}