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:
@@ -168,21 +168,33 @@ fn time_rung(
|
||||
started.elapsed().as_secs_f64()
|
||||
);
|
||||
|
||||
let shape: Vec<usize> = session.inputs()[0]
|
||||
.dtype()
|
||||
.tensor_shape()
|
||||
.ok_or("model input is not a tensor")?
|
||||
// Zeros for every input the model declares, by name — the denoiser
|
||||
// takes two (mosaic and σ), and a probe that fed only the first failed
|
||||
// every rung and left it on the CPU.
|
||||
let feeds: Vec<(String, Vec<usize>)> = session
|
||||
.inputs()
|
||||
.iter()
|
||||
.map(|&d| if d > 0 { d as usize } else { 1 })
|
||||
.collect();
|
||||
let zeros = vec![0f32; shape.iter().product()];
|
||||
.map(|i| {
|
||||
let shape = i
|
||||
.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 input = ort::value::Tensor::from_array((shape.clone(), zeros.clone()))
|
||||
.map_err(|e| e.to_string())?;
|
||||
let mut inputs: Vec<(String, ort::session::SessionInputValue)> = Vec::new();
|
||||
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 out = session
|
||||
.run(ort::inputs![input])
|
||||
.map_err(|e| e.to_string())?;
|
||||
let out = session.run(inputs).map_err(|e| e.to_string())?;
|
||||
let _ = out[0]
|
||||
.try_extract_tensor::<f32>()
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Reference in New Issue
Block a user