use crate::errors::unsupported_operation;
use crate::export::{ExportConfig, ExportFormat, ExportPrecision, ModelExporter};
use crate::traits::Model;
use anyhow::{anyhow, Result};
pub const NNEF_UNSUPPORTED_REASON: &str =
"an NNEF package requires the model's operation graph, which the `Model` trait \
does not expose (`named_tensors` yields parameters only). TrustformeRS will not \
emit a synthesized graph under a real model's name.";
#[derive(Clone)]
pub struct NNEFExporter {
version: String,
extensions: Vec<String>,
}
impl NNEFExporter {
pub fn new() -> Self {
Self {
version: "1.0".to_string(),
extensions: vec!["KHR_enable_fragment_definitions".to_string()],
}
}
pub fn with_config(version: String, extensions: Vec<String>) -> Self {
Self {
version,
extensions,
}
}
pub fn version(&self) -> &str {
&self.version
}
pub fn extensions(&self) -> &[String] {
&self.extensions
}
pub fn get_input_shape(&self, config: &ExportConfig) -> Vec<i64> {
let batch_size = config.batch_size.unwrap_or(1) as i64;
if let Some(ref input_shape) = config.input_shape {
if input_shape.len() == 4 {
return input_shape.iter().map(|&x| x as i64).collect();
} else if input_shape.len() == 3 && input_shape[2] > 50 {
return vec![
batch_size,
input_shape[2] as i64,
input_shape[0] as i64,
input_shape[1] as i64,
];
}
}
if config.sequence_length.unwrap_or(512) > 8192 {
return vec![batch_size, config.sequence_length.unwrap_or(16000) as i64];
}
let sequence_length = config.sequence_length.unwrap_or(512) as i64;
if let Some(ref task_type) = config.task_type {
if task_type.to_lowercase().contains("multimodal")
|| task_type.to_lowercase().contains("vision")
{
return vec![batch_size, sequence_length, 3, 224, 224]; }
}
vec![batch_size, sequence_length]
}
pub fn get_output_shape(&self, config: &ExportConfig) -> Vec<i64> {
let batch_size = config.batch_size.unwrap_or(1) as i64;
let sequence_length = config.sequence_length.unwrap_or(512) as i64;
if let Some(ref task_type) = config.task_type {
match task_type.to_lowercase().as_str() {
"classification" | "text-classification" => {
let num_classes = config.vocab_size.unwrap_or(2) as i64; vec![batch_size, num_classes]
},
"token-classification" | "ner" => {
let num_labels = config.vocab_size.unwrap_or(9) as i64; vec![batch_size, sequence_length, num_labels]
},
"question-answering" | "qa" => {
vec![batch_size, sequence_length, 2]
},
"image-classification" => {
let num_classes = config.vocab_size.unwrap_or(1000) as i64; vec![batch_size, num_classes]
},
"object-detection" => {
vec![batch_size, 100, 6] },
"generation" | "text-generation" | "causal-lm" => {
let vocab_size = config.vocab_size.unwrap_or(50257) as i64; vec![batch_size, sequence_length, vocab_size]
},
"masked-lm" | "mlm" => {
let vocab_size = config.vocab_size.unwrap_or(30522) as i64; vec![batch_size, sequence_length, vocab_size]
},
"embedding" | "feature-extraction" => {
let hidden_size = 768; vec![batch_size, sequence_length, hidden_size]
},
"similarity" | "sentence-similarity" => {
let hidden_size = 768;
vec![batch_size, hidden_size]
},
_ => {
vec![batch_size, sequence_length, 768]
},
}
} else {
let input_shape = self.get_input_shape(config);
match input_shape.len() {
2 => {
vec![batch_size, sequence_length, 768]
},
3 => {
vec![batch_size, sequence_length, 768]
},
4 => {
vec![batch_size, 1000] },
_ => {
vec![batch_size, sequence_length, 768]
},
}
}
}
pub fn precision_to_dtype(&self, precision: ExportPrecision) -> &'static str {
match precision {
ExportPrecision::FP32 => "real32",
ExportPrecision::FP16 => "real16",
ExportPrecision::INT8 => "integer8",
ExportPrecision::INT4 => "integer4",
}
}
pub fn validate_config(&self, config: &ExportConfig) -> Result<()> {
if config.format != ExportFormat::NNEF {
return Err(anyhow!(
"Invalid format for NNEF exporter: {:?}",
config.format
));
}
match config.precision {
ExportPrecision::FP32 | ExportPrecision::FP16 => {},
ExportPrecision::INT8 | ExportPrecision::INT4 => {
if config.quantization.is_none() {
return Err(anyhow!(
"Quantization config required for integer precision"
));
}
},
}
Ok(())
}
}
impl ModelExporter for NNEFExporter {
fn export<M: Model>(&self, model: &M, config: &ExportConfig) -> Result<()> {
self.validate_config(config)?;
let _tensors = crate::export::collect_model_tensors(model)?;
Err(unsupported_operation("NNEF graph export", NNEF_UNSUPPORTED_REASON).into())
}
fn supported_formats(&self) -> Vec<ExportFormat> {
vec![ExportFormat::NNEF]
}
fn validate_model<M: Model>(&self, _model: &M, format: ExportFormat) -> Result<()> {
if format != ExportFormat::NNEF {
return Err(anyhow!("NNEF exporter only supports NNEF format"));
}
Ok(())
}
}
impl Default for NNEFExporter {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::export::test_support::TestModel;
#[test]
fn test_nnef_exporter_creation() {
let exporter = NNEFExporter::new();
let formats = exporter.supported_formats();
assert_eq!(formats, vec![ExportFormat::NNEF]);
}
#[test]
fn test_nnef_exporter_with_config() {
let exporter = NNEFExporter::with_config(
"1.0".to_string(),
vec!["KHR_enable_fragment_definitions".to_string()],
);
assert_eq!(exporter.version, "1.0");
assert_eq!(exporter.extensions.len(), 1);
}
#[test]
fn test_precision_to_dtype() {
let exporter = NNEFExporter::new();
assert_eq!(exporter.precision_to_dtype(ExportPrecision::FP32), "real32");
assert_eq!(exporter.precision_to_dtype(ExportPrecision::FP16), "real16");
assert_eq!(
exporter.precision_to_dtype(ExportPrecision::INT8),
"integer8"
);
assert_eq!(
exporter.precision_to_dtype(ExportPrecision::INT4),
"integer4"
);
}
#[test]
fn test_input_output_shapes() {
let exporter = NNEFExporter::new();
let config = ExportConfig {
format: ExportFormat::NNEF,
batch_size: Some(2),
sequence_length: Some(128),
..Default::default()
};
let input_shape = exporter.get_input_shape(&config);
assert_eq!(input_shape, vec![2, 128]);
let output_shape = exporter.get_output_shape(&config);
assert_eq!(output_shape, vec![2, 128, 768]);
}
#[test]
fn export_refuses_to_write_a_synthesized_package() {
let dir = std::env::temp_dir().join("trustformers_nnef_export_test");
std::fs::create_dir_all(&dir).expect("temp dir");
let output = dir.join("model");
let exporter = NNEFExporter::new();
let model = TestModel::with_seed(3.0);
let config = ExportConfig {
format: ExportFormat::NNEF,
output_path: output.to_string_lossy().to_string(),
..Default::default()
};
let err = exporter.export(&model, &config).expect_err("must not fabricate a graph");
assert!(
err.to_string().contains("Unsupported operation"),
"expected UnsupportedOperation, got: {err}"
);
assert!(
!output.with_extension("nnef").exists(),
"no NNEF package directory may be produced"
);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn export_reports_missing_weights_before_missing_topology() {
let exporter = NNEFExporter::new();
let model = TestModel::empty();
let config = ExportConfig {
format: ExportFormat::NNEF,
..Default::default()
};
let err = exporter.export(&model, &config).expect_err("no weights, no export");
assert!(
err.to_string().contains("named_tensors"),
"expected the missing-weights diagnostic, got: {err}"
);
}
#[test]
fn test_validate_config_success() {
let exporter = NNEFExporter::new();
let config = ExportConfig {
format: ExportFormat::NNEF,
precision: ExportPrecision::FP32,
..Default::default()
};
assert!(exporter.validate_config(&config).is_ok());
}
#[test]
fn test_validate_config_wrong_format() {
let exporter = NNEFExporter::new();
let config = ExportConfig {
format: ExportFormat::ONNX,
..Default::default()
};
assert!(exporter.validate_config(&config).is_err());
}
#[test]
fn test_validate_model_success() {
let exporter = NNEFExporter::new();
let model = TestModel::with_seed(1.0);
assert!(exporter.validate_model(&model, ExportFormat::NNEF).is_ok());
}
#[test]
fn test_validate_model_wrong_format() {
let exporter = NNEFExporter::new();
let model = TestModel::with_seed(1.0);
assert!(exporter.validate_model(&model, ExportFormat::ONNX).is_err());
}
}