use super::{WeightCache, WeightRef, resolve_device};
use crate::{Device, Session};
use rlx_driver::ProcessGroup;
use rlx_ir::Graph;
use serde::{Deserialize, Serialize};
const TAG_TRAIN: u32 = 42;
#[derive(Serialize, Deserialize, Clone, Debug)]
pub struct DataRef {
pub input: String,
pub uri: String,
pub elem: usize,
pub shard_start: usize,
pub shard_len: usize,
}
#[derive(Serialize, Deserialize, Clone, Debug)]
pub struct TrainSpec {
pub graph: Graph,
pub params: Vec<WeightRef>,
pub grad_start: usize,
pub loss_index: usize,
pub data: Vec<DataRef>,
pub seed_input: Option<String>,
pub momentum: f32,
pub lr_per_epoch: Vec<f32>,
pub batch: usize,
pub device: String,
pub grad_group: u64,
#[serde(default)]
pub push_data: bool,
}
#[derive(Clone, Debug)]
pub struct TrainMetrics {
pub device: Device,
pub lanes: Vec<Device>,
pub inventory: Vec<Device>,
pub samples: usize,
pub wall_s: f64,
pub compute_s: f64,
pub comm_s: f64,
pub first_loss: f32,
pub last_loss: f32,
}
pub fn ship_train(group: &ProcessGroup, rank: u32, spec: &TrainSpec) -> Result<(), String> {
let bytes = serde_json::to_vec(spec).map_err(|e| e.to_string())?;
group
.transport()
.send_bytes(rank, TAG_TRAIN, &bytes)
.map_err(|e| format!("ship_train: {e}"))
}
pub fn recv_train(group: &ProcessGroup) -> Result<TrainSpec, String> {
let bytes = group
.transport()
.recv_bytes(0, TAG_TRAIN)
.map_err(|e| format!("recv_train: {e}"))?;
serde_json::from_slice(&bytes).map_err(|e| format!("TrainSpec: {e}"))
}
const TAG_DATA_BASE: u32 = 45;
const TAG_PARAM_BASE: u32 = 60;
pub fn push_shards<W: FnMut(&str) -> Vec<f32>>(
group: &ProcessGroup,
rank: u32,
spec: &TrainSpec,
mut resolve: W,
) -> Result<(), String> {
let mut cache = WeightCache::new();
for (i, w) in spec.params.iter().enumerate() {
let vals = cache.f32(&w.uri).unwrap_or_else(|_| resolve(&w.uri));
group
.send_f32(rank, TAG_PARAM_BASE + i as u32, &vals)
.map_err(|e| format!("push param {} → {rank}: {e}", w.name))?;
}
for (i, d) in spec.data.iter().enumerate() {
let shard = resolve_shard(
&d.uri,
d.elem,
d.shard_start,
d.shard_len,
&mut cache,
&mut resolve,
);
group
.send_f32(rank, TAG_DATA_BASE + i as u32, &shard)
.map_err(|e| format!("push_shards[{i}] → {rank}: {e}"))?;
}
Ok(())
}
pub fn pull_shards(
group: &ProcessGroup,
spec: &mut TrainSpec,
dir: &std::path::Path,
) -> Result<(), String> {
let write_bin = |name: &str, vals: &[f32]| -> Result<String, String> {
let path = dir.join(name);
let mut bytes = Vec::with_capacity(vals.len() * 4);
for v in vals {
bytes.extend_from_slice(&v.to_le_bytes());
}
std::fs::write(&path, &bytes).map_err(|e| format!("write {name}: {e}"))?;
Ok(format!("file://{}", path.display()))
};
for i in 0..spec.params.len() {
let vals = group
.recv_f32(0, TAG_PARAM_BASE + i as u32)
.map_err(|e| format!("pull param[{i}]: {e}"))?;
spec.params[i].uri = write_bin(&format!("rlx_pushed_param_{i}.bin"), &vals)?;
}
for i in 0..spec.data.len() {
let vals = group
.recv_f32(0, TAG_DATA_BASE + i as u32)
.map_err(|e| format!("pull_shards[{i}]: {e}"))?;
spec.data[i].uri = write_bin(&format!("rlx_pushed_shard_{i}.bin"), &vals)?;
spec.data[i].shard_start = 0; }
Ok(())
}
fn training_lanes(spec_device: &str, available: &[Device], graph: &Graph) -> Vec<Device> {
if spec_device.eq_ignore_ascii_case("all") {
let native = crate::devices_for(graph);
let mut ds: Vec<Device> = available
.iter()
.copied()
.filter(|d| !matches!(d, Device::Ane))
.collect();
ds.sort_by_key(|d| !native.contains(d));
if ds.is_empty() {
vec![crate::fastest_device()]
} else {
ds
}
} else {
vec![resolve_device(spec_device)]
}
}
pub fn run_train<W, R>(
spec: &TrainSpec,
world: u32,
mut resolve: W,
mut reduce: R,
log: bool,
) -> Result<(TrainMetrics, Vec<(String, Vec<f32>)>), String>
where
W: FnMut(&str) -> Vec<f32>,
R: FnMut(&[f32]) -> Vec<f32>,
{
use std::time::Instant;
let inventory = crate::available_devices();
let mut cache = WeightCache::new();
let mut params: Vec<Vec<f32>> = spec
.params
.iter()
.map(|w| cache.f32(&w.uri).unwrap_or_else(|_| resolve(&w.uri)))
.collect();
let names: Vec<String> = spec.params.iter().map(|w| w.name.clone()).collect();
let mut vel: Vec<Vec<f32>> = params.iter().map(|p| vec![0.0; p.len()]).collect();
let feeds: Vec<Feed> = spec
.data
.iter()
.map(|d| Feed {
input: d.input.clone(),
elem: d.elem,
data: resolve_shard(
&d.uri,
d.elem,
d.shard_start,
d.shard_len,
&mut cache,
&mut resolve,
),
})
.collect();
let shard_len = spec.data.first().map(|d| d.shard_len).unwrap_or(0);
let batches = shard_len / spec.batch.max(1);
let devices: Vec<Device> =
training_lanes(&spec.device, &crate::available_devices(), &spec.graph);
if world > 1 {
assert_uniform_shards(&mut reduce, batches)?;
}
let np = params.len();
let (mut compute_s, mut comm_s) = (0.0f64, 0.0f64);
let (mut first_loss, mut last_loss) = (f32::NAN, f32::NAN);
let mut samples = 0usize;
let wall = Instant::now();
if devices.len() > 1 {
if log {
let names: Vec<&str> = devices.iter().map(|d| crate::device_label(*d)).collect();
eprintln!(" intra-node lanes engaged: {names:?} (sample-weighted DP)");
}
let mut lanes = spawn_lanes(spec, &devices);
let ndev = lanes.len();
for (epoch, &lr) in spec.lr_per_epoch.iter().enumerate() {
let mut epoch_loss = 0.0f64;
for s in 0..batches {
let t = Instant::now();
for (i, lane) in lanes.iter().enumerate() {
let b = (s * ndev + i) % batches;
let fdata: Vec<Vec<f32>> = feeds
.iter()
.map(|f| {
let off = b * spec.batch * f.elem;
f.data[off..off + spec.batch * f.elem].to_vec()
})
.collect();
lane.job.send(LaneJob::Step(params.clone(), fdata)).ok();
}
let mut sum_grads: Vec<Vec<f32>> =
params.iter().map(|p| vec![0.0f32; p.len()]).collect();
let mut loss_sum = 0.0f64;
for lane in &lanes {
let (loss, grads) = lane
.res
.recv()
.map_err(|e| format!("lane {}: {e}", crate::device_label(lane.dev)))?;
loss_sum += loss as f64;
for i in 0..np {
for (a, g) in sum_grads[i].iter_mut().zip(&grads[i]) {
*a += g;
}
}
}
compute_s += t.elapsed().as_secs_f64();
epoch_loss += loss_sum / ndev as f64;
samples += spec.batch * ndev;
weighted_sync(
&mut reduce,
&mut params,
&mut vel,
&sum_grads,
ndev,
spec.momentum,
lr,
&mut comm_s,
);
}
record_epoch(
epoch,
(epoch_loss / batches.max(1) as f64) as f32,
&mut first_loss,
&mut last_loss,
log,
spec.lr_per_epoch.len(),
);
}
for lane in &mut lanes {
lane.job.send(LaneJob::Stop).ok();
}
for lane in lanes {
lane.handle.join().ok();
}
} else {
let mut compiled = Session::new(devices[0]).compile(spec.graph.clone());
for (n, p) in names.iter().zip(¶ms) {
compiled.set_param(n, p);
}
let one = [1.0f32];
let mut sum_grads: Vec<Vec<f32>> = params.iter().map(|p| vec![0.0f32; p.len()]).collect();
for (epoch, &lr) in spec.lr_per_epoch.iter().enumerate() {
let mut epoch_loss = 0.0f64;
for b in 0..batches {
let mut feed: Vec<(&str, &[f32])> = feeds
.iter()
.map(|f| {
let off = b * spec.batch * f.elem;
(f.input.as_str(), &f.data[off..off + spec.batch * f.elem])
})
.collect();
if let Some(seed) = &spec.seed_input {
feed.push((seed.as_str(), &one));
}
let t = Instant::now();
let outs = compiled.run(&feed);
compute_s += t.elapsed().as_secs_f64();
epoch_loss += outs[spec.loss_index][0] as f64;
samples += spec.batch;
for i in 0..np {
sum_grads[i].copy_from_slice(&outs[spec.grad_start + i]);
}
weighted_sync(
&mut reduce,
&mut params,
&mut vel,
&sum_grads,
1,
spec.momentum,
lr,
&mut comm_s,
);
for (n, p) in names.iter().zip(¶ms) {
compiled.set_param(n, p);
}
}
record_epoch(
epoch,
(epoch_loss / batches.max(1) as f64) as f32,
&mut first_loss,
&mut last_loss,
log,
spec.lr_per_epoch.len(),
);
}
}
let metrics = TrainMetrics {
device: crate::device_ext::fastest_among(&devices),
lanes: devices,
inventory,
samples,
wall_s: wall.elapsed().as_secs_f64(),
compute_s,
comm_s,
first_loss,
last_loss,
};
let final_params = names.into_iter().zip(params).collect();
Ok((metrics, final_params))
}
struct Feed {
input: String,
elem: usize,
data: Vec<f32>,
}
fn resolve_shard<W: FnMut(&str) -> Vec<f32>>(
uri: &str,
elem: usize,
start: usize,
len: usize,
cache: &mut WeightCache,
resolve: &mut W,
) -> Vec<f32> {
if let Some(path) = uri.strip_prefix("file://") {
use std::io::{Read, Seek, SeekFrom};
if let Ok(mut f) = std::fs::File::open(path) {
let mut buf = vec![0u8; len * elem * 4];
if f.seek(SeekFrom::Start((start * elem * 4) as u64)).is_ok()
&& f.read_exact(&mut buf).is_ok()
{
return buf
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
}
}
}
let full = cache.f32(uri).unwrap_or_else(|_| resolve(uri));
full.get(start * elem..(start + len) * elem)
.map(<[f32]>::to_vec)
.unwrap_or(full)
}
fn assert_uniform_shards<R: FnMut(&[f32]) -> Vec<f32>>(
reduce: &mut R,
batches: usize,
) -> Result<(), String> {
let b = batches as f32;
let stats = reduce(&[b, b * b]);
let mean_b = stats.first().copied().unwrap_or(b);
let mean_b2 = stats.get(1).copied().unwrap_or(b * b);
let var = (mean_b2 - mean_b * mean_b).max(0.0);
if var > 0.25 {
return Err(format!(
"run_train: uneven shards across the cluster — this rank has {batches} batches, \
cluster mean {mean_b:.2} (var {var:.2}). The master must ship an equal shard_len \
(÷ batch) to every rank, else the gradient all-reduce deadlocks."
));
}
Ok(())
}
fn weighted_sync<R: FnMut(&[f32]) -> Vec<f32>>(
reduce: &mut R,
params: &mut [Vec<f32>],
vel: &mut [Vec<f32>],
sum_grads: &[Vec<f32>],
lane_count: usize,
momentum: f32,
lr: f32,
comm_s: &mut f64,
) {
use std::time::Instant;
let total: usize = sum_grads.iter().map(|g| g.len()).sum();
let mut flat = Vec::with_capacity(total + 1);
for g in sum_grads {
flat.extend_from_slice(g);
}
flat.push(lane_count as f32);
let t = Instant::now();
let reduced = reduce(&flat);
*comm_s += t.elapsed().as_secs_f64();
let denom = reduced.get(total).copied().unwrap_or(lane_count as f32);
let inv = if denom != 0.0 { 1.0 / denom } else { 0.0 };
apply_sgd(params, vel, &reduced, inv, momentum, lr);
}
fn record_epoch(
epoch: usize,
mean: f32,
first_loss: &mut f32,
last_loss: &mut f32,
log: bool,
total_epochs: usize,
) {
if epoch == 0 {
*first_loss = mean;
}
*last_loss = mean;
if log {
eprintln!(" epoch {}/{total_epochs}: mean loss {mean:.4}", epoch + 1);
}
}
fn apply_sgd(
params: &mut [Vec<f32>],
vel: &mut [Vec<f32>],
reduced: &[f32],
scale: f32,
momentum: f32,
lr: f32,
) {
let mut off = 0usize;
for (p, v) in params.iter_mut().zip(vel.iter_mut()) {
for j in 0..p.len() {
let g = reduced[off + j] * scale;
v[j] = momentum * v[j] + g;
p[j] -= lr * v[j];
}
off += p.len();
}
}
struct DevLane {
dev: Device,
job: std::sync::mpsc::Sender<LaneJob>,
res: std::sync::mpsc::Receiver<(f32, Vec<Vec<f32>>)>,
handle: std::thread::JoinHandle<()>,
}
enum LaneJob {
Step(Vec<Vec<f32>>, Vec<Vec<f32>>),
Stop,
}
fn spawn_lanes(spec: &TrainSpec, devices: &[Device]) -> Vec<DevLane> {
devices
.iter()
.map(|&dev| {
let (jtx, jrx) = std::sync::mpsc::channel::<LaneJob>();
let (rtx, rrx) = std::sync::mpsc::channel();
let graph = spec.graph.clone();
let names: Vec<String> = spec.params.iter().map(|w| w.name.clone()).collect();
let inputs: Vec<String> = spec.data.iter().map(|d| d.input.clone()).collect();
let seed = spec.seed_input.clone();
let (grad_start, loss_index, np) =
(spec.grad_start, spec.loss_index, spec.params.len());
let handle = std::thread::spawn(move || {
let mut compiled = Session::new(dev).compile(graph);
let one = [1.0f32];
while let Ok(job) = jrx.recv() {
match job {
LaneJob::Step(params, fdata) => {
for (n, p) in names.iter().zip(¶ms) {
compiled.set_param(n, p);
}
let mut feed: Vec<(&str, &[f32])> = inputs
.iter()
.zip(&fdata)
.map(|(n, d)| (n.as_str(), d.as_slice()))
.collect();
if let Some(s) = &seed {
feed.push((s.as_str(), &one));
}
let outs = compiled.run(&feed);
let loss = outs[loss_index][0];
let grads: Vec<Vec<f32>> =
(0..np).map(|i| outs[grad_start + i].clone()).collect();
let _ = rtx.send((loss, grads));
}
LaneJob::Stop => break,
}
}
});
DevLane {
dev,
job: jtx,
res: rrx,
handle,
}
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn training_lanes_all_fans_out_cpu_and_gpu() {
use rlx_ir::{DType, Graph, Shape};
let mut g = Graph::new("t");
let x = g.input("x", Shape::new(&[2, 2], DType::F32));
let w = g.param("w", Shape::new(&[2, 2], DType::F32));
let mm = g.matmul(x, w, Shape::new(&[2, 2], DType::F32));
g.set_outputs(vec![mm]);
let avail = [Device::Cpu, Device::Metal];
assert_eq!(training_lanes("all", &avail, &g).len(), 2);
assert_eq!(training_lanes("all", &[Device::Cpu], &g).len(), 1);
assert_eq!(training_lanes("cpu", &avail, &g).len(), 1);
}
#[test]
fn weighted_sync_unbiased_across_uneven_lane_counts() {
let p = 2usize; let ga_sum = vec![6.0f32, 12.0]; let gb = [10.0f32, 20.0]; let bucket_a = [ga_sum[0], ga_sum[1], 3.0];
let bucket_b = [gb[0], gb[1], 1.0];
let mean: Vec<f32> = (0..p + 1)
.map(|i| (bucket_a[i] + bucket_b[i]) / 2.0)
.collect();
let mut mean_iter = mean.clone();
let mut reduce = move |buf: &[f32]| -> Vec<f32> {
assert_eq!(buf.len(), p + 1);
std::mem::take(&mut mean_iter)
};
let mut params = vec![vec![0.0f32; p]];
let mut vel = vec![vec![0.0f32; p]];
let sum_grads = vec![ga_sum.clone()]; let mut comm = 0.0;
weighted_sync(
&mut reduce,
&mut params,
&mut vel,
&sum_grads,
3,
0.0,
1.0,
&mut comm,
);
let expect = [(6.0 + 10.0) / 4.0, (12.0 + 20.0) / 4.0];
assert!(
(params[0][0] + expect[0]).abs() < 1e-6,
"got {:?}",
params[0]
);
assert!(
(params[0][1] + expect[1]).abs() < 1e-6,
"got {:?}",
params[0]
);
}
}