use std::env;
use std::fs;
use std::path::Path;
fn main() {
println!("cargo:rerun-if-changed=resources/models.yaml");
let out_dir = env::var("OUT_DIR").unwrap();
let dest_path = Path::new(&out_dir).join("generated_models.rs");
let yaml_content =
fs::read_to_string("resources/models.yaml").expect("Failed to read models.yaml");
let models: Vec<YamlModel> =
serde_yaml::from_str(&yaml_content).expect("Failed to parse models.yaml");
let code = generate_models_code(&models);
fs::write(&dest_path, code).expect("Failed to write generated code");
}
fn generate_models_code(models: &[YamlModel]) -> String {
let mut code = String::from("// Auto-generated from resources/models.yaml - do not edit\n\n");
for model in models {
code.push_str(&generate_model_struct(model));
}
code.push_str(
"pub fn register_all_models(factory: &mut crate::models::base::ModelFactory) {\n",
);
for model in models {
let struct_name = model_id_to_struct_name(&model.id);
code.push_str(&format!(
" factory.register::<{}>(\"{}\");\n",
struct_name, model.id
));
}
code.push_str("}\n");
code
}
fn generate_model_struct(model: &YamlModel) -> String {
let struct_name = model_id_to_struct_name(&model.id);
let mut s = String::new();
s.push_str(&format!("pub struct {struct_name} {{}}\n\n"));
s.push_str(&format!(
"impl crate::registry::ConfigConstructable for {struct_name} {{\n"
));
s.push_str(" fn new(_cfg: &serde_json::Value) -> Self { Self {} }\n");
s.push_str("}\n\n");
s.push_str(&format!(
"impl crate::models::base::Model for {struct_name} {{\n"
));
s.push_str(&format!(
" fn family(&self) -> &str {{ {:?} }}\n",
model.family
));
s.push_str(&format!(
" fn version(&self) -> &str {{ {:?} }}\n",
model.version
));
s.push_str(&format!(" fn size(&self) -> u64 {{ {} }}\n", model.size));
s.push_str(&format!(
" fn context_length(&self) -> u64 {{ {} }}\n",
model.context_length
));
s.push_str(&format!(" fn model_type(&self) -> &crate::models::base::ModelType {{ &crate::models::base::ModelType::{} }}\n", model.model_type));
s.push_str(&format!(
" fn huggingface_repo(&self) -> &str {{ {:?} }}\n",
model.huggingface_repo
));
s.push_str(&format!(
" fn native_dtype(&self) -> &str {{ {:?} }}\n",
model.native_dtype
));
s.push_str(" fn architecture(&self) -> &crate::models::base::ModelArchitecture {\n");
s.push_str(" static ARCHITECTURE: std::sync::LazyLock<crate::models::base::ModelArchitecture> = std::sync::LazyLock::new(|| ");
s.push_str(&generate_architecture_literal(&model.architecture));
s.push_str(");\n");
s.push_str(" &ARCHITECTURE\n");
s.push_str(" }\n");
s.push_str(" fn variants(&self) -> &[crate::models::base::ModelVariant] {\n");
s.push_str(" static VARIANTS: std::sync::LazyLock<Vec<crate::models::base::ModelVariant>> = std::sync::LazyLock::new(|| vec![\n");
for variant in &model.variants {
s.push_str(" crate::models::base::ModelVariant {\n");
s.push_str(&format!(
" format: {:?}.to_string(),\n",
variant.format
));
s.push_str(&format!(
" precision: {:?}.to_string(),\n",
variant.precision
));
s.push_str(&format!(
" size_gb: {},\n",
format_float(variant.size_gb)
));
s.push_str(&format!(
" url: {:?}.to_string(),\n",
variant.url
));
s.push_str(" },\n");
}
s.push_str(" ]);\n");
s.push_str(" &VARIANTS\n");
s.push_str(" }\n");
s.push_str(" fn description(&self) -> Option<&str> {\n");
if let Some(ref desc) = model.description {
s.push_str(&format!(" Some({desc:?})\n"));
} else {
s.push_str(" None\n");
}
s.push_str(" }\n");
s.push_str(" fn tags(&self) -> &[String] {\n");
s.push_str(" static TAGS: std::sync::LazyLock<Vec<String>> = std::sync::LazyLock::new(|| vec![\n");
for tag in &model.tags {
s.push_str(&format!(" {tag:?}.to_string(),\n"));
}
s.push_str(" ]);\n");
s.push_str(" &TAGS\n");
s.push_str(" }\n");
s.push_str(" fn supported_functions(&self) -> &[crate::models::base::ModelFunction] {\n");
s.push_str(" static FUNCS: std::sync::LazyLock<Vec<crate::models::base::ModelFunction>> = std::sync::LazyLock::new(|| vec![\n");
for func in &model.supported_functions {
s.push_str(&format!(
" crate::models::base::ModelFunction::{func},\n"
));
}
s.push_str(" ]);\n");
s.push_str(" &FUNCS\n");
s.push_str(" }\n");
s.push_str("}\n\n");
s.push_str(&format!(
"impl crate::models::base::HasModelMetadata for {struct_name} {{\n"
));
s.push_str(" fn metadata() -> crate::models::base::ModelMetadata {\n");
s.push_str(&generate_metadata_literal(model));
s.push_str(" }\n");
s.push_str("}\n\n");
s
}
fn generate_metadata_literal(model: &YamlModel) -> String {
let mut s = String::new();
s.push_str(" crate::models::base::ModelMetadata {\n");
s.push_str(&format!(
" family: {:?}.to_string(),\n",
model.family
));
s.push_str(&format!(
" version: {:?}.to_string(),\n",
model.version
));
s.push_str(&format!(" size: {},\n", model.size));
s.push_str(&format!(
" context_length: {},\n",
model.context_length
));
s.push_str(&format!(
" model_type: crate::models::base::ModelType::{},\n",
model.model_type
));
s.push_str(&format!(
" huggingface_repo: {:?}.to_string(),\n",
model.huggingface_repo
));
s.push_str(&format!(
" native_dtype: {:?}.to_string(),\n",
model.native_dtype
));
s.push_str(" architecture: ");
s.push_str(&generate_architecture_literal(&model.architecture));
s.push_str(",\n");
s.push_str(" variants: vec![\n");
for variant in &model.variants {
s.push_str(" crate::models::base::ModelVariant {\n");
s.push_str(&format!(
" format: {:?}.to_string(),\n",
variant.format
));
s.push_str(&format!(
" precision: {:?}.to_string(),\n",
variant.precision
));
s.push_str(&format!(
" size_gb: {},\n",
format_float(variant.size_gb)
));
s.push_str(&format!(
" url: {:?}.to_string(),\n",
variant.url
));
s.push_str(" },\n");
}
s.push_str(" ],\n");
if let Some(ref desc) = model.description {
s.push_str(&format!(
" description: Some({desc:?}.to_string()),\n"
));
} else {
s.push_str(" description: None,\n");
}
s.push_str(" tags: vec![\n");
for tag in &model.tags {
s.push_str(&format!(" {tag:?}.to_string(),\n"));
}
s.push_str(" ],\n");
s.push_str(" supported_functions: vec![\n");
for func in &model.supported_functions {
s.push_str(&format!(
" crate::models::base::ModelFunction::{func},\n"
));
}
s.push_str(" ],\n");
s.push_str(" }\n");
s
}
fn generate_architecture_literal(arch: &YamlArchitecture) -> String {
let mut s = String::new();
s.push_str("crate::models::base::ModelArchitecture {\n");
s.push_str(&format!(
" num_hidden_layers: {},\n",
arch.num_hidden_layers
));
s.push_str(&format!(" hidden_size: {},\n", arch.hidden_size));
s.push_str(&format!(
" num_attention_heads: {},\n",
arch.num_attention_heads
));
s.push_str(&format!(
" num_key_value_heads: {},\n",
arch.num_key_value_heads
));
s.push_str(&format!(" head_dim: {},\n", arch.head_dim));
s.push_str(" layer_types: vec![\n");
for ltc in &arch.layer_types {
s.push_str(" crate::models::base::LayerTypeCount {\n");
s.push_str(&format!(
" kind: {},\n",
generate_layer_kind_literal(ltc)
));
s.push_str(&format!(" count: {},\n", ltc.count));
s.push_str(" },\n");
}
s.push_str(" ],\n");
s.push_str(" }");
s
}
fn generate_layer_kind_literal(ltc: &YamlLayerTypeCount) -> String {
match ltc.kind.as_str() {
"full_attention" => "crate::models::base::LayerKind::FullAttention".to_string(),
"sliding_attention" => {
let window = ltc
.window
.unwrap_or_else(|| panic!("sliding_attention layer_type missing `window` field"));
format!("crate::models::base::LayerKind::SlidingAttention {{ window: {window} }}")
}
"recurrent" => {
let mamba = ltc
.mamba
.as_ref()
.unwrap_or_else(|| panic!("recurrent layer_type missing `mamba` shape"));
format!(
"crate::models::base::LayerKind::Recurrent(crate::models::base::MambaShape {{ d_conv: {}, d_state: {}, d_inner: {}, n_groups: {} }})",
mamba.d_conv, mamba.d_state, mamba.d_inner, mamba.n_groups
)
}
other => panic!("Unknown layer_types kind: {other:?}"),
}
}
fn format_float(value: f64) -> String {
if value.fract() == 0.0 {
format!("{value:.1}")
} else {
value.to_string()
}
}
fn model_id_to_struct_name(id: &str) -> String {
id.split('-')
.map(|part| {
if part.contains('.') {
part.replace('.', "")
} else {
let mut chars = part.chars();
match chars.next() {
None => String::new(),
Some(first) => first.to_uppercase().collect::<String>() + chars.as_str(),
}
}
})
.collect::<String>() }
#[derive(serde::Deserialize)]
struct YamlModel {
id: String,
family: String,
version: String,
size: u64,
context_length: u64,
model_type: String,
huggingface_repo: String,
native_dtype: String,
architecture: YamlArchitecture,
variants: Vec<YamlModelVariant>,
description: Option<String>,
tags: Vec<String>,
supported_functions: Vec<String>,
}
#[derive(serde::Deserialize)]
struct YamlModelVariant {
format: String,
precision: String,
size_gb: f64,
url: String,
}
#[derive(serde::Deserialize)]
struct YamlArchitecture {
num_hidden_layers: u64,
hidden_size: u64,
num_attention_heads: u64,
num_key_value_heads: u64,
head_dim: u64,
layer_types: Vec<YamlLayerTypeCount>,
}
#[derive(serde::Deserialize)]
struct YamlLayerTypeCount {
kind: String,
count: u64,
window: Option<u64>,
mamba: Option<YamlMambaShape>,
}
#[derive(serde::Deserialize)]
struct YamlMambaShape {
d_conv: u64,
d_state: u64,
d_inner: u64,
n_groups: u64,
}