Files
rustytorch/crates/training/rtx-distributed/src/dtensor/factory.rs
T
2026-03-04 00:08:42 +00:00

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)
}