1use anyhow::Result;
28use std::path::Path;
29
30use spark_runtime::gpu::GpuBackend;
31use spark_runtime::weights::{WeightLoader, WeightStore, parse_expert_index};
32
33use crate::weight_peer::WeightTensorRecord;
34
35pub struct RdmaWeightLoader {
37 pub peer_addr: String,
39 pub model_id: Option<String>,
42 pub ep_rank: usize,
43 pub ep_world_size: usize,
44 pub num_experts: usize,
45 pub stream_all_experts: bool,
48 pub peak_memory_multiplier: Option<f64>,
50}
51
52impl RdmaWeightLoader {
53 pub fn new(peer_addr: String) -> Self {
54 Self {
55 peer_addr,
56 model_id: None,
57 ep_rank: 0,
58 ep_world_size: 1,
59 num_experts: 0,
60 stream_all_experts: false,
61 peak_memory_multiplier: None,
62 }
63 }
64
65 pub fn with_ep(
66 peer_addr: String,
67 ep_rank: usize,
68 ep_world_size: usize,
69 num_experts: usize,
70 ) -> Self {
71 Self {
72 peer_addr,
73 model_id: None,
74 ep_rank,
75 ep_world_size,
76 num_experts,
77 stream_all_experts: false,
78 peak_memory_multiplier: None,
79 }
80 }
81
82 #[cfg_attr(not(atlas_rdma_verbs), allow(dead_code))]
89 fn should_skip_tensor(&self, rec: &WeightTensorRecord) -> bool {
90 if rec.extra {
91 return false;
92 }
93 if rec.name.starts_with("mtp.") {
95 return false;
96 }
97 if self.stream_all_experts && parse_expert_index(&rec.name).is_some() {
99 return true;
100 }
101 if self.ep_world_size <= 1 {
102 return false;
103 }
104 if let Some(idx) = parse_expert_index(&rec.name) {
105 let per_rank = self.num_experts / self.ep_world_size;
106 let local_start = self.ep_rank * per_rank;
107 let local_end = if self.ep_rank == self.ep_world_size - 1 {
108 self.num_experts
109 } else {
110 local_start + per_rank
111 };
112 idx < local_start || idx >= local_end
113 } else {
114 false
115 }
116 }
117}
118
119impl WeightLoader for RdmaWeightLoader {
120 fn load(
121 &self,
122 model_dir: &Path,
123 gpu: &dyn GpuBackend,
124 oom_reserve_bytes: usize,
125 ) -> Result<WeightStore> {
126 self.load_impl(model_dir, gpu, oom_reserve_bytes)
127 }
128}
129
130#[cfg(not(atlas_rdma_verbs))]
131impl RdmaWeightLoader {
132 fn load_impl(
133 &self,
134 _model_dir: &Path,
135 _gpu: &dyn GpuBackend,
136 _oom_reserve_bytes: usize,
137 ) -> Result<WeightStore> {
138 anyhow::bail!(
139 "$ATLAS_WEIGHT_PEER is set but this build has no rdma-core (atlas_rdma_verbs \
140 cfg); rebuild with rdma-core, or unset ATLAS_WEIGHT_PEER to load from disk"
141 )
142 }
143}
144
145#[cfg(atlas_rdma_verbs)]
146impl RdmaWeightLoader {
147 fn load_impl(
148 &self,
149 model_dir: &Path,
150 gpu: &dyn GpuBackend,
151 oom_reserve_bytes: usize,
152 ) -> Result<WeightStore> {
153 use std::collections::HashMap;
154 use std::ffi::c_void;
155 use std::io::Write;
156 use std::net::TcpStream;
157
158 use anyhow::{Context, bail};
159
160 use crate::expert_peer::MODE_VERBS;
161 use crate::weight_peer::{
162 rail_for_tensor, read_weight_manifest, tensor_remote_addr, write_model_request,
163 };
164 use atlas_rdma::env::{first_nonempty, first_set_u32};
165 use atlas_rdma::railset::{RailSet, RailSpec};
166 use atlas_rdma::verbs::Verbs;
167 use spark_runtime::weights::{WeightDtype, WeightTensor};
168
169 let mut stream = TcpStream::connect(&self.peer_addr)
171 .with_context(|| format!("connect weight peer {}", self.peer_addr))?;
172 stream.set_nodelay(true).ok();
173 let model_id = self
174 .model_id
175 .clone()
176 .unwrap_or_else(|| model_dir.to_string_lossy().into_owned());
177 write_model_request(&mut stream, &model_id).context("send model request")?;
178 let manifest = read_weight_manifest(&mut stream).context("read weight manifest")?;
179 let num_shards = manifest.num_shards();
180
181 let retained: Vec<&WeightTensorRecord> = manifest
184 .tensors
185 .iter()
186 .filter(|t| !self.should_skip_tensor(t))
187 .collect();
188
189 {
192 let est: u64 = retained.iter().map(|t| t.len).sum();
193 let fp8: u64 = retained
194 .iter()
195 .filter(|t| t.dtype == "F8_E4M3")
196 .map(|t| t.len)
197 .sum();
198 let fp8_frac = if est > 0 {
199 fp8 as f64 / est as f64
200 } else {
201 0.0
202 };
203 let mult =
204 self.peak_memory_multiplier
205 .unwrap_or(if fp8_frac > 0.5 { 1.5 } else { 1.3 });
206 let peak = (est as f64 * mult) as usize;
207 let free = gpu.free_memory()?;
208 let gib = |b: usize| b as f64 / (1024.0 * 1024.0 * 1024.0);
209 tracing::info!(
210 "RDMA weight load pre-flight: {:.2} GB manifest, {:.1}x = {:.2} GB peak, \
211 {:.2} GB free, {:.1} GB reserve (FP8 {:.0}%)",
212 gib(est as usize),
213 mult,
214 gib(peak),
215 gib(free),
216 gib(oom_reserve_bytes),
217 fp8_frac * 100.0,
218 );
219 if peak + oom_reserve_bytes > free {
220 bail!(
221 "OOM pre-flight (RDMA weight peer): peak {:.2} GB + {:.2} GB reserve > {:.2} GB free",
222 gib(peak),
223 gib(oom_reserve_bytes),
224 gib(free),
225 );
226 }
227 }
228
229 let spec =
235 |dev: String, gid: u32| RailSpec::new(dev, gid, rand::random::<u32>() & 0xff_ffff);
236 let rail0 = spec(
237 first_nonempty(
238 &["ATLAS_WEIGHT_RDMA_DEV", "ATLAS_EXPERT_RDMA_DEV"],
239 "roceP2p1s0f1",
240 ),
241 first_set_u32(&["ATLAS_WEIGHT_RDMA_GID", "ATLAS_EXPERT_RDMA_GID"], 3),
242 );
243 let dual = std::env::var("ATLAS_WEIGHT_DUAL_RAIL").ok().as_deref() == Some("1");
244 let specs: Vec<RailSpec> = if dual {
245 let rail1 = spec(
246 first_nonempty(
247 &["ATLAS_WEIGHT_RAIL2_DEV", "ATLAS_EXPERT_RAIL2_DEV"],
248 "rocep1s0f1",
249 ),
250 first_set_u32(&["ATLAS_WEIGHT_RAIL2_GID", "ATLAS_EXPERT_RAIL2_GID"], 3),
251 );
252 vec![rail0, rail1]
253 } else {
254 vec![rail0]
255 };
256 let n_rails = specs.len();
257
258 stream.write_all(&[MODE_VERBS]).context("send verbs mode")?;
259 let mut rs = RailSet::begin(&mut stream, &specs)?;
261
262 let max_len = retained.iter().map(|t| t.len).max().unwrap_or(0);
266 if max_len > u32::MAX as u64 {
267 bail!(
268 "tensor of {} bytes exceeds the 4 GiB single-WR RDMA READ limit \
269 (per-tensor chunking not implemented)",
270 max_len
271 );
272 }
273 let bounce_len = (max_len as usize).max(1);
274
275 let mut pinned: Vec<*mut u8> = Vec::with_capacity(n_rails);
278 let mut bounce_lkeys: Vec<u32> = Vec::with_capacity(n_rails);
279 for rail in &mut rs.rails {
280 let ptr = gpu
281 .alloc_host_pinned(bounce_len)
282 .context("alloc pinned RDMA landing bounce")?;
283 let keys = unsafe { rail.verbs.reg_mr(ptr as *mut c_void, bounce_len, false) }
286 .context("register RDMA landing bounce")?;
287 pinned.push(ptr);
288 bounce_lkeys.push(keys.lkey);
289 }
290
291 let server = rs
294 .read_server_ro(&mut stream)
295 .context("read verbs server params")?;
296 for sp in &server {
297 if sp.layers.len() != num_shards {
298 bail!(
299 "peer published {} shard MRs but manifest has {num_shards} shards",
300 sp.layers.len()
301 );
302 }
303 }
304
305 rs.complete(&mut stream, &server, "weight peer")?;
307 struct Rail {
308 verbs: Verbs,
309 bounce_ptr: *mut u8,
310 bounce_lkey: u32,
311 }
312 let mut rails: Vec<Rail> = rs
313 .into_verbs()
314 .into_iter()
315 .zip(&pinned)
316 .zip(&bounce_lkeys)
317 .map(|((verbs, &bounce_ptr), &bounce_lkey)| Rail {
318 verbs,
319 bounce_ptr,
320 bounce_lkey,
321 })
322 .collect();
323 tracing::info!(
324 "RDMA weight loader connected to {} ({} shards, {} resident tensors, {n_rails} rail(s))",
325 manifest.model_id,
326 num_shards,
327 retained.len(),
328 );
329
330 let mut weights: HashMap<String, WeightTensor> = HashMap::new();
335 let mut offload_logged = false;
336 for (idx, rec) in retained.iter().enumerate() {
337 let ri = rail_for_tensor(idx, n_rails);
338 let sp = &server[ri];
339 let (shard_base, rkey) = *sp
340 .layers
341 .get(rec.shard_index as usize)
342 .with_context(|| format!("no shard MR {} for {}", rec.shard_index, rec.name))?;
343 let remote_addr = tensor_remote_addr(shard_base, rec.offset_in_shard);
344 let len = rec.len as usize;
345 let wr_id = idx as u64;
346
347 let rail = &mut rails[ri];
348 unsafe {
352 rail.verbs
353 .post_read(
354 rail.bounce_ptr as *mut c_void,
355 rail.bounce_lkey,
356 remote_addr,
357 rkey,
358 len as u32,
359 wr_id,
360 )
361 .with_context(|| format!("post_read {}", rec.name))?;
362 }
363 match rail.verbs.poll() {
364 Ok(got) if got == wr_id => {}
365 Ok(got) => bail!(
366 "completion wr_id {got:#x} != expected {wr_id:#x} ({})",
367 rec.name
368 ),
369 Err(e) => return Err(e).with_context(|| format!("poll {}", rec.name)),
370 }
371
372 let src = unsafe { std::slice::from_raw_parts(rail.bounce_ptr, len) };
374 let dtype = WeightDtype::from_safetensors_str(&rec.dtype)
375 .with_context(|| format!("tensor {}", rec.name))?;
376 let shape: Vec<usize> = rec.shape.iter().map(|&d| d as usize).collect();
377
378 let ptr = match gpu.alloc(len) {
379 Ok(p) => {
380 gpu.copy_h2d(src, p)?;
381 p
382 }
383 Err(_) => {
384 if !offload_logged {
385 tracing::warn!(
386 "GPU alloc failed for {} ({len} bytes) — switching to managed (UVM) memory",
387 rec.name
388 );
389 offload_logged = true;
390 }
391 let p = gpu.alloc_managed(len)?;
392 unsafe {
396 std::ptr::copy_nonoverlapping(src.as_ptr(), p.0 as *mut u8, len);
397 }
398 p
399 }
400 };
401 weights.insert(rec.name.clone(), WeightTensor { ptr, shape, dtype });
402 }
403
404 drop(rails);
407 for ptr in pinned {
408 let _ = gpu.free_host_pinned(ptr, bounce_len);
409 }
410
411 tracing::info!("RDMA-loaded {} weight tensors", weights.len());
412 Ok(WeightStore::from_map(weights))
413 }
414}