Feed every input a model declares when probing a rung

The probe built one zero tensor from the first input and ran the session
with it. Every model so far had one input; the denoiser has two (mosaic and
sigma), so every rung failed with "Missing Input: sigma" and the role was
left on the CPU: 14.4 s for a 20 MP frame where TensorRT fp16 takes 3.1 s.
Zeros now go to each input by name.
This commit is contained in:
2026-10-03 11:15:49 -04:00
parent 6b0d29cc15
commit 20b7bd7663
2 changed files with 27 additions and 15 deletions
+24 -12
View File
@@ -168,21 +168,33 @@ fn time_rung(
started.elapsed().as_secs_f64() started.elapsed().as_secs_f64()
); );
let shape: Vec<usize> = session.inputs()[0] // Zeros for every input the model declares, by name — the denoiser
.dtype() // takes two (mosaic and σ), and a probe that fed only the first failed
.tensor_shape() // every rung and left it on the CPU.
.ok_or("model input is not a tensor")? let feeds: Vec<(String, Vec<usize>)> = session
.inputs()
.iter() .iter()
.map(|&d| if d > 0 { d as usize } else { 1 }) .map(|i| {
.collect(); let shape = i
let zeros = vec![0f32; shape.iter().product()]; .dtype()
.tensor_shape()
.ok_or("model input is not a tensor")?
.iter()
.map(|&d| if d > 0 { d as usize } else { 1 })
.collect();
Ok((i.name().to_string(), shape))
})
.collect::<Result<_, &str>>()?;
let run = |session: &mut ort::session::Session| -> Result<f64, String> { let run = |session: &mut ort::session::Session| -> Result<f64, String> {
let input = ort::value::Tensor::from_array((shape.clone(), zeros.clone())) let mut inputs: Vec<(String, ort::session::SessionInputValue)> = Vec::new();
.map_err(|e| e.to_string())?; for (name, shape) in &feeds {
let zeros = vec![0f32; shape.iter().product()];
let t = ort::value::Tensor::from_array((shape.clone(), zeros))
.map_err(|e| e.to_string())?;
inputs.push((name.clone(), t.into()));
}
let t = Instant::now(); let t = Instant::now();
let out = session let out = session.run(inputs).map_err(|e| e.to_string())?;
.run(ort::inputs![input])
.map_err(|e| e.to_string())?;
let _ = out[0] let _ = out[0]
.try_extract_tensor::<f32>() .try_extract_tensor::<f32>()
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
File diff suppressed because one or more lines are too long