use super::*;
#[test]
fn repeat_pad_f32_empty_source_returns_zeros() {
assert_eq!(repeat_pad_f32(&[], 5), vec![0.0; 5]);
}
#[test]
fn repeat_pad_f32_empty_source_and_zero_target_returns_empty() {
assert_eq!(repeat_pad_f32(&[], 0), Vec::<f32>::new());
}
#[test]
fn repeat_pad_f32_exact_length_is_identity() {
assert_eq!(
repeat_pad_f32(&[1.0, 2.0, 3.0, 4.0], 4),
vec![1.0, 2.0, 3.0, 4.0]
);
}
#[test]
fn repeat_pad_f32_short_source_tiles_periodically() {
assert_eq!(
repeat_pad_f32(&[1.0, 2.0, 3.0], 10),
vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 1.0]
);
}
#[test]
fn repeat_pad_f32_single_element_source_fills_uniformly() {
assert_eq!(repeat_pad_f32(&[7.0], 4), vec![7.0, 7.0, 7.0, 7.0]);
}
#[test]
fn repeat_pad_f32_longer_source_truncates() {
assert_eq!(
repeat_pad_f32(&[1.0, 2.0, 3.0, 4.0, 5.0], 3),
vec![1.0, 2.0, 3.0]
);
}
#[test]
fn repeat_pad_f32_zero_target_returns_empty() {
assert_eq!(repeat_pad_f32(&[1.0, 2.0], 0), Vec::<f32>::new());
}
fn doubling_copy_simulation(source: &[f32], target_len: usize) -> Vec<f32> {
let mut buf = vec![0.0f32; target_len];
let mut sample_count = source.len().min(target_len);
buf[..sample_count].copy_from_slice(&source[..sample_count]);
if sample_count == 0 {
return buf;
}
while sample_count < target_len {
let copy_count = sample_count.min(target_len - sample_count);
let (filled, rest) = buf.split_at_mut(sample_count);
rest[..copy_count].copy_from_slice(&filled[..copy_count]);
sample_count += copy_count;
}
buf
}
#[test]
fn doubling_copy_simulation_matches_repeat_pad_f32_non_power_of_two_lengths() {
let cases: &[(&[f32], usize)] = &[
(&[1.0, 2.0, 3.0], 10),
(&[1.0, 2.0, 3.0, 4.0, 5.0], 13),
(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0], 20),
(&[9.0], 7),
(&[1.0, 2.0], 2),
(&[1.0, 2.0, 3.0, 4.0, 5.0], 3),
(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], 1),
];
for &(source, target_len) in cases {
assert_eq!(
repeat_pad_f32(source, target_len),
doubling_copy_simulation(source, target_len),
"mismatch for source={source:?}, target_len={target_len}"
);
}
}
#[test]
fn doubling_copy_simulation_matches_repeat_pad_f32_on_empty_source() {
assert_eq!(repeat_pad_f32(&[], 6), doubling_copy_simulation(&[], 6));
}
#[test]
fn mask_row_f32_converts_true_and_false() {
assert_eq!(
mask_row_f32(&[true, false, true, true, false]),
vec![1.0, 0.0, 1.0, 1.0, 0.0]
);
}
#[test]
fn mask_row_f32_empty_mask_is_empty() {
assert_eq!(mask_row_f32(&[]), Vec::<f32>::new());
}
#[test]
fn mask_row_f32_all_false() {
assert_eq!(mask_row_f32(&[false, false]), vec![0.0, 0.0]);
}
#[test]
fn mask_row_f32_all_true() {
assert_eq!(mask_row_f32(&[true, true, true]), vec![1.0, 1.0, 1.0]);
}
#[test]
fn mask_row_padding_short_mask_tiles_after_conversion() {
let converted = mask_row_f32(&[true, false]);
assert_eq!(repeat_pad_f32(&converted, 5), vec![1.0, 0.0, 1.0, 0.0, 1.0]);
}
#[test]
fn build_waveform_repeats_the_same_row_in_every_slot() {
let out = build_waveform(&[1.0, 2.0]);
assert_eq!(
out.len(),
EMBED_SLOTS * crate::audio::speaker::segment::SEG_CHUNK_SAMPLES
);
for slot in out
.as_chunks::<{ crate::audio::speaker::segment::SEG_CHUNK_SAMPLES }>()
.0
{
assert_eq!(slot[0], 1.0);
assert_eq!(slot[1], 2.0);
assert_eq!(slot[2], 1.0); }
}
#[test]
fn build_masks_pads_each_slot_independently() {
let mask_a = [true, false, true];
let mask_b: [bool; 0] = [];
let mask_c = [true];
let masks: [&[bool]; EMBED_SLOTS] = [&mask_a, &mask_b, &mask_c];
let out = build_masks(&masks, 4);
assert_eq!(out.len(), EMBED_SLOTS * 4);
assert_eq!(&out[0..4], &[1.0, 0.0, 1.0, 1.0]); assert_eq!(&out[4..8], &[0.0, 0.0, 0.0, 0.0]); assert_eq!(&out[8..12], &[1.0, 1.0, 1.0, 1.0]); }
#[test]
fn check_mask_active_accepts_one_active_frame() {
assert_eq!(check_mask_active(&[false, true, false]), Ok(()));
}
#[test]
fn check_mask_active_accepts_all_active() {
assert_eq!(check_mask_active(&[true, true]), Ok(()));
}
#[test]
fn check_mask_active_rejects_all_false() {
assert_eq!(
check_mask_active(&[false, false, false]),
Err(InferError::EmptyMask)
);
}
#[test]
fn check_mask_active_rejects_empty_mask() {
assert_eq!(check_mask_active(&[]), Err(InferError::EmptyMask));
}
#[test]
fn check_finite_input_accepts_all_finite() {
assert_eq!(check_finite_input(&[0.0, 1.0, -1.0]), Ok(()));
}
#[test]
fn check_finite_input_rejects_nan_at_reported_index() {
assert_eq!(
check_finite_input(&[0.0, f32::NAN, 2.0]),
Err(InferError::NonFiniteInput(1))
);
}
#[test]
fn check_finite_input_rejects_positive_infinity() {
assert_eq!(
check_finite_input(&[f32::INFINITY]),
Err(InferError::NonFiniteInput(0))
);
}
#[test]
fn check_finite_input_rejects_negative_infinity() {
assert_eq!(
check_finite_input(&[0.0, 0.0, f32::NEG_INFINITY]),
Err(InferError::NonFiniteInput(2))
);
}
#[test]
fn check_finite_output_accepts_all_finite() {
assert_eq!(check_finite_output(&[0.0, 1.0, -1.0]), Ok(()));
}
#[test]
fn check_finite_output_rejects_nan_at_reported_index() {
assert_eq!(
check_finite_output(&[0.0, f32::NAN, 2.0]),
Err(InferError::NonFiniteOutput(1))
);
}
#[test]
fn check_finite_output_reports_first_offending_index() {
assert_eq!(
check_finite_output(&[f32::NAN, f32::INFINITY]),
Err(InferError::NonFiniteOutput(0))
);
}
#[test]
fn check_output_shape_accepts_correct_shape() {
assert_eq!(check_output_shape(&[EMBED_SLOTS, EMBEDDING_DIM]), Ok(()));
}
#[test]
fn check_output_shape_rejects_swapped_axes() {
assert_eq!(
check_output_shape(&[EMBEDDING_DIM, EMBED_SLOTS]),
Err(InferError::OutputShape(OutputShape::new(
vec![EMBEDDING_DIM, EMBED_SLOTS],
vec![EMBED_SLOTS, EMBEDDING_DIM]
)))
);
}
#[test]
fn check_output_shape_rejects_wrong_rank() {
assert_eq!(
check_output_shape(&[EMBED_SLOTS * EMBEDDING_DIM]),
Err(InferError::OutputShape(OutputShape::new(
vec![EMBED_SLOTS * EMBEDDING_DIM],
vec![EMBED_SLOTS, EMBEDDING_DIM]
)))
);
}
#[test]
fn check_output_shape_rejects_wrong_slot_count() {
assert_eq!(
check_output_shape(&[2, EMBEDDING_DIM]),
Err(InferError::OutputShape(OutputShape::new(
vec![2, EMBEDDING_DIM],
vec![EMBED_SLOTS, EMBEDDING_DIM]
)))
);
}
#[test]
fn check_output_shape_rejects_wrong_embedding_dim() {
assert_eq!(
check_output_shape(&[EMBED_SLOTS, EMBEDDING_DIM - 1]),
Err(InferError::OutputShape(OutputShape::new(
vec![EMBED_SLOTS, EMBEDDING_DIM - 1],
vec![EMBED_SLOTS, EMBEDDING_DIM]
)))
);
}
#[test]
fn options_new_defaults_to_all_compute() {
assert_eq!(EmbedModelOptions::new().compute(), ComputeUnits::All);
}
#[test]
fn options_default_matches_new() {
assert_eq!(EmbedModelOptions::default(), EmbedModelOptions::new());
}
#[test]
fn options_with_compute_overrides() {
let options = EmbedModelOptions::new().with_compute(ComputeUnits::CpuOnly);
assert_eq!(options.compute(), ComputeUnits::CpuOnly);
}
#[test]
fn options_set_compute_in_place() {
let mut options = EmbedModelOptions::new();
options.set_compute(ComputeUnits::CpuAndNeuralEngine);
assert_eq!(options.compute(), ComputeUnits::CpuAndNeuralEngine);
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_missing_compute_defaults_to_all() {
let options: EmbedModelOptions = serde_json::from_str("{}").unwrap();
assert_eq!(options.compute(), ComputeUnits::All);
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_round_trips_explicit_compute() {
let options: EmbedModelOptions = serde_json::from_str(r#"{"compute":"cpu_only"}"#).unwrap();
assert_eq!(options.compute(), ComputeUnits::CpuOnly);
let json = serde_json::to_string(&options).unwrap();
assert!(json.contains("cpu_only"), "round-tripped json: {json}");
}
fn models_dir() -> std::path::PathBuf {
std::env::var_os("SPEAKERKIT_TEST_MODELS").map_or_else(
|| crate::tests::models_root().join("speakerkit"),
std::path::PathBuf::from,
)
}
fn embed_v2_path() -> std::path::PathBuf {
models_dir().join("wespeaker_v2.mlmodelc")
}
fn embed_fp32_path() -> std::path::PathBuf {
models_dir().join("wespeaker.mlmodelc")
}
fn load_embed_model(path: std::path::PathBuf) -> EmbedModel {
EmbedModel::from_file_with(
path,
EmbedModelOptions::new().with_compute(ComputeUnits::CpuOnly),
)
.expect("load embedding model")
}
fn synthetic_samples(len: usize) -> Vec<f32> {
(0..len)
.map(|i| {
let t = i as f32 / 16_000.0; 0.2 * (2.0 * core::f32::consts::PI * 220.0 * t).sin()
+ 0.1 * (2.0 * core::f32::consts::PI * 440.0 * t).sin()
})
.collect()
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn from_file_loads_and_reports_mask_frame_count_v2() {
let model = load_embed_model(embed_v2_path());
assert_eq!(model.num_mask_frames(), 589);
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn from_file_loads_and_reports_mask_frame_count_fp32() {
let model = load_embed_model(embed_fp32_path());
assert_eq!(model.num_mask_frames(), 589);
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn from_file_rejects_wrong_contract_model() {
let path = models_dir().join("pyannote_segmentation.mlmodelc");
let err = EmbedModel::from_file(path).expect_err("wrong contract must be rejected");
assert!(matches!(
err,
ModelError::ContractMismatch(m) if m.feature() == "waveform"
));
}
fn embed_chunk_produces_correctly_shaped_finite_embeddings(path: std::path::PathBuf) {
let model = load_embed_model(path);
let samples = synthetic_samples(crate::audio::speaker::segment::SEG_CHUNK_SAMPLES);
let mask = vec![true; model.num_mask_frames()];
let masks: [&[bool]; EMBED_SLOTS] = [&mask, &mask, &mask];
let out = model
.embed_chunk(&samples, &masks)
.expect("embed_chunk on real audio");
for row in out.iter() {
assert_eq!(row.len(), EMBEDDING_DIM);
assert!(
row.iter().all(|v| v.is_finite()),
"all embedding values finite"
);
}
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_produces_correctly_shaped_finite_embeddings_v2() {
embed_chunk_produces_correctly_shaped_finite_embeddings(embed_v2_path());
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_produces_correctly_shaped_finite_embeddings_fp32() {
embed_chunk_produces_correctly_shaped_finite_embeddings(embed_fp32_path());
}
fn embed_chunk_is_deterministic_across_repeated_calls(path: std::path::PathBuf) {
let model = load_embed_model(path);
let samples = synthetic_samples(crate::audio::speaker::segment::SEG_CHUNK_SAMPLES);
let mask = vec![true; model.num_mask_frames()];
let masks: [&[bool]; EMBED_SLOTS] = [&mask, &mask, &mask];
let first = model
.embed_chunk(&samples, &masks)
.expect("first embed_chunk");
let second = model
.embed_chunk(&samples, &masks)
.expect("second embed_chunk");
assert_eq!(
first, second,
"repeated embed_chunk must be bit-identical WITHIN one artifact"
);
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_is_deterministic_across_repeated_calls_v2() {
embed_chunk_is_deterministic_across_repeated_calls(embed_v2_path());
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_is_deterministic_across_repeated_calls_fp32() {
embed_chunk_is_deterministic_across_repeated_calls(embed_fp32_path());
}
fn embed_chunk_with_frame_mask_is_raw_not_unit_norm(path: std::path::PathBuf) {
let model = load_embed_model(path);
let samples = synthetic_samples(crate::audio::speaker::segment::SEG_CHUNK_SAMPLES);
let mask = vec![true; model.num_mask_frames()];
let embedding = model
.embed_chunk_with_frame_mask(&samples, &mask)
.expect("embed_chunk_with_frame_mask on real audio");
let norm: f32 = embedding.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() > 1e-3,
"raw WeSpeaker output must not be unit-normalized (dia normalizes downstream, not here); got norm {norm}"
);
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_with_frame_mask_is_raw_not_unit_norm_v2() {
embed_chunk_with_frame_mask_is_raw_not_unit_norm(embed_v2_path());
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_with_frame_mask_is_raw_not_unit_norm_fp32() {
embed_chunk_with_frame_mask_is_raw_not_unit_norm(embed_fp32_path());
}
fn embed_chunk_with_frame_mask_matches_batched_slot_zero(path: std::path::PathBuf) {
let model = load_embed_model(path);
let samples = synthetic_samples(crate::audio::speaker::segment::SEG_CHUNK_SAMPLES);
let mut mask = vec![false; model.num_mask_frames()];
for m in mask.iter_mut().step_by(3) {
*m = true;
}
let veneer = model
.embed_chunk_with_frame_mask(&samples, &mask)
.expect("embed_chunk_with_frame_mask");
let empty: &[bool] = &[];
let masks: [&[bool]; EMBED_SLOTS] = [&mask, empty, empty];
let batched = model.embed_chunk(&samples, &masks).expect("embed_chunk");
assert_eq!(
veneer, batched[0],
"embed_chunk_with_frame_mask must equal embed_chunk's slot 0 for the same mask"
);
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_with_frame_mask_matches_batched_slot_zero_v2() {
embed_chunk_with_frame_mask_matches_batched_slot_zero(embed_v2_path());
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_with_frame_mask_matches_batched_slot_zero_fp32() {
embed_chunk_with_frame_mask_matches_batched_slot_zero(embed_fp32_path());
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_with_frame_mask_rejects_all_false_mask() {
let model = load_embed_model(embed_v2_path());
let samples = synthetic_samples(crate::audio::speaker::segment::SEG_CHUNK_SAMPLES);
let mask = vec![false; model.num_mask_frames()];
let err = model
.embed_chunk_with_frame_mask(&samples, &mask)
.expect_err("all-false mask must be rejected");
assert_eq!(err, InferError::EmptyMask);
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_rejects_non_finite_samples() {
let model = load_embed_model(embed_v2_path());
let mut samples = synthetic_samples(crate::audio::speaker::segment::SEG_CHUNK_SAMPLES);
samples[1234] = f32::NAN;
let mask = vec![true; model.num_mask_frames()];
let masks: [&[bool]; EMBED_SLOTS] = [&mask, &mask, &mask];
let err = model
.embed_chunk(&samples, &masks)
.expect_err("NaN samples must be rejected");
assert_eq!(err, InferError::NonFiniteInput(1234));
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn embed_chunk_handles_short_padded_input() {
let model = load_embed_model(embed_v2_path());
let samples = synthetic_samples(40_000); let mask = vec![true; model.num_mask_frames()];
let masks: [&[bool]; EMBED_SLOTS] = [&mask, &mask, &mask];
let out = model
.embed_chunk(&samples, &masks)
.expect("embed_chunk on a short, repeat-padded chunk");
for row in out.iter() {
assert!(row.iter().all(|v| v.is_finite()));
}
}
use crate::{AxisRange, FeatureInfo, ModelDescription, ShapeConstraint, model::RawShapeConstraint};
fn fixed(name: &str, shape: &[usize], dtype: DataType) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
false,
Some(RawShapeConstraint::new(
2,
vec![shape.to_vec()],
shape.iter().map(|d| AxisRange::new(*d, 1)).collect(),
)),
)
}
fn ranged(name: &str, shape: &[usize], dtype: DataType, ranges: &[AxisRange]) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
false,
Some(RawShapeConstraint::new(3, Vec::new(), ranges.to_vec())),
)
}
fn wespeaker_description() -> ModelDescription {
ModelDescription::from_parts(
vec![
fixed(
names::MASK,
&[EMBED_SLOTS, MEASURED_MASK_FRAMES],
DataType::F32,
),
fixed(
names::WAVEFORM,
&[
EMBED_SLOTS,
crate::audio::speaker::segment::SEG_CHUNK_SAMPLES,
],
DataType::F32,
),
],
vec![
FeatureInfo::from_parts(
"constant".to_string(),
Vec::new(),
Some(DataType::F32),
false,
Some(RawShapeConstraint::new(1, Vec::new(), Vec::new())),
),
fixed(
names::EMBEDDING,
&[EMBED_SLOTS, EMBEDDING_DIM],
DataType::F32,
),
],
Vec::new(),
)
}
const MEASURED_MASK_FRAMES: usize = 589;
fn check(description: &ModelDescription) -> Result<(), ModelError> {
crate::model::contract::check_load_contract(description, &embed_contract())
.map_err(crate::audio::speaker::error::contract_violation)
}
#[test]
fn the_contract_accepts_the_staged_wespeaker_description() {
let description = wespeaker_description();
assert_eq!(check(&description), Ok(()));
assert_eq!(
description.input(names::MASK).expect("mask").shape()[1],
MEASURED_MASK_FRAMES
);
assert_eq!(
description
.output("constant")
.expect("constant")
.shape_constraint(),
Some(ShapeConstraint::Unspecified)
);
}
#[test]
fn the_contract_refuses_a_flexible_mask_whose_default_is_the_artifacts_589() {
let description = ModelDescription::from_parts(
vec![
ranged(
names::MASK,
&[EMBED_SLOTS, MEASURED_MASK_FRAMES],
DataType::F32,
&[
AxisRange::new(EMBED_SLOTS, 1),
AxisRange::inclusive(1, 4096),
],
),
fixed(
names::WAVEFORM,
&[
EMBED_SLOTS,
crate::audio::speaker::segment::SEG_CHUNK_SAMPLES,
],
DataType::F32,
),
],
vec![fixed(
names::EMBEDDING,
&[EMBED_SLOTS, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
assert_eq!(
description.input(names::MASK).expect("mask").shape()[1],
MEASURED_MASK_FRAMES
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, ModelError::ContractMismatch(m)
if m.feature() == names::MASK && m.actual() == "range" && m.expected() == "fixed"),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_zero_frame_mask() {
let description = ModelDescription::from_parts(
vec![
fixed(names::MASK, &[EMBED_SLOTS, 0], DataType::F32),
fixed(
names::WAVEFORM,
&[
EMBED_SLOTS,
crate::audio::speaker::segment::SEG_CHUNK_SAMPLES,
],
DataType::F32,
),
],
vec![fixed(
names::EMBEDDING,
&[EMBED_SLOTS, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, ModelError::ContractMismatch(m) if m.feature() == names::MASK),
"{err}"
);
}
#[test]
fn the_contract_refuses_an_extra_required_input() {
let mut inputs = wespeaker_description().inputs().to_vec();
inputs.push(fixed("speaker_ids", &[EMBED_SLOTS], DataType::I32));
let description = ModelDescription::from_parts(
inputs,
wespeaker_description().outputs().to_vec(),
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, ModelError::UnsatisfiableInput(name) if name == "speaker_ids"),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_graph_that_declares_state() {
let base = wespeaker_description();
let description = ModelDescription::from_parts(
base.inputs().to_vec(),
base.outputs().to_vec(),
vec![fixed("kv_cache", &[1, 8], DataType::F32)],
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, ModelError::UnsatisfiableState(name) if name == "kv_cache"),
"{err}"
);
}
#[test]
fn the_embed_contract_refuses_the_vendored_silero_bundle() {
let bundle = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../Models/vadkit/silero-vad-unified-256ms-v6.2.1.mlmodelc");
assert!(
bundle.is_dir(),
"the vendored silero bundle is committed, so this gate is NOT model-gated; looked for {}",
bundle.display()
);
let err = EmbedModel::from_file(&bundle).expect_err("silero is not this door's model");
assert!(
matches!(&err, ModelError::ContractMismatch(m)
if m.feature() == names::WAVEFORM && m.actual() == "missing"),
"{err}"
);
}
fn with_axis_bumped(base: &ModelDescription, feature: &str, axis: usize) -> ModelDescription {
let bump = |declared: &FeatureInfo| -> FeatureInfo {
if declared.name() != feature {
return declared.clone();
}
let mut shape = declared.shape().to_vec();
shape[axis] += 1;
fixed(
declared.name(),
&shape,
declared.data_type().expect("a multi-array feature"),
)
};
ModelDescription::from_parts(
base.inputs().iter().map(bump).collect(),
base.outputs().iter().map(bump).collect(),
base.states().to_vec(),
)
}
#[test]
fn every_axis_is_pinned_except_the_frame_count_the_door_reads_back() {
const FREE: &[(&str, usize)] = &[(names::MASK, 1)];
let base = wespeaker_description();
let named = [names::MASK, names::WAVEFORM, names::EMBEDDING];
let mut perturbations = 0_usize;
for declared in base.inputs().iter().chain(base.outputs()) {
if !named.contains(&declared.name()) {
continue;
}
for axis in 0..declared.shape().len() {
let perturbed = with_axis_bumped(&base, declared.name(), axis);
let free = FREE.contains(&(declared.name(), axis));
assert_eq!(
check(&perturbed).is_ok(),
free,
"`{}` axis {axis}: the contract {} it",
declared.name(),
if free { "must accept" } else { "must refuse" }
);
perturbations += 1;
}
}
assert_eq!(perturbations, 6);
}
#[test]
fn every_named_features_element_type_is_pinned() {
let base = wespeaker_description();
let named = [names::MASK, names::WAVEFORM, names::EMBEDDING];
let mut checked = 0_usize;
for declared in base.inputs().iter().chain(base.outputs()) {
if !named.contains(&declared.name()) {
continue;
}
let swap = |other: &FeatureInfo| -> FeatureInfo {
if other.name() == declared.name() {
fixed(other.name(), other.shape(), DataType::F16)
} else {
other.clone()
}
};
let perturbed = ModelDescription::from_parts(
base.inputs().iter().map(swap).collect(),
base.outputs().iter().map(swap).collect(),
base.states().to_vec(),
);
assert!(
matches!(check(&perturbed), Err(ModelError::ContractMismatch(m))
if m.feature() == declared.name()
&& m.expected() == "float32"
&& m.actual() == "float16"),
"`{}` re-declared float16 must be refused",
declared.name()
);
checked += 1;
}
assert_eq!(checked, 3);
}