26 lines
866 B
Rust
26 lines
866 B
Rust
//! Factory functions for creating distributed tensors.
|
|
|
|
use crate::device_mesh::DeviceMesh;
|
|
use crate::error::Result;
|
|
use rtx_tensor::{Device, Tensor};
|
|
use std::sync::Arc;
|
|
|
|
use super::dtensor_core::DTensor;
|
|
use super::spec::TensorSpec;
|
|
|
|
/// Create DTensor filled with zeros.
|
|
pub fn zeros(_shape: Vec<usize>, spec: TensorSpec, mesh: Arc<DeviceMesh>) -> Result<DTensor> {
|
|
let local_shape = spec.local_shape(&mesh)?;
|
|
let device = Device::cpu();
|
|
let tensor = Tensor::zeros(&local_shape, &device)?;
|
|
DTensor::from_local(tensor, spec, mesh)
|
|
}
|
|
|
|
/// Create DTensor filled with ones.
|
|
pub fn ones(_shape: Vec<usize>, spec: TensorSpec, mesh: Arc<DeviceMesh>) -> Result<DTensor> {
|
|
let local_shape = spec.local_shape(&mesh)?;
|
|
let device = Device::cpu();
|
|
let tensor = Tensor::ones(&local_shape, &device)?;
|
|
DTensor::from_local(tensor, spec, mesh)
|
|
}
|