diff --git a/crates/bunsen-firehose-image/src/colortype_support.rs b/crates/bunsen-firehose-image/src/colortype_support.rs index 17c08d8..d3cddba 100644 --- a/crates/bunsen-firehose-image/src/colortype_support.rs +++ b/crates/bunsen-firehose-image/src/colortype_support.rs @@ -260,43 +260,43 @@ mod tests { assert_eq!( convert_to_colortype(source.clone(), ColorType::L8), - DynamicImage::from(source.clone().to_luma8()) + DynamicImage::from(source.to_luma8()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::La8), - DynamicImage::from(source.clone().to_luma_alpha8()) + DynamicImage::from(source.to_luma_alpha8()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::Rgb8), - DynamicImage::from(source.clone().to_rgb8()) + DynamicImage::from(source.to_rgb8()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::Rgba8), - DynamicImage::from(source.clone().to_rgba8()) + DynamicImage::from(source.to_rgba8()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::L16), - DynamicImage::from(source.clone().to_luma16()) + DynamicImage::from(source.to_luma16()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::La16), - DynamicImage::from(source.clone().to_luma_alpha16()) + DynamicImage::from(source.to_luma_alpha16()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::Rgb16), - DynamicImage::from(source.clone().to_rgb16()) + DynamicImage::from(source.to_rgb16()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::Rgba16), - DynamicImage::from(source.clone().to_rgba16()) + DynamicImage::from(source.to_rgba16()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::Rgb32F), - DynamicImage::from(source.clone().to_rgb32f()) + DynamicImage::from(source.to_rgb32f()) ); assert_eq!( convert_to_colortype(source.clone(), ColorType::Rgba32F), - DynamicImage::from(source.clone().to_rgba32f()) + DynamicImage::from(source.to_rgba32f()) ); } } diff --git a/crates/bunsen/src/blocks/images/drop/drop_path.rs b/crates/bunsen/src/blocks/images/drop/drop_path.rs index 73e5da0..d5d3cbc 100644 --- a/crates/bunsen/src/blocks/images/drop/drop_path.rs +++ b/crates/bunsen/src/blocks/images/drop/drop_path.rs @@ -282,7 +282,7 @@ mod tests { Tensor::::from_data([[[1.0]], [[0.0]], [[1.0]]], device) }, ); - res.to_data().assert_eq(&x.clone().to_data(), true); + res.to_data().assert_eq(&x.to_data(), true); // No-op case: training, but drop_prob = 0.0 let training = true; @@ -299,7 +299,7 @@ mod tests { Tensor::::from_data([[[1.0]], [[0.0]], [[1.0]]], device) }, ); - res.to_data().assert_eq(&x.clone().to_data(), true); + res.to_data().assert_eq(&x.to_data(), true); // Training, but no scaling let training = true; diff --git a/crates/bunsen/src/blocks/rnn/lstm/ext_lstm_state.rs b/crates/bunsen/src/blocks/rnn/lstm/ext_lstm_state.rs index 9497297..5469787 100644 --- a/crates/bunsen/src/blocks/rnn/lstm/ext_lstm_state.rs +++ b/crates/bunsen/src/blocks/rnn/lstm/ext_lstm_state.rs @@ -267,8 +267,8 @@ mod tests { let device = Default::default(); let state: ExtLstmState = random_state([2, 3, 4], &device); - let expected_cell = state.cell.clone().to_data(); - let expected_hidden = state.hidden.clone().to_data(); + let expected_cell = state.cell.to_data(); + let expected_hidden = state.hidden.to_data(); let (cell, hidden) = state.unpack(); @@ -354,7 +354,7 @@ mod tests { let device = Default::default(); let state: ExtLstmState = random_state([2, 3], &device); - let expected_cell = state.cell.clone().to_data(); + let expected_cell = state.cell.to_data(); let roundtrip: ExtLstmState = state.unsqueeze_dim::<3>(1).squeeze_dim(1); @@ -384,8 +384,8 @@ mod tests { let device = Default::default(); let state: ExtLstmState = random_state([2, 3], &device); - let expected_cell = state.cell.clone().to_data(); - let expected_hidden = state.hidden.clone().to_data(); + let expected_cell = state.cell.to_data(); + let expected_hidden = state.hidden.to_data(); // The shape passed here is intentionally different; it must be ignored. let unwrapped = Some(state).unwrap_or_initial([9, 9], &device); diff --git a/crates/bunsen/src/burner/tensor/tensor_op_ext.rs b/crates/bunsen/src/burner/tensor/tensor_op_ext.rs index e6ab0b2..6294e33 100644 --- a/crates/bunsen/src/burner/tensor/tensor_op_ext.rs +++ b/crates/bunsen/src/burner/tensor/tensor_op_ext.rs @@ -12,6 +12,7 @@ use burn::{ BasicOps, Bool, DType, + DataError, Element, Float, Int, @@ -80,13 +81,19 @@ where K: BasicOps, K::Elem: Element, { - /// Copies the current `Tensor` into `TensorData`; converts the dtype. + /// Copies the current `Tensor` into a `TensorData`; converts the dtype. + /// + /// By contract, this will yield the same result as + /// `tensor.to_data().convert::(dtype)`. /// /// The conversion is a no-op if the dtype is the same as the current dtype. - fn to_data_convert(&self) -> TensorData; + fn to_data_as(&self) -> TensorData; /// Copies the current `Tensor` into `TensorData`; converts the dtype. /// + /// By contract, this will yield the same result as + /// `tensor.to_data().convert_dtype(dtype)`. + /// /// The conversion is a no-op if the dtype is the same as the current dtype. fn to_data_cast( &self, @@ -95,11 +102,17 @@ where /// Converts the current `Tensor` into `TensorData`; converts the dtype. /// + /// By contract, this will yield the same result as + /// `tensor.into_data().convert::(dtype)`. + /// /// The conversion is a no-op if the dtype is the same as the current dtype. - fn into_data_convert(self) -> TensorData; + fn into_data_as(self) -> TensorData; /// Converts the current `Tensor` into `TensorData`; converts the dtype. /// + /// By contract, this will yield the same result as + /// `tensor.into_data().convert_dtype(dtype)`. + /// /// The conversion is a no-op if the dtype is the same as the current dtype. fn into_data_cast( self, @@ -113,8 +126,8 @@ where K: BasicOps, K::Elem: Element, { - fn to_data_convert(&self) -> TensorData { - self.to_data().convert::() + fn to_data_as(&self) -> TensorData { + self.to_data_cast(E::dtype()) } fn to_data_cast( @@ -124,8 +137,8 @@ where self.to_data().convert_dtype(dtype) } - fn into_data_convert(self) -> TensorData { - self.into_data().convert::() + fn into_data_as(self) -> TensorData { + self.into_data_cast(E::dtype()) } fn into_data_cast( @@ -278,6 +291,74 @@ where } } +/// Extension trait for `TensorData` that provides additional methods. +pub trait TensorDataToVecAsExt { + /// Cast the data to a new dtype. + /// + /// TODO: Implement proper error handling in `TensorData`. + /// + /// # Returns + /// Ok(data) on success, (Currently) panics on failure. + fn try_cast( + self, + dtype: DType, + ) -> Result; + + /// Convert the data to a new dtype. + /// + /// TODO: Implement proper error handling in `TensorData`. + /// + /// By contract, this is equivalent to: + /// `data.try_cast(E::dtype())` + /// + /// # Returns + /// Ok(data) on success, (Currently) panics on failure. + fn try_convert(self) -> Result; + + /// Copy and convert the data to a [`Vec`]. + /// + /// By contract, this is equivalent to: + /// `data.clone().into_vec_as::()` + /// + /// Particular conversions may provide more efficient implementations. + /// + /// # Returns + /// `Ok(vec)` on success, or an error if the conversion fails. + fn to_vec_as(&self) -> Result, DataError>; + + /// Convert the data to [`Vec`]. + /// + /// By contract, this is equivalent to: + /// `data.try_convert::()?.to_vec::()` + /// + /// Particular conversions may provide more efficient implementations. + /// + /// # Returns + /// `Ok(vec)` on success, or an error if the conversion fails. + fn into_vec_as(self) -> Result, DataError>; +} + +impl TensorDataToVecAsExt for TensorData { + fn try_cast( + self, + dtype: DType, + ) -> Result { + Ok(self.convert_dtype(dtype)) + } + + fn try_convert(self) -> Result { + self.try_cast(E::dtype()) + } + + fn to_vec_as(&self) -> Result, DataError> { + self.clone().into_vec_as::() + } + + fn into_vec_as(self) -> Result, DataError> { + self.try_convert::()?.to_vec::() + } +} + #[cfg(test)] mod tests { use burn::tensor::{ @@ -381,4 +462,41 @@ mod tests { .to_data() .assert_eq(&TensorData::from([3]), false); } + + #[test] + fn test_to_vec_as() { + let data = TensorData::from([0.0f32, 1.0, 2.5]); + + // Same-dtype copy. + assert_eq!(data.to_vec_as::().unwrap(), vec![0.0f32, 1.0, 2.5]); + + // Widening cast (different element size). + assert_eq!(data.to_vec_as::().unwrap(), vec![0.0f64, 1.0, 2.5]); + + // Float to int cast (same element size) truncates. + assert_eq!(data.to_vec_as::().unwrap(), vec![0i32, 1, 2]); + + // The source data is borrowed, not consumed. + data.assert_eq(&TensorData::from([0.0f32, 1.0, 2.5]), true); + } + + #[test] + fn test_into_vec_as() { + let data = TensorData::from([0i32, 1, 2, 3]); + + // Same-dtype conversion. + assert_eq!( + data.clone().into_vec_as::().unwrap(), + vec![0i32, 1, 2, 3] + ); + + // Int to float cast. + assert_eq!( + data.clone().into_vec_as::().unwrap(), + vec![0.0f32, 1.0, 2.0, 3.0] + ); + + // Narrowing int cast. + assert_eq!(data.into_vec_as::().unwrap(), vec![0u8, 1, 2, 3]); + } } diff --git a/crates/bunsen/src/kits/bimm/resnet/blocks/layer_block.rs b/crates/bunsen/src/kits/bimm/resnet/blocks/layer_block.rs index 2d8bf7d..8df5ff5 100644 --- a/crates/bunsen/src/kits/bimm/resnet/blocks/layer_block.rs +++ b/crates/bunsen/src/kits/bimm/resnet/blocks/layer_block.rs @@ -467,6 +467,7 @@ mod tests { use crate::{ contracts::assert_shape_contract, kits::bimm::resnet::blocks::BasicBlockConfig, + prelude::*, support::testing::PerformanceBackend, }; @@ -556,7 +557,7 @@ mod tests { expected = block.forward(expected); } output - .to_data() - .assert_approx_eq::(&expected.to_data(), Tolerance::default()); + .to_data_as::() + .assert_approx_eq::(&expected.to_data_as::(), Tolerance::default()); } } diff --git a/crates/bunsen/src/kits/bimm/swin/v2/blocks/window_attention/pos_grid.rs b/crates/bunsen/src/kits/bimm/swin/v2/blocks/window_attention/pos_grid.rs index 5335758..dc886fe 100644 --- a/crates/bunsen/src/kits/bimm/swin/v2/blocks/window_attention/pos_grid.rs +++ b/crates/bunsen/src/kits/bimm/swin/v2/blocks/window_attention/pos_grid.rs @@ -284,7 +284,7 @@ mod tests { let device = Default::default(); let rel = window_attention_relative_position_index::(window_shape, &device); - rel.clone().to_data().assert_eq( + rel.to_data().assert_eq( &TensorData::from([ [7, 6, 5, 2, 1, 0], [8, 7, 6, 3, 2, 1], diff --git a/crates/bunsen/src/kits/sims/conway/life2d.rs b/crates/bunsen/src/kits/sims/conway/life2d.rs index 38ca63d..581d3af 100644 --- a/crates/bunsen/src/kits/sims/conway/life2d.rs +++ b/crates/bunsen/src/kits/sims/conway/life2d.rs @@ -88,7 +88,7 @@ where state .slice(slices) - .to_data_convert::() + .to_data_as::() .to_vec() .unwrap() .chunks(w) diff --git a/crates/bunsen/src/kits/sims/lbm/d2q9/collision.rs b/crates/bunsen/src/kits/sims/lbm/d2q9/collision.rs index 6861c1a..9e273a3 100644 --- a/crates/bunsen/src/kits/sims/lbm/d2q9/collision.rs +++ b/crates/bunsen/src/kits/sims/lbm/d2q9/collision.rs @@ -119,7 +119,7 @@ mod tests { // Invariant: density(collision(dist, param)) == density(dist) density(col_dist.clone()) .to_data() - .assert_approx_eq::(&rho.clone().to_data(), Tolerance::default()); + .assert_approx_eq::(&rho.to_data(), Tolerance::default()); // With Correction { diff --git a/crates/bunsen/src/kits/sims/lbm/d2q9/space.rs b/crates/bunsen/src/kits/sims/lbm/d2q9/space.rs index 83da824..b8332e3 100644 --- a/crates/bunsen/src/kits/sims/lbm/d2q9/space.rs +++ b/crates/bunsen/src/kits/sims/lbm/d2q9/space.rs @@ -395,7 +395,7 @@ mod tests { let momentum = macroscopic_momentum(dist.clone(), lbm_tables.e_vec()); - momentum.clone().to_data().assert_approx_eq::( + momentum.to_data().assert_approx_eq::( &Tensor::::from_data([[[-1., -1.], [15., 5.]]], &device).to_data(), Tolerance::default(), ); @@ -404,7 +404,7 @@ mod tests { let rho_data = rho.to_data().to_vec::().unwrap(); - u.clone().to_data().assert_approx_eq::( + u.to_data().assert_approx_eq::( &Tensor::::from_data( [[ [-1. / rho_data[0], -1. / rho_data[0]], @@ -418,7 +418,7 @@ mod tests { let v_sq = velocity_squared(u.clone()); - v_sq.clone().to_data().assert_approx_eq::( + v_sq.to_data().assert_approx_eq::( &Tensor::::from_data( [[ (1. + 1.) / rho_data[0].powi(2), diff --git a/crates/bunsen/src/kits/sims/lbm/d2q9/thermal.rs b/crates/bunsen/src/kits/sims/lbm/d2q9/thermal.rs index ff86693..82b94dc 100644 --- a/crates/bunsen/src/kits/sims/lbm/d2q9/thermal.rs +++ b/crates/bunsen/src/kits/sims/lbm/d2q9/thermal.rs @@ -165,7 +165,7 @@ mod tests { let parts = ldv_projection(e.clone(), u.clone()); - parts.clone().to_data().assert_approx_eq::( + parts.to_data().assert_approx_eq::( &Tensor::::from_data( [[ [ @@ -187,7 +187,7 @@ mod tests { let e_u = lattice_dot_velocity(u.clone(), e.clone()); - e_u.clone().to_data().assert_approx_eq::( + e_u.to_data().assert_approx_eq::( &parts.sum_dim(4).squeeze_dims::<4>(&[4]).to_data(), Tolerance::default(), ); diff --git a/crates/bunsen/src/kits/speech/silero_vad/blocks/module.rs b/crates/bunsen/src/kits/speech/silero_vad/blocks/module.rs index df68689..a2e3c76 100644 --- a/crates/bunsen/src/kits/speech/silero_vad/blocks/module.rs +++ b/crates/bunsen/src/kits/speech/silero_vad/blocks/module.rs @@ -858,7 +858,10 @@ mod tests { }; use super::*; - use crate::support::testing::PerformanceBackend; + use crate::{ + prelude::*, + support::testing::PerformanceBackend, + }; type B = PerformanceBackend; @@ -1060,11 +1063,11 @@ mod tests { let tol = Tolerance::::default(); seq_probs - .into_data() - .assert_approx_eq::(&step_probs.into_data(), tol); + .into_data_as::() + .assert_approx_eq::(&step_probs.into_data_as::(), tol); seq_state - .into_data() - .assert_approx_eq::(&state.into_data(), tol); + .into_data_as::() + .assert_approx_eq::(&state.into_data_as::(), tol); } #[test] diff --git a/crates/bunsen/src/kits/speech/silero_vad/cross_test.rs b/crates/bunsen/src/kits/speech/silero_vad/cross_test.rs index 13eb963..3bd2262 100644 --- a/crates/bunsen/src/kits/speech/silero_vad/cross_test.rs +++ b/crates/bunsen/src/kits/speech/silero_vad/cross_test.rs @@ -26,6 +26,7 @@ mod tests { SileroVadMeta, reference::ReferenceModel, }, + prelude::*, support::testing::PerformanceBackend, }; @@ -65,12 +66,12 @@ mod tests { s_out .reshape([batch, 1]) - .to_data() - .assert_approx_eq::(&r_out.to_data(), Tolerance::permissive()); + .to_data_as::() + .assert_approx_eq::(&r_out.to_data_as::(), Tolerance::permissive()); s_state - .to_data() - .assert_approx_eq::(&r_state.to_data(), Tolerance::permissive()); + .to_data_as::() + .assert_approx_eq::(&r_state.to_data_as::(), Tolerance::permissive()); } } diff --git a/crates/bunsen/src/kits/speech/ten_vad/cross_test.rs b/crates/bunsen/src/kits/speech/ten_vad/cross_test.rs index d36a504..42c4a2b 100644 --- a/crates/bunsen/src/kits/speech/ten_vad/cross_test.rs +++ b/crates/bunsen/src/kits/speech/ten_vad/cross_test.rs @@ -13,6 +13,7 @@ mod test { TenVad, reference::ReferenceModel, }, + prelude::*, support::testing::PerformanceBackend, }; @@ -50,31 +51,27 @@ mod test { mod_prob .unsqueeze_dim::<3>(2) - .to_data() - .assert_approx_eq::(&ref_prob.to_data(), Tolerance::permissive()); + .to_data_as::() + .assert_approx_eq::(&ref_prob.to_data_as::(), Tolerance::permissive()); mod_lstm1_state .hidden - .clone() - .to_data() - .assert_approx_eq::(&ref_lstm1_hidden.to_data(), Tolerance::permissive()); + .to_data_as::() + .assert_approx_eq::(&ref_lstm1_hidden.to_data_as::(), Tolerance::permissive()); mod_lstm1_state .cell - .clone() - .to_data() - .assert_approx_eq::(&ref_lstm1_cell.to_data(), Tolerance::permissive()); + .to_data_as::() + .assert_approx_eq::(&ref_lstm1_cell.to_data_as::(), Tolerance::permissive()); mod_lstm2_state .hidden - .clone() - .to_data() - .assert_approx_eq::(&ref_lstm2_hidden.to_data(), Tolerance::permissive()); + .to_data_as::() + .assert_approx_eq::(&ref_lstm2_hidden.to_data_as::(), Tolerance::permissive()); mod_lstm2_state .cell - .clone() - .to_data() - .assert_approx_eq::(&ref_lstm2_cell.to_data(), Tolerance::permissive()); + .to_data_as::() + .assert_approx_eq::(&ref_lstm2_cell.to_data_as::(), Tolerance::permissive()); } } diff --git a/crates/bunsen/src/kits/speech/whisper/blocks/decoder_block.rs b/crates/bunsen/src/kits/speech/whisper/blocks/decoder_block.rs index aab7a1b..ca4e5f5 100644 --- a/crates/bunsen/src/kits/speech/whisper/blocks/decoder_block.rs +++ b/crates/bunsen/src/kits/speech/whisper/blocks/decoder_block.rs @@ -281,12 +281,12 @@ mod tests { .output .clone() .into_data() - .assert_approx_eq::(&expected.output.clone().into_data(), Default::default()); + .assert_approx_eq::(&expected.output.to_data(), Default::default()); result .ca_weights .clone() .into_data() - .assert_approx_eq::(&expected.ca_weights.clone().into_data(), Default::default()); + .assert_approx_eq::(&expected.ca_weights.to_data(), Default::default()); assert_shape_contract!( ["batch", "seq_len", "d_model"], diff --git a/crates/bunsen/src/ops/arange.rs b/crates/bunsen/src/ops/arange.rs index 2eb6188..dc1efb4 100644 --- a/crates/bunsen/src/ops/arange.rs +++ b/crates/bunsen/src/ops/arange.rs @@ -134,9 +134,12 @@ mod tests { }; use super::*; - use crate::support::testing::{ - CpuBackend, - assert_close_to_vec, + use crate::{ + prelude::*, + support::testing::{ + CpuBackend, + assert_close_to_vec, + }, }; type B = CpuBackend; type F = ::FloatElem; @@ -248,7 +251,7 @@ mod tests { let actual = tensor_linspace::(start, end, num, &device); - actual.to_data().assert_approx_eq::( + actual.to_data_as::().assert_approx_eq::( &TensorData::from([0.0, 0.25, 0.5, 0.75, 1.0]), Tolerance::default(), ); @@ -264,7 +267,7 @@ mod tests { let actual = tensor_linspace::(start, end, num, &device); - actual.to_data().assert_approx_eq::( + actual.to_data_as::().assert_approx_eq::( &TensorData::from([1.0, 0.7, 0.4, 0.1, -0.2]), Tolerance::default(), ); @@ -281,7 +284,7 @@ mod tests { let actual = tensor_linspace::(start, end, num, &device); actual - .to_data() + .to_data_as::() .assert_approx_eq::(&TensorData::from([0.0]), Tolerance::default()); } } diff --git a/crates/bunsen/src/ops/signal/cosine_window.rs b/crates/bunsen/src/ops/signal/cosine_window.rs index dc4932c..5477c25 100644 --- a/crates/bunsen/src/ops/signal/cosine_window.rs +++ b/crates/bunsen/src/ops/signal/cosine_window.rs @@ -235,6 +235,7 @@ mod tests { use super::*; use crate::{ ops::signal::testing::assert_builder_impls_match, + prelude::*, support::testing::CpuBackend, }; @@ -255,8 +256,11 @@ mod tests { info!("checking hann_window reference implementation"); hann_window::(size, periodic, options.clone()) - .to_data() - .assert_approx_eq::(&TensorData::from(expected), Tolerance::default()); + .to_data_as::() + .assert_approx_eq::( + &TensorData::from(expected).convert::(), + Tolerance::default(), + ); info!("cross-checking vec/tensor impls"); assert_builder_impls_match::(&cfg, expected, options.clone()); @@ -312,8 +316,11 @@ mod tests { info!("checking blackman_window reference implementation"); blackman_window::(size, periodic, options.clone()) - .to_data() - .assert_approx_eq::(&TensorData::from(expected), Tolerance::default()); + .to_data_as::() + .assert_approx_eq::( + &TensorData::from(expected).convert::(), + Tolerance::default(), + ); info!("cross-checking vec/tensor impls"); assert_builder_impls_match::(&cfg, expected, options.clone()); diff --git a/crates/bunsen/src/ops/signal/sliding_stft.rs b/crates/bunsen/src/ops/signal/sliding_stft.rs index c40e283..7dd9753 100644 --- a/crates/bunsen/src/ops/signal/sliding_stft.rs +++ b/crates/bunsen/src/ops/signal/sliding_stft.rs @@ -474,7 +474,10 @@ mod tests { }; use super::*; - use crate::support::testing::CpuBackend; + use crate::{ + prelude::*, + support::testing::CpuBackend, + }; type B = CpuBackend; type F = ::FloatElem; @@ -662,10 +665,12 @@ mod tests { .flat_map(|(host, row)| host.push(row)) .collect(); - out.cast(DType::F64).to_data().assert_approx_eq::( - &TensorData::new(expected, [batch, n_bins, 2]), - Tolerance::permissive(), - ); + out.cast(DType::F64) + .to_data_as::() + .assert_approx_eq::( + &TensorData::new(expected, [batch, n_bins, 2]).convert::(), + Tolerance::permissive(), + ); } } @@ -707,12 +712,12 @@ mod tests { let tol = Tolerance::::permissive(); seq_out - .to_data() - .assert_approx_eq::(&step_out.to_data(), tol); + .to_data_as::() + .assert_approx_eq::(&step_out.to_data_as::(), tol); seq_stft .queue - .to_data() - .assert_approx_eq::(&step_stft.queue.to_data(), tol); + .to_data_as::() + .assert_approx_eq::(&step_stft.queue.to_data_as::(), tol); } } @@ -729,8 +734,9 @@ mod tests { stft.forward(hop); stft.reset(); - stft.queue - .to_data() - .assert_eq(&Tensor::::zeros([1, 48], &device).to_data(), true); + stft.queue.to_data_as::().assert_eq( + &Tensor::::zeros([1, 48], &device).to_data_as::(), + true, + ); } } diff --git a/examples/conway_vis/src/main.rs b/examples/conway_vis/src/main.rs index 77cdcee..c218f67 100644 --- a/examples/conway_vis/src/main.rs +++ b/examples/conway_vis/src/main.rs @@ -246,7 +246,7 @@ impl Simulation { export_duration: Duration, ) -> Self { let shutdown = Arc::new(AtomicBool::new(false)); - let frame_handle_1 = Arc::new(Mutex::new(conway.state.clone().into_data())); + let frame_handle_1 = Arc::new(Mutex::new(conway.state.to_data())); let frame_handle_2 = frame_handle_1.clone(); let shutdown_clone = shutdown.clone(); @@ -284,7 +284,7 @@ impl Simulation { if t1 - last_export > export_duration { last_export = t1; - let frame = conway.state.clone().into_data_convert::(); + let frame = conway.state.to_data_as::(); *frame_handle_1.lock().unwrap() = frame; t1 = std::time::Instant::now(); diff --git a/examples/lbm2d_vis/src/main.rs b/examples/lbm2d_vis/src/main.rs index 2cac0d9..b379b5d 100644 --- a/examples/lbm2d_vis/src/main.rs +++ b/examples/lbm2d_vis/src/main.rs @@ -190,7 +190,7 @@ fn run( let cells = ((cells / scale) + 1.0) / 2.0; // let cells = cells.mul_scalar(std::f64::consts::PI / 2.0).sin(); - *vis_cells_publish.lock().unwrap() = cells.to_data_convert::(); + *vis_cells_publish.lock().unwrap() = cells.to_data_as::(); last_export = std::time::Instant::now(); } @@ -212,7 +212,7 @@ fn run( let mut app = FlowVisApp { gl: GlGraphics::new(opengl), cell_data: vis_cells, - solid_mask: solid_mask.to_data_convert::(), + solid_mask: solid_mask.to_data_as::(), opacity: args.opacity, };