use super::*;
use std::net::Ipv4Addr;
const TEST_SALT: SessionSalt = [
0x42, 0x42, 0x42, 0x42, 0x42, 0x42, 0x42, 0x42,
0x42, 0x42, 0x42, 0x42, 0x42, 0x42, 0x42, 0x42,
];
fn fake_relay(
port: u16,
ranks: Vec<u32>,
salt: SessionSalt,
per_rank_frames: Vec<Vec<RoundFrame>>,
) -> Result<Vec<Vec<RoundFrame>>> {
assert_eq!(ranks.len(), per_rank_frames.len());
let n_rounds = per_rank_frames.first().map(|v| v.len()).unwrap_or(0);
let addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), port);
let mut stream = TcpStream::connect(addr)
.map_err(|e| TensorError::new(&format!("fake_relay: connect: {e}")))?;
stream.set_nodelay(true).ok();
crate::distributed::wire::write_channel_magic(
&mut stream,
crate::distributed::wire::CHANNEL_MAGIC_DATA,
)?;
MuxRecord::control(RelayControlMsg::Hello {
host: "test-host".into(),
ranks: ranks.clone(),
})
.write_to(&mut stream, &salt)?;
match MuxRecord::read_from(&mut stream, &salt)? {
Some(MuxRecord::Control(RelayControlMsg::HelloAck)) => {}
other => {
return Err(TensorError::new(&format!(
"fake_relay: expected HelloAck, got {other:?}"
)));
}
}
let mut received: Vec<Vec<RoundFrame>> = ranks.iter().map(|_| Vec::new()).collect();
for r in 0..n_rounds {
let round_frames: Vec<&RoundFrame> =
per_rank_frames.iter().map(|frames| &frames[r]).collect();
let folded = sum_frames(&round_frames)?;
let mut buf = Vec::new();
write_round_frame(&mut buf, &folded, &salt)?;
MuxRecord::host_frame(buf).write_to(&mut stream, &salt)?;
match MuxRecord::read_from(&mut stream, &salt)? {
Some(MuxRecord::Broadcast { payload }) => {
let frame = read_round_frame(&mut payload.as_slice(), &salt)?
.ok_or_else(|| {
TensorError::new("fake_relay: truncated consensus frame")
})?;
for slot in received.iter_mut() {
slot.push(frame.clone());
}
}
other => {
return Err(TensorError::new(&format!(
"fake_relay: expected Broadcast reply, got {other:?}"
)));
}
}
}
Ok(received)
}
fn one_tensor_frame(data: &[f32]) -> RoundFrame {
RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![data.len() as u32],
bytes: f32_to_bytes(data),
}],
weight: 1.0,
..Default::default()
}
}
fn two_tensor_frame(a: &[f32], b: &[f32]) -> RoundFrame {
RoundFrame {
tensors: vec![
TensorPayload {
dtype: DTYPE_F32,
shape: vec![a.len() as u32],
bytes: f32_to_bytes(a),
},
TensorPayload {
dtype: DTYPE_F32,
shape: vec![b.len() as u32],
bytes: f32_to_bytes(b),
},
],
weight: 1.0,
..Default::default()
}
}
#[test]
fn two_rank_average_one_round() {
let avg = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
2,
TEST_SALT,
)
.unwrap();
let port = avg.port();
let recv = fake_relay(
port,
vec![0, 1],
TEST_SALT,
vec![
vec![one_tensor_frame(&[1.0, 2.0, 3.0])],
vec![one_tensor_frame(&[3.0, 4.0, 5.0])],
],
)
.unwrap();
avg.shutdown().unwrap();
let expected = bytes_as_f32(&recv[0][0].tensors[0].bytes).unwrap();
assert_eq!(expected, vec![2.0, 3.0, 4.0]);
assert_eq!(recv[0][0], recv[1][0]);
}
#[test]
fn three_rank_average_multi_round_multi_tensor() {
let avg = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
3,
TEST_SALT,
)
.unwrap();
let port = avg.port();
let r0_frames = vec![
two_tensor_frame(&[0.0, 10.0], &[1.0, 1.0]),
two_tensor_frame(&[0.0, 5.0], &[0.5, 0.5]),
];
let r1_frames = vec![
two_tensor_frame(&[10.0, 20.0], &[2.0, 2.0]),
two_tensor_frame(&[5.0, 10.0], &[1.0, 1.0]),
];
let r2_frames = vec![
two_tensor_frame(&[20.0, 30.0], &[3.0, 3.0]),
two_tensor_frame(&[10.0, 15.0], &[1.5, 1.5]),
];
let recv = fake_relay(
port,
vec![0, 1, 2],
TEST_SALT,
vec![r0_frames, r1_frames, r2_frames],
)
.unwrap();
avg.shutdown().unwrap();
assert_eq!(recv[0].len(), 2, "rank 0 should receive 2 averaged frames");
assert_eq!(recv[1].len(), 2);
assert_eq!(recv[2].len(), 2);
let r1_t0 = bytes_as_f32(&recv[0][0].tensors[0].bytes).unwrap();
let r1_t1 = bytes_as_f32(&recv[0][0].tensors[1].bytes).unwrap();
assert_eq!(r1_t0, vec![10.0, 20.0]);
assert_eq!(r1_t1, vec![2.0, 2.0]);
let r2_t0 = bytes_as_f32(&recv[0][1].tensors[0].bytes).unwrap();
let r2_t1 = bytes_as_f32(&recv[0][1].tensors[1].bytes).unwrap();
assert_eq!(r2_t0, vec![5.0, 10.0]);
assert_eq!(r2_t1, vec![1.0, 1.0]);
assert_eq!(recv[0], recv[1]);
assert_eq!(recv[1], recv[2]);
}
#[test]
fn rejects_non_hello_first_record() {
let avg = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
1,
TEST_SALT,
)
.unwrap();
let port = avg.port();
let mut s =
TcpStream::connect(SocketAddr::new(Ipv4Addr::LOCALHOST.into(), port)).unwrap();
crate::distributed::wire::write_channel_magic(
&mut s,
crate::distributed::wire::CHANNEL_MAGIC_DATA,
)
.unwrap();
MuxRecord::data(0, vec![1, 2, 3])
.write_to(&mut s, &TEST_SALT)
.unwrap();
let ack = MuxRecord::read_from(&mut s, &TEST_SALT);
assert!(
matches!(ack, Ok(None)) || ack.is_err(),
"controller should drop the connection, got {ack:?}"
);
drop(s);
let _ = avg.shutdown(); }
#[test]
fn rejects_rank_out_of_range() {
let avg = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
1,
TEST_SALT,
)
.unwrap();
let port = avg.port();
let mut s =
TcpStream::connect(SocketAddr::new(Ipv4Addr::LOCALHOST.into(), port)).unwrap();
crate::distributed::wire::write_channel_magic(
&mut s,
crate::distributed::wire::CHANNEL_MAGIC_DATA,
)
.unwrap();
MuxRecord::control(RelayControlMsg::Hello {
host: "rogue".into(),
ranks: vec![5], })
.write_to(&mut s, &TEST_SALT)
.unwrap();
let ack = MuxRecord::read_from(&mut s, &TEST_SALT);
assert!(
matches!(ack, Ok(None)) || ack.is_err(),
"controller should reject out-of-range rank, got {ack:?}"
);
drop(s);
let _ = avg.shutdown();
}
#[test]
fn rejects_non_f32_dtype_in_reduce() {
let frames = vec![
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: 7, shape: vec![2],
bytes: vec![0; 8],
}],
..Default::default()
}),
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: 7,
shape: vec![2],
bytes: vec![0; 8],
}],
..Default::default()
}),
];
let err = reduce_realized_work(&frames).unwrap_err();
assert!(
err.to_string().contains("dtype tag 7"),
"expected dtype-tag-7-not-supported, got: {err}"
);
}
#[test]
fn rejects_shape_mismatch_across_ranks() {
let frames = vec![
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![2],
bytes: f32_to_bytes(&[1.0, 2.0]),
}],
..Default::default()
}),
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![3],
bytes: f32_to_bytes(&[1.0, 2.0, 3.0]),
}],
..Default::default()
}),
];
let err = reduce_realized_work(&frames).unwrap_err();
assert!(err.to_string().contains("shape"), "got: {err}");
}
#[test]
fn reduce_realized_work_normalizes_by_accepted_mass_only() {
let frames = vec![
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![2],
bytes: f32_to_bytes(&[3.0, 6.0]), }],
weight: 3.0,
..Default::default()
}),
None, Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![2],
bytes: f32_to_bytes(&[5.0, 10.0]), }],
weight: 1.0,
..Default::default()
}),
];
let out = reduce_realized_work(&frames).unwrap();
let consensus = bytes_as_f32(&out.tensors[0].bytes).unwrap();
assert!((consensus[0] - 2.0).abs() < 1e-6, "got {consensus:?}");
assert!((consensus[1] - 4.0).abs() < 1e-6, "got {consensus:?}");
assert!(
(out.weight - 4.0).abs() < 1e-9,
"accepted mass, got {}",
out.weight
);
}
#[test]
fn reduce_realized_work_control_is_pure_sum() {
let frames = vec![
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![2],
bytes: f32_to_bytes(&[3.0, 0.0]),
}],
kind: RoundKind::Control,
..Default::default()
}),
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![2],
bytes: f32_to_bytes(&[0.0, 5.0]),
}],
kind: RoundKind::Control,
..Default::default()
}),
];
let out = reduce_realized_work(&frames).unwrap();
let sum = bytes_as_f32(&out.tensors[0].bytes).unwrap();
assert!((sum[0] - 3.0).abs() < 1e-6, "got {sum:?}");
assert!((sum[1] - 5.0).abs() < 1e-6, "got {sum:?}");
}
#[test]
fn reduce_realized_work_zero_mass_returns_untouched_sum() {
let frames = vec![
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![1],
bytes: f32_to_bytes(&[0.0]),
}],
weight: 0.0,
..Default::default()
}),
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![1],
bytes: f32_to_bytes(&[0.0]),
}],
weight: 0.0,
..Default::default()
}),
];
let out = reduce_realized_work(&frames).unwrap();
assert_eq!(out.weight, 0.0);
assert!(bytes_as_f32(&out.tensors[0].bytes).unwrap()[0].abs() < 1e-9);
}
#[test]
fn reduce_realized_work_rejects_all_dead() {
let frames: Vec<Option<RoundFrame>> = vec![None, None];
let err = reduce_realized_work(&frames).unwrap_err();
assert!(
err.to_string().contains("no accepted frames"),
"got: {err}"
);
}
#[test]
fn averager_zero_world_size_errors() {
let err = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
0,
TEST_SALT,
)
.unwrap_err();
assert!(err.to_string().contains("world_size"), "got: {err}");
}
#[test]
fn rejects_round_frame_with_wrong_inner_salt() {
use crate::distributed::wire::SESSION_SALT_BYTES;
let controller_salt = TEST_SALT;
let rogue_salt: SessionSalt = [0xAAu8; SESSION_SALT_BYTES];
assert_ne!(controller_salt, rogue_salt);
let avg = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
1,
controller_salt,
)
.unwrap();
let port = avg.port();
let mut stream =
TcpStream::connect(SocketAddr::new(Ipv4Addr::LOCALHOST.into(), port)).unwrap();
crate::distributed::wire::write_channel_magic(
&mut stream,
crate::distributed::wire::CHANNEL_MAGIC_DATA,
)
.unwrap();
MuxRecord::control(RelayControlMsg::Hello {
host: "test".into(),
ranks: vec![0],
})
.write_to(&mut stream, &controller_salt)
.unwrap();
match MuxRecord::read_from(&mut stream, &controller_salt).unwrap() {
Some(MuxRecord::Control(RelayControlMsg::HelloAck)) => {}
other => panic!("expected HelloAck, got {other:?}"),
}
let mut buf = Vec::new();
write_round_frame(&mut buf, &one_tensor_frame(&[1.0, 2.0, 3.0]), &rogue_salt).unwrap();
MuxRecord::host_frame(buf)
.write_to(&mut stream, &controller_salt)
.unwrap();
let _ = MuxRecord::read_from(&mut stream, &controller_salt);
drop(stream);
let err = avg.shutdown().expect_err(
"controller's reduce loop must propagate an inner-RoundFrame HMAC error",
);
assert!(
err.to_string().contains("HMAC verification failed"),
"expected HMAC verification failure, got: {err}"
);
}
#[test]
fn elastic_scatter_declares_broken_connection_dead_and_continues() {
use std::net::{Shutdown, TcpListener, TcpStream};
let listener =
TcpListener::bind(SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0)).unwrap();
let addr = listener.local_addr().unwrap();
let pair = || {
let client = TcpStream::connect(addr).unwrap();
let (server, _) = listener.accept().unwrap();
(server, client)
};
let (ctrl0, _peer0) = pair();
let (ctrl1, mut peer1) = pair();
ctrl0.shutdown(Shutdown::Both).unwrap();
let mut conn_writes = vec![ctrl0, ctrl1];
let conn_ranks = vec![vec![0usize], vec![1usize, 2usize]];
let dead = DeadRanks::new(3);
let frames: Vec<Option<RoundFrame>> = vec![
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![1],
bytes: f32_to_bytes(&[0.0]),
}],
weight: 1.0,
..Default::default()
}),
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![1],
bytes: f32_to_bytes(&[3.0]),
}],
weight: 2.0,
..Default::default()
}),
];
average_and_scatter(
&frames,
&mut conn_writes,
&conn_ranks,
&dead,
&TEST_SALT,
None,
None,
)
.expect("elastic scatter must not propagate a per-connection failure");
assert!(dead.is_dead(0), "broken connection's rank must be declared dead");
assert!(!dead.is_dead(1) && !dead.is_dead(2), "survivors stay alive");
match MuxRecord::read_from(&mut peer1, &TEST_SALT).unwrap() {
Some(MuxRecord::Broadcast { payload }) => {
let frame = read_round_frame(&mut payload.as_slice(), &TEST_SALT)
.unwrap()
.expect("scattered frame");
let vals = bytes_as_f32(&frame.tensors[0].bytes).unwrap();
assert!((vals[0] - 1.0).abs() < 1e-6, "consensus, got {vals:?}");
assert!((frame.weight - 3.0).abs() < 1e-9, "accepted mass rides down");
}
other => panic!("expected Broadcast record, got {other:?}"),
}
}
#[test]
fn sum_frames_is_a_pure_sum_with_summed_mass() {
let a = RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![2],
bytes: f32_to_bytes(&[3.0, 6.0]),
}],
weight: 3.0,
..Default::default()
};
let b = RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![2],
bytes: f32_to_bytes(&[5.0, 10.0]),
}],
weight: 1.0,
..Default::default()
};
let folded = sum_frames(&[&a, &b]).unwrap();
let vals = bytes_as_f32(&folded.tensors[0].bytes).unwrap();
assert_eq!(vals, vec![8.0, 16.0], "values sum, never divide");
assert!((folded.weight - 4.0).abs() < 1e-9, "masses sum");
assert_eq!(folded.kind, RoundKind::Model, "kind preserved");
}
#[test]
fn sum_frames_rejects_kind_mismatch() {
let model = one_tensor_frame(&[1.0]);
let control = RoundFrame {
kind: RoundKind::Control,
..one_tensor_frame(&[1.0])
};
let err = sum_frames(&[&model, &control]).unwrap_err();
assert!(err.to_string().contains("kind"), "got: {err}");
}
fn one_tensor_frame_bf16(data: &[f32], weight: f64) -> RoundFrame {
RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_BF16,
shape: vec![data.len() as u32],
bytes: f32_slice_to_payload_bytes(data, DTYPE_BF16).unwrap(),
}],
kind: RoundKind::Model,
weight,
}
}
#[test]
fn bf16_bits_round_to_nearest_even() {
use super::round_frame::{bf16_bits_to_f32, f32_to_bf16_bits};
for v in [0.0f32, 1.0, -2.0, 0.5, 256.0, -0.09375] {
assert_eq!(bf16_bits_to_f32(f32_to_bf16_bits(v)), v, "{v}");
}
let halfway = f32::from_bits(0x3F80_8000);
assert_eq!(bf16_bits_to_f32(f32_to_bf16_bits(halfway)), 1.0);
let above = f32::from_bits(0x3F80_8001);
assert_eq!(
bf16_bits_to_f32(f32_to_bf16_bits(above)),
f32::from_bits(0x3F81_0000)
);
assert_eq!(bf16_bits_to_f32(f32_to_bf16_bits(f32::INFINITY)), f32::INFINITY);
assert_eq!(
bf16_bits_to_f32(f32_to_bf16_bits(f32::NEG_INFINITY)),
f32::NEG_INFINITY
);
assert!(bf16_bits_to_f32(f32_to_bf16_bits(f32::NAN)).is_nan());
for b in [0u16, 1, 0x3F80, 0x7F80, 0x8000, 0x4049] {
assert_eq!(f32_to_bf16_bits(bf16_bits_to_f32(b)), b, "bits {b:#06x}");
}
}
#[test]
fn payload_codec_round_trips() {
let data = [1.5f32, -3.0, 0.0, 42.0];
for dtype in [DTYPE_F32, DTYPE_BF16] {
let p = TensorPayload {
dtype,
shape: vec![4],
bytes: f32_slice_to_payload_bytes(&data, dtype).unwrap(),
};
assert_eq!(payload_to_f32(&p).unwrap(), data.to_vec(), "dtype {dtype}");
}
assert!(f32_slice_to_payload_bytes(&data, 9).is_err());
}
#[test]
fn round_frame_wire_len_matches_serialized_frame() {
for dtype in [DTYPE_F32, DTYPE_BF16] {
let frame = RoundFrame {
tensors: vec![
TensorPayload {
dtype,
shape: vec![2, 3],
bytes: f32_slice_to_payload_bytes(&[1.5f32; 6], dtype).unwrap(),
},
TensorPayload {
dtype,
shape: vec![4],
bytes: f32_slice_to_payload_bytes(&[-2.0f32; 4], dtype).unwrap(),
},
],
kind: RoundKind::Model,
weight: 3.0,
};
let parts: Vec<PayloadPart<'_>> = frame
.tensors
.iter()
.map(|t| PayloadPart {
dtype: t.dtype,
shape: &t.shape,
nbytes: t.bytes.len() as u64,
})
.collect();
let mut buf = Vec::new();
write_round_frame(&mut buf, &frame, &TEST_SALT).unwrap();
assert_eq!(
buf.len() as u64,
round_frame_wire_len(&parts),
"dtype {dtype}"
);
let parsed = read_round_frame(&mut buf.as_slice(), &TEST_SALT)
.unwrap()
.unwrap();
assert_eq!(parsed, frame, "dtype {dtype}");
}
}
#[test]
fn zeros_emitter_streams_the_same_bytes_a_zeros_model_would() {
let frame = RoundFrame {
tensors: vec![
TensorPayload {
dtype: DTYPE_F32,
shape: vec![2, 3],
bytes: f32_slice_to_payload_bytes(&[0.0f32; 6], DTYPE_F32).unwrap(),
},
TensorPayload {
dtype: DTYPE_F32,
shape: vec![4],
bytes: f32_slice_to_payload_bytes(&[0.0f32; 4], DTYPE_F32).unwrap(),
},
],
kind: RoundKind::Control,
weight: 0.0,
};
let mut via_model = Vec::new();
write_round_frame(&mut via_model, &frame, &TEST_SALT).unwrap();
let parts: Vec<PayloadPart<'_>> = frame
.tensors
.iter()
.map(|t| PayloadPart {
dtype: t.dtype,
shape: &t.shape,
nbytes: t.bytes.len() as u64,
})
.collect();
let zeros = [0u8; 5];
let mut via_emitter = Vec::new();
write_round_frame_streamed(
&mut via_emitter,
RoundKind::Control,
0.0,
&parts,
&TEST_SALT,
&mut |ti, tee| {
let mut left = parts[ti].nbytes;
while left > 0 {
let n = left.min(zeros.len() as u64) as usize;
std::io::Write::write_all(tee, &zeros[..n])
.map_err(|e| TensorError::new(&e.to_string()))?;
left -= n as u64;
}
Ok(())
},
)
.unwrap();
assert_eq!(via_model, via_emitter, "zeros wire bytes must be identical");
}
#[test]
fn streamed_writer_rejects_emitter_byte_count_mismatch() {
let shape = [2u32];
let parts = [PayloadPart {
dtype: DTYPE_F32,
shape: &shape,
nbytes: 8,
}];
let mut buf = Vec::new();
let err = write_round_frame_streamed(
&mut buf,
RoundKind::Model,
1.0,
&parts,
&TEST_SALT,
&mut |_, tee| {
use std::io::Write;
tee.write_all(&[0u8; 4])
.map_err(|e| TensorError::new(&e.to_string()))
},
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("declared") && msg.contains("4"),
"expected loud emitter byte-count mismatch, got: {msg}"
);
}
#[test]
fn sum_frames_bf16_accumulates_in_f32() {
let frames: Vec<RoundFrame> =
(0..512).map(|_| one_tensor_frame_bf16(&[1.0], 1.0)).collect();
let refs: Vec<&RoundFrame> = frames.iter().collect();
let folded = sum_frames(&refs).unwrap();
assert_eq!(folded.tensors[0].dtype, DTYPE_BF16, "fold preserves the wire dtype");
assert_eq!(payload_to_f32(&folded.tensors[0]).unwrap(), vec![512.0]);
assert!((folded.weight - 512.0).abs() < 1e-9);
}
#[test]
fn sum_frames_rejects_dtype_mix() {
let bf16 = one_tensor_frame_bf16(&[1.0], 1.0);
let f32f = one_tensor_frame(&[1.0]);
let err = sum_frames(&[&bf16, &f32f]).unwrap_err();
assert!(err.to_string().contains("dtype"), "got: {err}");
}
#[test]
fn reduce_realized_work_bf16_normalizes() {
let frames = vec![
Some(one_tensor_frame_bf16(&[6.0, 12.0], 3.0)),
Some(one_tensor_frame_bf16(&[4.0, 8.0], 1.0)),
];
let reduced = reduce_realized_work(&frames).unwrap();
assert_eq!(reduced.tensors[0].dtype, DTYPE_BF16);
assert_eq!(
payload_to_f32(&reduced.tensors[0]).unwrap(),
vec![2.5, 5.0]
);
assert!((reduced.weight - 4.0).abs() < 1e-9);
}
#[test]
fn bf16_reduce_tracks_f32_reduce_within_wire_precision() {
let a = [0.123_456_7f32, -1.234_567_8, 3.317_742_9];
let b = [2.941_385_2f32, -0.577_215_7, 1.489_306_4];
let f32_frames = vec![
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![3],
bytes: f32_slice_to_payload_bytes(&a, DTYPE_F32).unwrap(),
}],
kind: RoundKind::Model,
weight: 2.0,
}),
Some(RoundFrame {
tensors: vec![TensorPayload {
dtype: DTYPE_F32,
shape: vec![3],
bytes: f32_slice_to_payload_bytes(&b, DTYPE_F32).unwrap(),
}],
kind: RoundKind::Model,
weight: 1.0,
}),
];
let bf16_frames = vec![
Some(one_tensor_frame_bf16(&a, 2.0)),
Some(one_tensor_frame_bf16(&b, 1.0)),
];
let exact = payload_to_f32(&reduce_realized_work(&f32_frames).unwrap().tensors[0]).unwrap();
let quant = payload_to_f32(&reduce_realized_work(&bf16_frames).unwrap().tensors[0]).unwrap();
for (e, q) in exact.iter().zip(&quant) {
assert!(
(e - q).abs() <= e.abs() * (2.0 / 256.0) + 1e-6,
"exact {e} vs bf16 {q}"
);
}
}
#[test]
fn host_fold_then_reduce_matches_flat_reduce() {
let r0 = one_tensor_frame(&[1.0, 2.0]);
let r1 = one_tensor_frame(&[3.0, 4.0]);
let r2 = one_tensor_frame(&[5.0, 6.0]);
let flat = reduce_realized_work(&[
Some(r0.clone()),
Some(r1.clone()),
Some(r2.clone()),
])
.unwrap();
let host_a = sum_frames(&[&r0, &r1]).unwrap();
let host_b = sum_frames(&[&r2]).unwrap();
let folded = reduce_realized_work(&[Some(host_a), Some(host_b)]).unwrap();
assert_eq!(
bytes_as_f32(&flat.tensors[0].bytes).unwrap(),
bytes_as_f32(&folded.tensors[0].bytes).unwrap(),
);
assert!((flat.weight - folded.weight).abs() < 1e-9);
}
#[test]
fn rejects_per_rank_data_record_on_data_channel() {
let avg = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
1,
TEST_SALT,
)
.unwrap();
let port = avg.port();
let mut stream =
TcpStream::connect(SocketAddr::new(Ipv4Addr::LOCALHOST.into(), port)).unwrap();
crate::distributed::wire::write_channel_magic(
&mut stream,
crate::distributed::wire::CHANNEL_MAGIC_DATA,
)
.unwrap();
MuxRecord::control(RelayControlMsg::Hello {
host: "stale".into(),
ranks: vec![0],
})
.write_to(&mut stream, &TEST_SALT)
.unwrap();
match MuxRecord::read_from(&mut stream, &TEST_SALT).unwrap() {
Some(MuxRecord::Control(RelayControlMsg::HelloAck)) => {}
other => panic!("expected HelloAck, got {other:?}"),
}
let mut buf = Vec::new();
write_round_frame(&mut buf, &one_tensor_frame(&[1.0]), &TEST_SALT).unwrap();
MuxRecord::data(0, buf).write_to(&mut stream, &TEST_SALT).unwrap();
let _ = MuxRecord::read_from(&mut stream, &TEST_SALT);
drop(stream);
let err = avg.shutdown().expect_err(
"a per-rank Data record on the data channel must surface as an error",
);
assert!(
err.to_string().contains("HostFrame"),
"expected the mixed-builds diagnostic, got: {err}"
);
}
#[test]
fn two_host_fold_average_one_round() {
let avg = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
3,
TEST_SALT,
)
.unwrap();
let port = avg.port();
let host_a = std::thread::spawn(move || {
fake_relay(
port,
vec![0, 1],
TEST_SALT,
vec![
vec![one_tensor_frame(&[1.0, 2.0])],
vec![one_tensor_frame(&[2.0, 3.0])],
],
)
});
let host_b = std::thread::spawn(move || {
fake_relay(
port,
vec![2],
TEST_SALT,
vec![vec![one_tensor_frame(&[3.0, 4.0])]],
)
});
let recv_a = host_a.join().unwrap().unwrap();
let recv_b = host_b.join().unwrap().unwrap();
avg.shutdown().unwrap();
let consensus = bytes_as_f32(&recv_a[0][0].tensors[0].bytes).unwrap();
assert_eq!(consensus, vec![2.0, 3.0]);
assert_eq!(recv_a[0][0], recv_a[1][0]);
assert_eq!(recv_a[0][0], recv_b[0][0]);
assert!((recv_b[0][0].weight - 3.0).abs() < 1e-9);
}