mod common;
use coremlit::{ComputeUnits, DataType, Features, Model, MultiArray};
#[test]
#[ignore = "requires local tiny model (see plan Task 8 Step 1)"]
fn loads_mel_model_and_reads_description() {
let model = Model::load(
common::tiny_dir().join("MelSpectrogram.mlmodelc"),
ComputeUnits::CpuOnly,
)
.unwrap();
let description = model.description();
let input = description
.input("audio")
.expect("mel input feature `audio`");
assert_eq!(input.shape(), &[480_000]);
let output = description
.output("melspectrogram_features")
.expect("mel output feature `melspectrogram_features`");
assert_eq!(output.shape(), &[1, 80, 1, 3000]);
assert_eq!(output.data_type(), Some(DataType::F16));
}
#[test]
#[ignore = "requires local tiny model (see plan Task 8 Step 1)"]
fn mel_predict_produces_f16_spectrogram() {
let model = Model::load(
common::tiny_dir().join("MelSpectrogram.mlmodelc"),
ComputeUnits::CpuOnly,
)
.unwrap();
let audio = MultiArray::zeros(&[480_000], DataType::F32).unwrap();
let outputs = model
.predict(&Features::new().with("audio", audio))
.unwrap();
let mel = outputs
.get("melspectrogram_features")
.expect("mel output present");
assert_eq!(mel.shape(), vec![1, 80, 1, 3000]);
assert_eq!(mel.data_type(), DataType::F16);
let mut values = vec![half::f16::from_f32(0.0); mel.count()];
mel.copy_into(&mut values).unwrap();
assert!(values.iter().all(|v| v.to_f32().is_finite()));
}
#[test]
#[ignore = "requires local tiny model (see plan Task 8 Step 1)"]
fn prewarm_loads_and_drops() {
coremlit::Model::prewarm(
common::tiny_dir().join("MelSpectrogram.mlmodelc"),
ComputeUnits::CpuOnly,
)
.unwrap();
}
#[test]
#[ignore = "requires local tiny model (see plan Task 8 Step 1)"]
fn stateless_model_accepts_state_prediction() {
let model = Model::load(
common::tiny_dir().join("MelSpectrogram.mlmodelc"),
ComputeUnits::CpuOnly,
)
.unwrap();
if !model.supports_state() {
eprintln!("skipping: MLState unavailable on this OS");
return;
}
let mut state = model.make_state().unwrap();
let stateful_audio = MultiArray::zeros(&[480_000], DataType::F32).unwrap();
let stateful_outputs = model
.predict_with_state(&Features::new().with("audio", stateful_audio), &mut state)
.unwrap();
let stateful_mel = stateful_outputs.get("melspectrogram_features").unwrap();
assert_eq!(stateful_mel.shape(), vec![1, 80, 1, 3000]);
let plain_audio = MultiArray::zeros(&[480_000], DataType::F32).unwrap();
let plain_outputs = model
.predict(&Features::new().with("audio", plain_audio))
.unwrap();
let plain_mel = plain_outputs.get("melspectrogram_features").unwrap();
let mut stateful_values = vec![half::f16::from_f32(0.0); stateful_mel.count()];
stateful_mel.copy_into(&mut stateful_values).unwrap();
let mut plain_values = vec![half::f16::from_f32(0.0); plain_mel.count()];
plain_mel.copy_into(&mut plain_values).unwrap();
assert_eq!(stateful_values, plain_values);
}
#[test]
#[ignore = "requires local tiny model (see plan Task 8 Step 1)"]
fn loads_through_non_ascii_symlinked_path() {
let link_dir = std::env::temp_dir().join("coremlit-tests");
std::fs::create_dir_all(&link_dir).unwrap();
let link = link_dir.join("modèle-mel.mlmodelc");
let _ = std::fs::remove_file(&link);
std::os::unix::fs::symlink(common::tiny_dir().join("MelSpectrogram.mlmodelc"), &link).unwrap();
let model = Model::load(&link, ComputeUnits::CpuOnly).unwrap();
assert!(model.description().input("audio").is_some());
let _ = std::fs::remove_file(&link);
}
#[test]
#[ignore = "requires local tiny model (see plan Task 8 Step 1)"]
fn mel_predict_with_borrowed_inputs() {
let model = Model::load(
common::tiny_dir().join("MelSpectrogram.mlmodelc"),
ComputeUnits::CpuOnly,
)
.unwrap();
let audio = MultiArray::zeros(&[480_000], DataType::F32).unwrap();
let outputs = model.predict_with(&[("audio", &audio)]).unwrap();
assert_eq!(
outputs.get("melspectrogram_features").unwrap().shape(),
vec![1, 80, 1, 3000]
);
assert_eq!(audio.count(), 480_000);
}