ruccl/in_process/collective/
memory.rs1use 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 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}