use super::Tensor;
use crate::{ExecutionError, TensorData};
use alloc::vec::Vec;
use burn_backend::ops::TransactionPrimitive;
use burn_dispatch::Dispatch;
pub struct Transaction {
opaque: transaction_opaque::Opaque,
}
burn_std::obfuscate!(
type: TransactionPrimitive<Dispatch>,
module: transaction_opaque,
derives: [Send, Sync],
);
impl Default for Transaction {
fn default() -> Self {
Self::from_op(TransactionPrimitive::<Dispatch>::default())
}
}
impl Transaction {
pub(crate) fn from_op(op: TransactionPrimitive<Dispatch>) -> Self {
Self {
opaque: transaction_opaque::Opaque::new(op),
}
}
pub(crate) fn as_op_mut(&mut self) -> &mut TransactionPrimitive<Dispatch> {
self.opaque.as_mut()
}
pub(crate) fn into_op(self) -> TransactionPrimitive<Dispatch> {
self.opaque.into_inner()
}
pub fn register<const D: usize, K: crate::kind::Transaction>(
mut self,
tensor: Tensor<D, K>,
) -> Self {
K::register_transaction(self.as_op_mut(), tensor.primitive);
self
}
pub fn execute(self) -> Vec<TensorData> {
burn_std::future::block_on(self.execute_async())
.expect("Error while reading data: use `try_execute` to handle error at runtime")
}
pub fn try_execute(self) -> Result<Vec<TensorData>, ExecutionError> {
burn_std::future::block_on(self.execute_async())
}
pub async fn execute_async(self) -> Result<Vec<TensorData>, ExecutionError> {
self.into_op().execute_async().await
}
}