Skip to main content

ruccl/in_process/collective/
memory.rs

1use super::*;
2
3impl<B, T, E> InProcessCollective<'_, T, crate::tensor_device::TensorDevice<B>, E>
4where
5    B: ruda_tensor::Backend,
6    T: crate::tensor_device::TensorElement,
7    E: From<crate::tensor_device::TensorDeviceError> + From<InProcessError>,
8{
9    /// Attach existing rank tensors without downloading and re-uploading values.
10    pub fn import_buffers(
11        &self,
12        buffers: Vec<crate::tensor_device::TensorBuffer<B, T>>,
13    ) -> Result<DistributedBuffer<T, crate::tensor_device::TensorDevice<B>>, E> {
14        if buffers.len() != self.world_size() {
15            return Err(InProcessError::InvalidLength("one tensor buffer is required per rank").into());
16        }
17        let length_per_rank = buffers.first().ok_or(InProcessError::EmptyWorld)?.len();
18        for (buffer, context) in buffers.iter().zip(self.contexts) {
19            if buffer.len() != length_per_rank {
20                return Err(InProcessError::InvalidLength("rank tensor lengths must match").into());
21            }
22            if buffer.device() != context.device() {
23                return Err(crate::tensor_device::TensorDeviceError::DeviceMismatch.into());
24            }
25        }
26        Ok(DistributedBuffer {
27            communicator_id: self.id,
28            marker: PhantomData,
29            buffers,
30            length_per_rank,
31        })
32    }
33}
34
35impl<T, D, E> InProcessCollective<'_, T, D, E>
36where
37    T: Copy + Send + Sync + 'static,
38    D: InProcessDevice<T>,
39    E: From<D::Error> + From<InProcessError>,
40{
41    pub fn allocate(&self, length_per_rank: usize) -> Result<DistributedBuffer<T, D>, E> {
42        let buffers = self
43            .contexts
44            .iter()
45            .map(|context| <D as InProcessDevice<T>>::alloc(context, length_per_rank))
46            .collect::<Result<Vec<_>, _>>()?;
47        Ok(DistributedBuffer {
48            communicator_id: self.id,
49            marker: PhantomData,
50            buffers,
51            length_per_rank,
52        })
53    }
54
55    pub fn allocate_rooted(&self, root: usize, length: usize) -> Result<RootedBuffer<T, D>, E> {
56        let context = self.context(root)?;
57        Ok(RootedBuffer {
58            communicator_id: self.id,
59            marker: PhantomData,
60            root,
61            buffer: <D as InProcessDevice<T>>::alloc(context, length)?,
62        })
63    }
64
65    pub fn allocate_variable(
66        &self,
67        lengths: &[usize],
68    ) -> Result<VariableDistributedBuffer<T, D>, E> {
69        if lengths.len() != self.world_size() {
70            return Err(InProcessError::InvalidLength(
71                "variable buffer lengths must contain one entry per rank",
72            )
73            .into());
74        }
75        let buffers = self
76            .contexts
77            .iter()
78            .zip(lengths)
79            .map(|(context, length)| {
80                if *length == 0 {
81                    Ok(None)
82                } else {
83                    <D as InProcessDevice<T>>::alloc(context, *length).map(Some)
84                }
85            })
86            .collect::<Result<Vec<_>, _>>()?;
87        Ok(VariableDistributedBuffer {
88            communicator_id: self.id,
89            marker: PhantomData,
90            buffers,
91            lengths: lengths.to_vec(),
92        })
93    }
94
95    pub fn upload_variable_rank(
96        &self,
97        buffer: &VariableDistributedBuffer<T, D>,
98        rank: usize,
99        values: &[T],
100    ) -> Result<(), E> {
101        self.validate_variable_buffer(buffer)?;
102        let context = self.context(rank)?;
103        let expected = buffer
104            .rank_length(rank)
105            .ok_or(InProcessError::RankOutOfRange {
106                rank,
107                world_size: self.world_size(),
108            })?;
109        if values.len() != expected {
110            return Err(InProcessError::InvalidLength(
111                "variable rank upload length does not match allocation",
112            )
113            .into());
114        }
115        if let Some(rank_buffer) = buffer.rank_buffer(rank) {
116            <D as InProcessDevice<T>>::copy_to_device(context, rank_buffer, values)?;
117        }
118        Ok(())
119    }
120
121    pub fn download_variable_rank(
122        &self,
123        buffer: &VariableDistributedBuffer<T, D>,
124        rank: usize,
125    ) -> Result<Vec<T>, E> {
126        self.validate_variable_buffer(buffer)?;
127        let context = self.context(rank)?;
128        match buffer.rank_buffer(rank) {
129            Some(rank_buffer) => Ok(<D as InProcessDevice<T>>::copy_from_device(
130                context,
131                rank_buffer,
132            )?),
133            None if buffer.rank_length(rank) == Some(0) => Ok(Vec::new()),
134            None => Err(InProcessError::RankOutOfRange {
135                rank,
136                world_size: self.world_size(),
137            }
138            .into()),
139        }
140    }
141
142    pub fn upload_root(&self, buffer: &RootedBuffer<T, D>, values: &[T]) -> Result<(), E> {
143        self.validate_rooted_buffer(buffer)?;
144        <D as InProcessDevice<T>>::copy_to_device(
145            self.context(buffer.root)?,
146            &buffer.buffer,
147            values,
148        )?;
149        Ok(())
150    }
151
152    pub fn download_root(&self, buffer: &RootedBuffer<T, D>) -> Result<Vec<T>, E> {
153        self.validate_rooted_buffer(buffer)?;
154        Ok(<D as InProcessDevice<T>>::copy_from_device(
155            self.context(buffer.root)?,
156            &buffer.buffer,
157        )?)
158    }
159
160    pub fn upload_rank(
161        &self,
162        buffer: &DistributedBuffer<T, D>,
163        rank: usize,
164        values: &[T],
165    ) -> Result<(), E> {
166        self.validate_buffer(buffer)?;
167        let context = self.context(rank)?;
168        let rank_buffer = buffer
169            .buffers
170            .get(rank)
171            .ok_or(InProcessError::RankOutOfRange {
172                rank,
173                world_size: self.world_size(),
174            })?;
175        <D as InProcessDevice<T>>::copy_to_device(context, rank_buffer, values)?;
176        Ok(())
177    }
178
179    pub fn download_rank(
180        &self,
181        buffer: &DistributedBuffer<T, D>,
182        rank: usize,
183    ) -> Result<Vec<T>, E> {
184        self.validate_buffer(buffer)?;
185        let context = self.context(rank)?;
186        let rank_buffer = buffer
187            .buffers
188            .get(rank)
189            .ok_or(InProcessError::RankOutOfRange {
190                rank,
191                world_size: self.world_size(),
192            })?;
193        Ok(<D as InProcessDevice<T>>::copy_from_device(
194            context,
195            rank_buffer,
196        )?)
197    }
198}