pub fn with_batch_policy<R>(
session: &mut dyn BackendSession,
policy: CpuBatchPolicy,
f: impl FnOnce(&mut dyn BackendSession) -> R,
) -> Result<R>Expand description
Run f on session with policy as the effective batch policy.
This is the scoped override of the precedence per-operation > scoped >
backend default; wrapping a single call expresses a per-operation choice.
The previous policy is restored when f returns, returns an error or
unwinds, and scopes nest with the innermost winning.
§Examples
use tenferro_cpu::{with_batch_policy, CpuBackend, CpuBatchPolicy, CpuBatchStrategy};
use tenferro_tensor::{BackendSessionHost, DotGeneralConfig, Tensor, TensorRead};
let mut backend = CpuBackend::with_threads(1)?;
let lhs = Tensor::from_vec_col_major(vec![2, 2, 3], vec![1.0_f64; 12])?;
let rhs = Tensor::from_vec_col_major(vec![2, 2, 3], vec![2.0_f64; 12])?;
let config = DotGeneralConfig {
lhs_contracting_dims: [1].as_slice().into(),
rhs_contracting_dims: [0].as_slice().into(),
lhs_batch_dims: [2].as_slice().into(),
rhs_batch_dims: [2].as_slice().into(),
};
let product = backend.with_backend_session(|session| {
// ProviderItems runs one provider GEMM per item with every provider;
// a forced Sequential is rejected by providers with their own threading.
with_batch_policy(session, CpuBatchPolicy::new(CpuBatchStrategy::ProviderItems), |session| {
session.dot_general_read(
TensorRead::from_tensor(&lhs),
TensorRead::from_tensor(&rhs),
&config,
)
})
})???;
assert_eq!(product.as_slice::<f64>()?, &[4.0; 12]);§Errors
Returns tenferro_tensor::Error::Unsupported without running f when
session is neither a CPU execution session nor a session that forwards
one through tenferro_tensor::BackendSession::native_session (a CUDA or
other custom session). f’s own result is returned unchanged inside Ok.
§Panics
A panic in f propagates after the previous policy is restored.