use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::path::Path;
use anyhow::{Context, Result, anyhow};
use crate::forward::{BatchCol, StageBatchOut, StageOut};
use crate::{GpuCtx, Lfm2Gpu, Weights};
const OP_STEP: u32 = 0;
const OP_RESET: u32 = 1;
const OP_SHUTDOWN: u32 = 2;
const OP_BINIT: u32 = 3;
const OP_BTAB: u32 = 4;
const OP_BSTEP: u32 = 5;
const OP_ZSLOT: u32 = 6;
const OP_BSTEPW: u32 = 7;
const OP_MTP: u32 = 8;
const OP_DNRESTORE: u32 = 9;
const OP_MTP_BATCH: u32 = 10;
const OP_MTP_SEED: u32 = 11;
const OP_CHAIN: u32 = 12;
const OP_INFO: u32 = 13;
const OP_LOAD: u32 = 14;
const OP_HELLO: u32 = 15;
const OP_SHIP_BEGIN: u32 = 16;
const OP_SHIP_CHUNK: u32 = 17;
const OP_SHIP_END: u32 = 18;
const OP_JOIN: u32 = 19;
const OP_VOCAB: u32 = 20;
const OP_SIGNAL: u32 = 21;
const SIG_OFFER: u32 = 0;
const SIG_ANSWER: u32 = 1;
const SIG_FINISH: u32 = 2;
const OP_RTC_SERVE: u32 = 22;
const OP_RTC_CHAIN: u32 = 23;
const TAG_HIDDEN: u32 = 0;
const TAG_TOKEN: u32 = 1;
const TAG_INFO: u32 = 2;
const TAG_DONE: u32 = 3;
fn write_msg_bytes(s: &mut Conn, header: &[u32], payload: &[u8]) -> Result<()> {
let mut buf = Vec::with_capacity(header.len() * 4 + payload.len());
for h in header {
buf.extend_from_slice(&h.to_le_bytes());
}
buf.extend_from_slice(payload);
s.write_all(&buf)?;
Ok(())
}
fn write_msg(s: &mut Conn, header: &[u32], payload: &[f32]) -> Result<()> {
write_msg_bytes(s, header, bytemuck::cast_slice(payload))
}
fn read_bytes(s: &mut Conn, n: usize) -> Result<Vec<u8>> {
let mut b = vec![0u8; n];
s.read_exact(&mut b)?;
Ok(b)
}
fn read_u32s(s: &mut Conn, n: usize) -> Result<Vec<u32>> {
let mut b = vec![0u8; n * 4];
s.read_exact(&mut b)?;
Ok(b.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect())
}
fn read_f32s(s: &mut Conn, n: usize) -> Result<Vec<f32>> {
if n == 0 {
return Ok(Vec::new());
}
let mut b = vec![0u8; n * 4];
s.read_exact(&mut b)?;
Ok(bytemuck::cast_slice(&b).to_vec())
}
pub(crate) enum Conn {
Tcp(TcpStream),
#[cfg(feature = "ws")]
Ws(Box<WsPipe>),
#[cfg(feature = "webrtc")]
Webrtc(Box<crate::webrtc_native::WebrtcPipe>),
}
impl Conn {
fn peer_addr(&self) -> Result<std::net::SocketAddr> {
match self {
Conn::Tcp(s) => Ok(s.peer_addr()?),
#[cfg(feature = "ws")]
Conn::Ws(w) => w.peer_addr(),
#[cfg(feature = "webrtc")]
Conn::Webrtc(_) => Err(anyhow!("webrtc data channel has no socket peer address")),
}
}
}
impl std::io::Read for Conn {
fn read(&mut self, out: &mut [u8]) -> std::io::Result<usize> {
match self {
Conn::Tcp(s) => s.read(out),
#[cfg(feature = "ws")]
Conn::Ws(w) => w.read(out),
#[cfg(feature = "webrtc")]
Conn::Webrtc(w) => w.read(out),
}
}
}
impl std::io::Write for Conn {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
match self {
Conn::Tcp(s) => s.write(buf),
#[cfg(feature = "ws")]
Conn::Ws(w) => w.write(buf),
#[cfg(feature = "webrtc")]
Conn::Webrtc(w) => w.write(buf),
}
}
fn flush(&mut self) -> std::io::Result<()> {
match self {
Conn::Tcp(s) => s.flush(),
#[cfg(feature = "ws")]
Conn::Ws(w) => w.flush(),
#[cfg(feature = "webrtc")]
Conn::Webrtc(w) => w.flush(),
}
}
}
#[cfg(feature = "ws")]
fn ws_config() -> tungstenite::protocol::WebSocketConfig {
tungstenite::protocol::WebSocketConfig::default()
.max_message_size(Some(256 * 1024 * 1024))
.max_frame_size(Some(64 * 1024 * 1024))
}
#[cfg(feature = "ws")]
pub(crate) struct WsPipe {
ws: tungstenite::WebSocket<tungstenite::stream::MaybeTlsStream<TcpStream>>,
buf: Vec<u8>,
pos: usize,
}
#[cfg(feature = "ws")]
impl WsPipe {
fn new(ws: tungstenite::WebSocket<tungstenite::stream::MaybeTlsStream<TcpStream>>) -> Self {
Self {
ws,
buf: Vec::new(),
pos: 0,
}
}
fn peer_addr(&self) -> Result<std::net::SocketAddr> {
match self.ws.get_ref() {
tungstenite::stream::MaybeTlsStream::Plain(s) => Ok(s.peer_addr()?),
_ => Err(anyhow!("unexpected TLS stream")),
}
}
}
#[cfg(feature = "ws")]
impl std::io::Read for WsPipe {
fn read(&mut self, out: &mut [u8]) -> std::io::Result<usize> {
while self.pos >= self.buf.len() {
match self.ws.read() {
Ok(tungstenite::Message::Binary(b)) => {
self.buf = b.as_ref().to_vec();
self.pos = 0;
}
Ok(tungstenite::Message::Close(_)) => return Ok(0),
Ok(_) => continue, Err(tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed) => {
return Ok(0);
}
Err(tungstenite::Error::Io(e)) => return Err(e),
Err(e) => return Err(std::io::Error::other(e)),
}
}
let n = out.len().min(self.buf.len() - self.pos);
out[..n].copy_from_slice(&self.buf[self.pos..self.pos + n]);
self.pos += n;
Ok(n)
}
}
#[cfg(feature = "ws")]
impl std::io::Write for WsPipe {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.ws
.send(tungstenite::Message::binary(buf.to_vec()))
.map_err(std::io::Error::other)?;
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
self.ws.flush().map_err(std::io::Error::other)
}
}
pub(crate) fn dial(addr: &str, timeout: std::time::Duration) -> Result<Conn> {
let tcp_target = |hostport: &str| -> Result<std::net::SocketAddr> {
std::net::ToSocketAddrs::to_socket_addrs(hostport)
.with_context(|| format!("resolve {hostport}"))?
.next()
.ok_or_else(|| anyhow!("{hostport} resolved to nothing"))
};
if let Some(rest) = addr.strip_prefix("ws://") {
#[cfg(feature = "ws")]
{
let hostport = rest.split('/').next().unwrap_or(rest);
let tcp = TcpStream::connect_timeout(&tcp_target(hostport)?, timeout)?;
tcp.set_nodelay(true).ok();
let (ws, _) = tungstenite::client::client_with_config(
addr,
tungstenite::stream::MaybeTlsStream::Plain(tcp),
Some(ws_config()),
)
.map_err(|e| anyhow!("ws handshake with {addr}: {e}"))?;
return Ok(Conn::Ws(Box::new(WsPipe::new(ws))));
}
#[cfg(not(feature = "ws"))]
{
let _ = rest;
anyhow::bail!("{addr}: this build lacks the `ws` feature");
}
}
let s = TcpStream::connect_timeout(&tcp_target(addr)?, timeout)?;
s.set_nodelay(true).ok();
Ok(Conn::Tcp(s))
}
pub(crate) fn accept_conn(s: TcpStream) -> Result<Conn> {
s.set_nodelay(true).ok();
let mut first = [0u8; 1];
let n = s.peek(&mut first)?;
if n == 1 && first[0] == b'G' {
#[cfg(feature = "ws")]
{
let ws = tungstenite::accept_with_config(
tungstenite::stream::MaybeTlsStream::Plain(s),
Some(ws_config()),
)
.map_err(|e| anyhow!("ws accept: {e}"))?;
return Ok(Conn::Ws(Box::new(WsPipe::new(ws))));
}
#[cfg(not(feature = "ws"))]
anyhow::bail!("WS peer, but this build lacks the `ws` feature");
}
Ok(Conn::Tcp(s))
}
struct ChainOut {
s: Conn,
sink: bool,
}
struct LoadedStage {
gpu: Lfm2Gpu,
mtp: Option<crate::forward::MtpEngine>,
plans: Vec<crate::forward::BatchPlan>,
bpw: Option<crate::forward::BatchPlan>,
start: usize,
end: usize,
stage_first: bool,
stage_last: bool,
total: usize,
src_fp: u64,
}
struct StageEngine {
loaded: Option<LoadedStage>,
chain: Option<ChainOut>,
}
struct WorkerShared {
ctx: GpuCtx,
model_dir: Option<std::path::PathBuf>,
cache_root: std::path::PathBuf,
total_layers: usize,
vram_mb: u64,
pinned: bool,
ckpt_fp: u64,
token: Option<String>,
shipping: std::sync::Mutex<std::collections::HashSet<std::path::PathBuf>>,
eng: std::sync::Mutex<StageEngine>,
}
struct ShipClaim {
sh: std::sync::Arc<WorkerShared>,
key: std::path::PathBuf,
}
impl Drop for ShipClaim {
fn drop(&mut self) {
self.sh
.shipping
.lock()
.expect("shipping set poisoned")
.remove(&self.key);
}
}
#[derive(Default)]
pub struct WorkerOptions {
pub layers: Option<(usize, usize)>,
pub vram_gb: Option<f64>,
pub token: Option<String>,
pub cache_dir: Option<std::path::PathBuf>,
pub join: Option<String>,
pub advertise: Option<String>,
}
fn fnv64_update(mut h: u64, bytes: &[u8]) -> u64 {
for b in bytes {
h ^= u64::from(*b);
h = h.wrapping_mul(0x100_0000_01b3);
}
h
}
fn fnv64(bytes: &[u8]) -> u64 {
fnv64_update(0xcbf2_9ce4_8422_2325, bytes)
}
fn checkpoint_fingerprint(dir: &Path) -> Result<u64> {
let mut h = fnv64(&std::fs::read(dir.join("config.json")).context("read config.json")?);
let single = dir.join("model.safetensors");
let header = if single.exists() {
use std::io::Read as _;
let mut f = std::fs::File::open(&single)?;
let mut lenb = [0u8; 8];
f.read_exact(&mut lenb)?;
let hn = u64::from_le_bytes(lenb).min(16 * 1024 * 1024) as usize;
let mut hb = vec![0u8; hn];
f.read_exact(&mut hb)?;
hb
} else {
std::fs::read(dir.join("model.safetensors.index.json")).context("read safetensors index")?
};
h ^= fnv64(&header).rotate_left(1);
Ok(h)
}
pub(crate) struct MiniCkpt {
pub header: Vec<u8>,
pub slices: Vec<(u64, u64)>,
pub total_bytes: u64,
}
fn tensor_layer(name: &str) -> Option<usize> {
let idx = name.find(".layers.")?;
if name[..idx].contains("mtp") {
return None;
}
let rest = &name[idx + ".layers.".len()..];
rest[..rest.find('.')?].parse().ok()
}
pub(crate) fn plan_mini_ckpt(
src: &Path,
start: usize,
end: usize,
total: usize,
) -> Result<MiniCkpt> {
use std::io::Read as _;
let mut f = std::fs::File::open(src).with_context(|| format!("open {}", src.display()))?;
let mut lenb = [0u8; 8];
f.read_exact(&mut lenb)?;
let hlen = u64::from_le_bytes(lenb);
anyhow::ensure!(hlen < 256 * 1024 * 1024, "implausible safetensors header");
let mut hb = vec![0u8; hlen as usize];
f.read_exact(&mut hb)?;
let table: serde_json::Value = serde_json::from_slice(&hb).context("safetensors header")?;
let table = table.as_object().context("header is not an object")?;
let data_base = 8 + hlen;
let touches_ends = start == 0 || end == total;
let mut picked: Vec<(&String, u64, u64, &serde_json::Value)> = Vec::new();
for (name, meta) in table {
if name == "__metadata__" {
continue;
}
let keep = match tensor_layer(name) {
Some(i) => start <= i && i < end,
None => touches_ends,
};
if !keep {
continue;
}
let offs = meta["data_offsets"]
.as_array()
.context("tensor data_offsets")?;
let (a, b) = (
offs[0].as_u64().context("offset")?,
offs[1].as_u64().context("offset")?,
);
picked.push((name, a, b, meta));
}
anyhow::ensure!(!picked.is_empty(), "no tensors selected for {start}..{end}");
picked.sort_by_key(|(_, a, _, _)| *a);
let mut out = serde_json::Map::new();
let mut cursor = 0u64;
let mut slices = Vec::with_capacity(picked.len());
for (name, a, b, meta) in picked {
let len = b - a;
let mut m = meta.clone();
m["data_offsets"] = serde_json::json!([cursor, cursor + len]);
out.insert(name.clone(), m);
slices.push((data_base + a, len));
cursor += len;
}
let mut hjson = serde_json::to_vec(&serde_json::Value::Object(out))?;
while hjson.len() % 8 != 0 {
hjson.push(b' ');
}
let mut header = Vec::with_capacity(8 + hjson.len());
header.extend_from_slice(&(hjson.len() as u64).to_le_bytes());
header.extend_from_slice(&hjson);
let total_bytes = header.len() as u64 + cursor;
Ok(MiniCkpt {
header,
slices,
total_bytes,
})
}
fn load_stage(
ctx: &GpuCtx,
model_dir: &Path,
start: usize,
end: usize,
total: usize,
src_fp: u64,
) -> Result<LoadedStage> {
let w = Weights::load_shard(ctx, model_dir, start, end)?;
let mtp = crate::forward::MtpEngine::new(ctx, &w);
let (stage_first, stage_last) = (w.cfg.stage_first, w.cfg.stage_last);
let gpu = Lfm2Gpu::new(ctx, w);
Ok(LoadedStage {
gpu,
mtp,
plans: Vec::new(),
bpw: None,
start,
end,
stage_first,
stage_last,
total,
src_fp,
})
}
fn cache_entry_meta(entry: &Path) -> Option<(usize, u64)> {
let fp = u64::from_str_radix(
std::fs::read_to_string(entry.join("source.fp"))
.ok()?
.trim(),
16,
)
.ok()?;
if !entry.join("model.safetensors").exists() {
return None;
}
let cfg =
crate::weights::Lfm2Config::from_json(&std::fs::read(entry.join("config.json")).ok()?)
.ok()?;
Some((cfg.n_layers, fp))
}
pub fn run_worker(addr: &str, model_dir: Option<&Path>, opts: WorkerOptions) -> Result<()> {
let ctx = GpuCtx::new()?;
let (total_layers, ckpt_fp) = match model_dir {
Some(d) => (
crate::weights::Lfm2Config::from_json(&std::fs::read(d.join("config.json"))?)?.n_layers,
checkpoint_fingerprint(d)?,
),
None => (0, 0),
};
let cache_root = opts
.cache_dir
.or_else(|| std::env::var_os("OSFKB_SHARD_CACHE").map(std::path::PathBuf::from))
.or_else(|| {
std::env::var_os("HOME").map(|h| std::path::PathBuf::from(h).join(".cache/lfm2-shards"))
})
.unwrap_or_else(|| std::env::temp_dir().join("lfm2-shards"));
let token = opts
.token
.or_else(|| std::env::var("OSFKB_SHARD_TOKEN").ok())
.filter(|t| !t.is_empty());
let vram_mb = opts.vram_gb.map_or(0, |g| (g * 1024.0) as u64);
let auth = if token.is_some() { ", token auth" } else { "" };
match (opts.layers, model_dir) {
(Some((s0, e0)), _) => eprintln!(
"shard worker [{s0}..{e0}) on {addr} backend {} (pinned{auth})",
ctx.backend
),
(None, Some(_)) => eprintln!(
"shard worker AUTO on {addr} backend {} (awaiting OP_LOAD; vram {vram_mb} MB{auth})",
ctx.backend
),
(None, None) => eprintln!(
"shard worker MODELLESS on {addr} backend {} (weights arrive by shipping; \
cache {}; vram {vram_mb} MB{auth})",
ctx.backend,
cache_root.display()
),
}
let loaded = match opts.layers {
Some((s0, e0)) => {
let d = model_dir.context("--layers requires --model")?;
anyhow::ensure!(
e0 > s0 && e0 <= total_layers,
"--layers {s0}:{e0} invalid for depth {total_layers}"
);
Some(load_stage(&ctx, d, s0, e0, total_layers, ckpt_fp)?)
}
None => None,
};
let listener = TcpListener::bind(addr)?;
eprintln!("shard worker ready");
let shared = std::sync::Arc::new(WorkerShared {
ctx,
model_dir: model_dir.map(Path::to_path_buf),
cache_root,
total_layers,
vram_mb,
pinned: opts.layers.is_some(),
ckpt_fp,
token,
shipping: std::sync::Mutex::new(std::collections::HashSet::new()),
eng: std::sync::Mutex::new(StageEngine {
loaded,
chain: None,
}),
});
if let Some(join_addr) = opts.join.clone() {
let data_port: u32 = addr
.rsplit_once(':')
.and_then(|(_, p)| p.parse().ok())
.context("--listen must be host:port")?;
let advertise = opts.advertise.clone().unwrap_or_default();
let sh = std::sync::Arc::clone(&shared);
std::thread::spawn(move || {
loop {
match dial(&join_addr, std::time::Duration::from_secs(5)) {
Ok(mut s) => {
let hello_ok = match sh.token.as_deref() {
Some(tok) => hop_hello(&mut s, tok)
.map_err(|e| eprintln!("shard worker: join auth failed: {e:#}"))
.is_ok(),
None => true,
};
if hello_ok
&& write_msg_bytes(
&mut s,
&[OP_JOIN, data_port, advertise.len() as u32],
advertise.as_bytes(),
)
.is_ok()
{
eprintln!("shard worker: joined coordinator at {join_addr}");
match serve_conn(&sh, s, true) {
Ok(()) => eprintln!("shard worker: session over; re-joining"),
Err(e) => eprintln!("shard worker: session error: {e:#}"),
}
}
}
Err(e) => {
eprintln!("shard worker: join {join_addr} unreachable ({e}); retrying")
}
}
std::thread::sleep(std::time::Duration::from_secs(2));
}
});
}
for conn in listener.incoming() {
let s = conn?;
let sh = std::sync::Arc::clone(&shared);
std::thread::spawn(move || {
let c = match accept_conn(s) {
Ok(c) => c,
Err(e) => {
eprintln!("shard worker: handshake failed: {e:#}");
return;
}
};
match serve_conn(&sh, c, false) {
Ok(()) => eprintln!("shard worker: session ended; awaiting next"),
Err(e) => eprintln!("shard worker: connection error: {e:#}"),
}
});
}
Ok(())
}
struct ShipFile {
entry: std::path::PathBuf,
fname: String,
fp: String,
part: std::path::PathBuf,
w: std::io::BufWriter<std::fs::File>,
fnv: u64,
_claim: ShipClaim,
}
const SHIP_FNV_INIT: u64 = 0xcbf2_9ce4_8422_2325;
const SHIP_CHUNK_BYTES: usize = 8 * 1024 * 1024;
fn resume_state(part: &Path) -> Result<(u64, u64)> {
use std::io::Read as _;
let mut f = match std::fs::File::open(part) {
Ok(f) => f,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok((0, SHIP_FNV_INIT)),
Err(e) => return Err(e.into()),
};
let mut fnv = SHIP_FNV_INIT;
let mut have = 0u64;
let mut buf = vec![0u8; 1 << 20];
loop {
let k = f.read(&mut buf)?;
if k == 0 {
break;
}
fnv = fnv64_update(fnv, &buf[..k]);
have += k as u64;
}
Ok((have, fnv))
}
fn need_loaded(loaded: &mut Option<LoadedStage>) -> Result<&mut LoadedStage> {
loaded.as_mut().ok_or_else(|| {
anyhow!(
"stage not loaded — pin --layers on the worker or connect an auto-split coordinator"
)
})
}
fn hop_hello(s: &mut Conn, token: &str) -> Result<()> {
write_msg_bytes(s, &[OP_HELLO, 0, token.len() as u32], token.as_bytes())?;
let hdr = read_u32s(s, 2)?;
let p = read_f32s(s, hdr[1] as usize)?;
anyhow::ensure!(
p.first().map(|x| x.to_bits()) == Some(0),
"next hop rejected the fleet token"
);
Ok(())
}
fn serve_conn(sh: &std::sync::Arc<WorkerShared>, mut s: Conn, pre_authed: bool) -> Result<()> {
let ctx = &sh.ctx;
let trace = std::env::var_os("OSFKB_STAGE_TRACE").is_some();
let mut dump_hidden = std::env::var("OSFKB_DUMP_HIDDEN").ok().map(|path| {
std::io::BufWriter::new(
std::fs::OpenOptions::new()
.append(true)
.create(true)
.open(path)
.expect("dump file"),
)
});
let (mut bstep_n, mut bstep_recv, mut bstep_gpu, mut bstep_send) = (0u64, 0f64, 0f64, 0f64);
let mut bstep_idle = 0f64;
let mut last_done = std::time::Instant::now();
let (mut donor_steps, mut donor_win_toks) = (0u64, 0u64);
let (mut donor_win_start, mut donor_last_step) =
(std::time::Instant::now(), std::time::Instant::now());
let mut last_stat = std::time::Instant::now();
let mut vocab: Option<crate::vocab::ByteVocab> = None;
let mut recent: std::collections::VecDeque<u32> = std::collections::VecDeque::new();
let mut authed = pre_authed || sh.token.is_none();
let mut ship: Option<ShipFile> = None;
#[cfg(feature = "webrtc")]
let mut peer: Option<crate::webrtc_native::NativeWebrtcPeer> = None;
loop {
let hdr = match read_u32s(&mut s, 3) {
Ok(h) => h,
Err(_) => break, };
let (op, pos, n) = (hdr[0], hdr[1] as usize, hdr[2] as usize);
if op == OP_HELLO {
let raw = read_bytes(&mut s, n)?;
let ok = sh
.token
.as_deref()
.is_none_or(|t| t.as_bytes() == raw.as_slice());
write_msg(&mut s, &[TAG_TOKEN, 1], &[f32::from_bits(u32::from(!ok))])?;
if !ok {
return Err(anyhow!("fleet token rejected"));
}
authed = true;
continue;
}
anyhow::ensure!(
authed,
"unauthenticated connection: this worker requires the fleet token (OSFKB_SHARD_TOKEN)"
);
#[cfg(feature = "webrtc")]
if op == OP_SIGNAL {
let payload = read_bytes(&mut s, n)?;
let (status, body) = match pos as u32 {
SIG_OFFER => crate::webrtc_native::do_offer(&mut peer),
SIG_ANSWER => crate::webrtc_native::do_answer(&mut peer, &payload),
SIG_FINISH => crate::webrtc_native::do_finish(&peer, &payload),
other => (1u32, format!("bad webrtc sub-kind {other}").into_bytes()),
};
write_msg_bytes(&mut s, &[status, body.len() as u32], &body)?;
continue;
}
#[cfg(feature = "webrtc")]
if op == OP_RTC_SERVE {
let _ = read_bytes(&mut s, n)?;
let (status, body) = match crate::webrtc_native::open_pipe_blocking(peer.take()) {
Ok(pipe) => {
let sh2 = std::sync::Arc::clone(sh);
std::thread::spawn(move || {
if let Err(e) = serve_conn(&sh2, Conn::Webrtc(Box::new(pipe)), true) {
eprintln!("shard worker: webrtc inbound session ended: {e:#}");
}
});
(0u32, Vec::new())
}
Err(e) => (1u32, e.into_bytes()),
};
write_msg_bytes(&mut s, &[status, body.len() as u32], &body)?;
continue;
}
#[cfg(feature = "webrtc")]
if op == OP_RTC_CHAIN {
let _ = read_bytes(&mut s, n)?;
let sink = pos & 1 == 1;
let (status, body) = match crate::webrtc_native::open_pipe_blocking(peer.take()) {
Ok(pipe) => {
sh.eng.lock().expect("engine lock poisoned").chain = Some(ChainOut {
s: Conn::Webrtc(Box::new(pipe)),
sink,
});
(0u32, Vec::new())
}
Err(e) => (1u32, e.into_bytes()),
};
write_msg_bytes(&mut s, &[status, body.len() as u32], &body)?;
continue;
}
let mut guard = sh.eng.lock().expect("engine lock poisoned");
let StageEngine { loaded, chain } = &mut *guard;
match op {
OP_BINIT => {
let LoadedStage {
gpu, plans, bpw, ..
} = need_loaded(loaded)?;
let (kcols, groups) = (pos & 0xFFFF, (pos >> 16).max(1));
let spec = std::env::var("OSFKB_MTP_SPEC").ok().as_deref() == Some("1");
*plans = (0..groups)
.map(|_| {
if spec {
gpu.make_batch_plan_spec(ctx, kcols)
} else {
gpu.make_batch_plan(ctx, kcols)
}
})
.collect();
*bpw = if n > 0 {
anyhow::ensure!(
ctx.subgroups,
"coordinator requested wide prefill; this stage's adapter lacks subgroups"
);
Some(gpu.make_batch_plan_wide(ctx, n))
} else {
None
};
write_msg(&mut s, &[TAG_TOKEN, 1], &[f32::from_bits(0)])?;
}
OP_ZSLOT => {
let LoadedStage { gpu, .. } = need_loaded(loaded)?;
gpu.zero_dn_slot(ctx, pos);
if let Some(c) = chain.as_mut().filter(|c| !c.sink) {
write_msg(&mut c.s, &hdr, &[])?;
}
}
OP_BTAB => {
let LoadedStage { gpu, mtp, .. } = need_loaded(loaded)?;
let raw = read_f32s(&mut s, n)?;
let blocks: Vec<u32> = raw.iter().map(|x| x.to_bits()).collect();
gpu.write_btab_row(ctx, pos as u32, &blocks);
if let Some(m) = mtp.as_ref() {
m.gpu.write_btab_row(ctx, pos as u32, &blocks);
}
if let Some(c) = chain.as_mut().filter(|c| !c.sink) {
write_msg(&mut c.s, &hdr, &raw)?;
}
}
OP_BSTEP | OP_BSTEPW => {
let LoadedStage {
gpu,
mtp,
plans,
bpw,
stage_first,
stage_last,
start,
end,
..
} = need_loaded(loaded)?;
let group = (pos >> 16) & 0xFF;
let spec_k = (pos >> 24) & 0x3F;
let spec_on = (pos >> 30) & 1 == 1;
let spec_verify = (pos >> 31) & 1 == 1;
let pos = pos & 0xFFFF;
if trace && last_done.elapsed().as_secs_f64() < 3600.0 {
bstep_idle += last_done.elapsed().as_secs_f64();
}
let t0 = std::time::Instant::now();
let colsu = read_u32s(&mut s, 4 * pos)?;
let hidden = read_f32s(&mut s, n)?;
let t_recv = t0.elapsed();
let plan = if op == OP_BSTEPW {
bpw.as_ref()
.ok_or_else(|| anyhow!("OP_BSTEPW without a wide plan"))?
} else {
plans
.get(group)
.ok_or_else(|| anyhow!("OP_BSTEP group {group} out of range"))?
};
let cols: Vec<BatchCol> = colsu
.chunks(4)
.map(|c| BatchCol::text(c[3], c[0], c[1], c[2] != 0))
.collect();
let hin = (!*stage_first).then_some(hidden.as_slice());
let t1 = std::time::Instant::now();
let out = gpu.batch_stage_step(ctx, plan, &cols, hin)?;
let t_gpu = t1.elapsed();
if *stage_last && let Some(dump) = &mut dump_hidden {
let hbuf = gpu.read_cur_for_tests(ctx, plan, cols.len())?;
let hsz = hbuf.len() / cols.len();
for (p, col) in cols.iter().enumerate() {
use std::io::Write;
dump.write_all(&col.pos.to_le_bytes())?;
dump.write_all(&col.btrow.to_le_bytes())?;
dump.write_all(&col.token.to_le_bytes())?;
dump.write_all(bytemuck::cast_slice(&hbuf[p * hsz..(p + 1) * hsz]))?;
}
}
if trace && op == OP_BSTEPW {
let br = gpu.take_step_breakdown();
eprintln!(
"bstepw ncols={pos} recv {:.1}ms gpu {:.1}ms [prep {:.1} enc {:.1} poll {:.1}] idle-before {:.1}ms",
t_recv.as_secs_f64() * 1e3,
t_gpu.as_secs_f64() * 1e3,
br.0 * 1e3,
br.1 * 1e3,
br.2 * 1e3,
last_done.elapsed().as_secs_f64() * 1e3
- t_gpu.as_secs_f64() * 1e3
- t_recv.as_secs_f64() * 1e3,
);
}
if trace && pos <= 8 {
let br = gpu.take_step_breakdown();
eprintln!(
"bstep small ncols={pos}: total {:.2}ms [prep {:.2} encode {:.2} poll+read {:.2}]",
t_gpu.as_secs_f64() * 1e3,
br.0 * 1e3,
br.1 * 1e3,
br.2 * 1e3
);
}
let t2 = std::time::Instant::now();
match out {
StageBatchOut::Hidden(h) => match chain.as_mut() {
None => write_msg(&mut s, &[TAG_HIDDEN, h.len() as u32], &h)?,
Some(c) if !c.sink => {
let mut fh = Vec::with_capacity(3 + colsu.len());
fh.extend_from_slice(&[op, hdr[1], h.len() as u32]);
fh.extend_from_slice(&colsu);
write_msg(&mut c.s, &fh, &h)?;
}
Some(_) => {
return Err(anyhow!("non-final stage chained to the return sink"));
}
},
StageBatchOut::Tokens(t) => {
let mut resp: Vec<u32> = t[..pos].to_vec();
if (spec_verify || spec_on) && mtp.is_some() {
let m = mtp.as_ref().expect("checked");
let plan = plans
.get(group)
.ok_or_else(|| anyhow!("spec group out of range"))?;
let mut reqs: Vec<crate::forward::DraftReq> = Vec::new();
let mut seeds: Vec<(u32, usize)> = Vec::new();
let mut pairs: Vec<crate::forward::DraftPair> = Vec::new();
let mut i = 0usize;
while i < cols.len() {
let mut j = i + 1;
while j < cols.len() && cols[j].btrow == cols[i].btrow {
j += 1;
}
let slot = cols[i].btrow;
if spec_verify {
let k = j - i - 1;
let mut a = 0usize;
while a < k && t[i + a] == cols[i + 1 + a].token {
a += 1;
}
seeds.push((slot, i + a));
reqs.push(crate::forward::DraftReq {
btrow: slot,
first_token: t[i + a],
pos0: cols[i + a].pos as usize,
mpos0: cols[i + a].mpos,
});
} else {
for p in i..j.saturating_sub(1) {
pairs.push(crate::forward::DraftPair {
btrow: slot,
cur_col: p,
token: cols[p + 1].token,
pos: cols[p].pos as usize,
mpos: cols[p + 1].mpos,
});
}
let lastc = j - 1;
if cols[lastc].need_logit {
seeds.push((slot, lastc));
reqs.push(crate::forward::DraftReq {
btrow: slot,
first_token: t[lastc],
pos0: cols[lastc].pos as usize,
mpos0: cols[lastc].mpos,
});
}
}
i = j;
}
for chunk in pairs.chunks(crate::forward::MtpEngine::K_DRAFT_BATCH) {
m.draft_pairs(ctx, plan, chunk)?;
}
if spec_k > 0 {
for (chunk, seed_chunk) in reqs
.chunks(crate::forward::MtpEngine::K_DRAFT_BATCH)
.zip(seeds.chunks(crate::forward::MtpEngine::K_DRAFT_BATCH))
{
let drafts = m.draft_chain_batch_seeded(
ctx,
Some(plan),
seed_chunk,
chunk,
spec_k,
)?;
resp.extend(drafts.into_iter().flatten());
}
} else {
for (slot, col) in &seeds {
m.seed_slot_from_col(ctx, plan, *col, *slot as usize);
}
}
}
let bits: Vec<f32> = resp.iter().map(|x| f32::from_bits(*x)).collect();
match chain.as_mut() {
None => {
write_msg(&mut s, &[TAG_TOKEN, bits.len() as u32], &bits)?;
}
Some(c) if c.sink => {
write_msg(
&mut c.s,
&[TAG_DONE, group as u32, bits.len() as u32],
&bits,
)?;
}
Some(_) => {
return Err(anyhow!("final stage must chain to the return sink"));
}
}
}
}
last_done = std::time::Instant::now();
if donor_last_step.elapsed().as_secs_f64() > 0.5 {
donor_win_toks = 0;
donor_win_start = std::time::Instant::now();
}
donor_last_step = std::time::Instant::now();
donor_steps += 1;
donor_win_toks += pos as u64;
for c in &cols {
recent.push_back(c.token);
}
while recent.len() > 64 {
recent.pop_front();
}
if last_stat.elapsed().as_secs_f64() >= 2.0
&& donor_win_start.elapsed().as_secs_f64() >= 0.25
{
let rate = donor_win_toks as f64 / donor_win_start.elapsed().as_secs_f64();
let tail = vocab.as_ref().map_or(String::new(), |v| {
let text = v.decode(&recent.iter().copied().collect::<Vec<_>>());
let chars: Vec<char> = text.chars().collect();
chars[chars.len().saturating_sub(48)..]
.iter()
.collect::<String>()
.replace('\n', " ")
});
eprintln!(
"shard worker: serving {}..{} · {donor_steps} steps · {rate:.1} tok/s · \"{tail}\"",
*start, *end
);
donor_win_toks = 0;
donor_win_start = std::time::Instant::now();
last_stat = std::time::Instant::now();
}
if trace {
bstep_n += 1;
bstep_recv += t_recv.as_secs_f64();
bstep_gpu += t_gpu.as_secs_f64();
bstep_send += t2.elapsed().as_secs_f64();
if bstep_n % 50 == 0 {
eprintln!(
"stage trace [{bstep_n}]: recv {:.1}ms gpu+sync {:.1}ms send {:.1}ms idle {:.1}ms (avg/step, ncols {})",
bstep_recv * 20.0,
bstep_gpu * 20.0,
bstep_send * 20.0,
bstep_idle * 20.0,
pos
);
bstep_recv = 0.0;
bstep_gpu = 0.0;
bstep_send = 0.0;
bstep_idle = 0.0;
}
}
}
OP_MTP_BATCH => {
let LoadedStage { mtp, .. } = need_loaded(loaded)?;
let (ncols, kdraft) = (pos, n);
let payload = read_f32s(&mut s, ncols * 3)?;
let m = mtp
.as_ref()
.ok_or_else(|| anyhow!("OP_MTP_BATCH without an MTP head"))?;
let reqs: Vec<(u32, u32, usize)> = payload
.chunks(3)
.map(|c| (c[0].to_bits(), c[1].to_bits(), c[2].to_bits() as usize))
.collect();
let drafts = m.draft_chain_batch(ctx, &reqs, kdraft)?;
let flat: Vec<f32> = drafts
.iter()
.flat_map(|d| d.iter().map(|x| f32::from_bits(*x)))
.collect();
write_msg(&mut s, &[TAG_TOKEN, flat.len() as u32], &flat)?;
}
OP_MTP_SEED => {
let LoadedStage { mtp, plans, .. } = need_loaded(loaded)?;
let (col, group) = (n & 0xFFFF, (n >> 16) & 0xFFFF);
let m = mtp
.as_ref()
.ok_or_else(|| anyhow!("OP_MTP_SEED without an MTP head"))?;
let plan = plans
.get(group)
.ok_or_else(|| anyhow!("OP_MTP_SEED group {group} out of range"))?;
m.seed_slot_from_col(ctx, plan, col, pos);
}
OP_MTP => {
let LoadedStage { mtp, plans, .. } = need_loaded(loaded)?;
let payload = read_f32s(&mut s, 1)?;
let first_tok = payload[0].to_bits();
let (col_j, kdraft) = (n >> 8, n & 0xff);
let m = mtp
.as_ref()
.ok_or_else(|| anyhow!("OP_MTP without an MTP head on this stage"))?;
let plan = plans
.first()
.ok_or_else(|| anyhow!("OP_MTP before OP_BINIT"))?;
let drafts = m.draft_chain_seeded(ctx, plan, col_j, first_tok, pos, kdraft)?;
let mut resp = drafts.clone();
if std::env::var("OSFKB_MTP_BETA_PROBE").is_ok() {
for t3 in m.last_top3() {
resp.extend_from_slice(&t3);
}
}
let bits: Vec<f32> = resp.iter().map(|x| f32::from_bits(*x)).collect();
write_msg(&mut s, &[TAG_TOKEN, bits.len() as u32], &bits)?;
}
OP_DNRESTORE => {
let LoadedStage { gpu, plans, .. } = need_loaded(loaded)?;
let (col, group) = (n & 0xFFFF, (n >> 16) & 0xFFFF);
let plan = plans
.get(group)
.ok_or_else(|| anyhow!("OP_DNRESTORE group {group} out of range"))?;
gpu.dn_restore(ctx, plan, pos as u32, col);
if let Some(c) = chain.as_mut().filter(|c| !c.sink) {
write_msg(&mut c.s, &hdr, &[])?;
}
}
OP_STEP => {
let LoadedStage { gpu, .. } = need_loaded(loaded)?;
let hidden = read_f32s(&mut s, n)?;
match gpu.stage_forward(ctx, pos, 0, Some(&hidden))? {
StageOut::Hidden(h) => match chain.as_mut() {
None => write_msg(&mut s, &[TAG_HIDDEN, h.len() as u32], &h)?,
Some(c) if !c.sink => {
write_msg(&mut c.s, &[OP_STEP, pos as u32, h.len() as u32], &h)?;
}
Some(_) => {
return Err(anyhow!("non-final stage chained to the return sink"));
}
},
StageOut::Token(t) => match chain.as_mut() {
None => write_msg(&mut s, &[TAG_TOKEN, 1], &[f32::from_bits(t)])?,
Some(c) if c.sink => {
write_msg(&mut c.s, &[TAG_DONE, 0, 1], &[f32::from_bits(t)])?;
}
Some(_) => {
return Err(anyhow!("final stage must chain to the return sink"));
}
},
}
}
OP_RESET => {
let LoadedStage { gpu, .. } = need_loaded(loaded)?;
gpu.reset(ctx);
write_msg(&mut s, &[TAG_TOKEN, 1], &[f32::from_bits(0)])?;
}
OP_CHAIN => {
let raw = read_bytes(&mut s, n)?;
*chain = None; let ok = if raw.is_empty() {
true
} else {
let want = String::from_utf8(raw).context("chain addr utf8")?;
let addr = if let Some(port) = want.strip_prefix(':') {
std::net::SocketAddr::new(s.peer_addr()?.ip(), port.parse()?).to_string()
} else {
want
};
match dial(&addr, std::time::Duration::from_secs(5)) {
Ok(mut t) => {
let hop_ok = match (sh.token.as_deref(), pos & 1 == 1) {
(Some(tok), true) => write_msg_bytes(
&mut t,
&[OP_HELLO, 0, tok.len() as u32],
tok.as_bytes(),
)
.is_ok(),
(Some(tok), false) => hop_hello(&mut t, tok)
.map_err(|e| {
eprintln!("shard worker: hop auth {addr} failed: {e:#}");
})
.is_ok(),
(None, _) => true,
};
if hop_ok {
*chain = Some(ChainOut {
s: t,
sink: pos & 1 == 1,
});
}
hop_ok
}
Err(err) => {
eprintln!("shard worker: chain dial {addr} failed: {err}");
false
}
}
};
write_msg(&mut s, &[TAG_TOKEN, 1], &[f32::from_bits(u32::from(!ok))])?;
}
OP_LOAD => {
let (start, end) = (pos, n);
let fpx = String::from_utf8(read_bytes(&mut s, 16)?).context("load fp utf8")?;
let want_fp = u64::from_str_radix(&fpx, 16).context("load fp hex")?;
let ack: u32 = if loaded
.as_ref()
.is_some_and(|l| l.start == start && l.end == end && l.src_fp == want_fp)
{
eprintln!("shard worker: layers {start}..{end} already resident — no reload");
0
} else {
let src: std::result::Result<(std::path::PathBuf, usize, u64), u32> =
if let Some(md) = &sh.model_dir {
if sh.ckpt_fp == want_fp {
Ok((md.clone(), sh.total_layers, sh.ckpt_fp))
} else {
eprintln!(
"shard worker: OP_LOAD wants {fpx} but --model is \
{:016x} — wrong checkpoint",
sh.ckpt_fp
);
Err(3)
}
} else {
let entry = sh.cache_root.join(&fpx).join(format!("{start}-{end}"));
match cache_entry_meta(&entry) {
Some((total, fp)) if fp == want_fp => Ok((entry, total, fp)),
_ => Err(2), }
};
match src {
Err(code) => code,
Ok((dir, total, fp)) if end > start && end <= total => {
*loaded = None; match load_stage(ctx, &dir, start, end, total, fp) {
Ok(l) => {
eprintln!("shard worker: loaded layers {start}..{end}");
*loaded = Some(l);
0
}
Err(e) => {
eprintln!("shard worker: OP_LOAD {start}..{end} failed: {e:#}");
1
}
}
}
Ok((_, total, _)) => {
eprintln!(
"shard worker: OP_LOAD {start}..{end} invalid for depth {total}"
);
1
}
}
};
write_msg(&mut s, &[TAG_TOKEN, 1], &[f32::from_bits(ack)])?;
}
OP_SHIP_BEGIN => {
let meta = String::from_utf8(read_bytes(&mut s, n)?).context("ship meta utf8")?;
let parts: Vec<&str> = meta.splitn(5, '|').collect();
anyhow::ensure!(parts.len() >= 4, "malformed ship meta: {meta}");
let entry = sh.cache_root.join(parts[0]).join(format!(
"{}-{}",
parts[1].parse::<usize>()?,
parts[2].parse::<usize>()?
));
std::fs::create_dir_all(&entry)?;
let fname = parts[3].to_string();
anyhow::ensure!(
fname == "config.json" || fname == "model.safetensors",
"unexpected ship filename {fname}"
);
let part = entry.join(format!("{fname}.part"));
let claimed = sh
.shipping
.lock()
.expect("shipping set poisoned")
.insert(part.clone());
if !claimed {
eprintln!("shard worker: ship {fname} busy (another shipper holds {part:?})");
write_msg(&mut s, &[u32::MAX, u32::MAX, 0, 0], &[])?;
continue;
}
let claim = ShipClaim {
sh: std::sync::Arc::clone(sh),
key: part.clone(),
};
let (have, fnv) = resume_state(&part)?;
let f = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&part)?;
eprintln!(
"shard worker: ship begin {} → {} (resume at {:.2} GB)",
fname,
entry.display(),
have as f64 / 1e9
);
ship = Some(ShipFile {
entry,
fname,
fp: parts[0].to_string(),
part,
w: std::io::BufWriter::new(f),
fnv,
_claim: claim,
});
write_msg(
&mut s,
&[
have as u32,
(have >> 32) as u32,
fnv as u32,
(fnv >> 32) as u32,
],
&[],
)?;
}
OP_SHIP_CHUNK => {
let sf = ship
.as_mut()
.ok_or_else(|| anyhow!("chunk without begin"))?;
let raw = read_bytes(&mut s, n)?;
sf.fnv = fnv64_update(sf.fnv, &raw);
std::io::Write::write_all(&mut sf.w, &raw)?;
std::io::Write::flush(&mut sf.w)?;
write_msg(&mut s, &[TAG_TOKEN, 0], &[])?;
}
OP_SHIP_END => {
let sf = ship.take().ok_or_else(|| anyhow!("end without begin"))?;
let want = (pos as u64) | ((n as u64) << 32);
let ok = sf.fnv == want;
if ok {
let mut w = sf.w;
std::io::Write::flush(&mut w)?;
drop(w);
std::fs::rename(&sf.part, sf.entry.join(&sf.fname))?;
if sf.fname == "model.safetensors" {
std::fs::write(sf.entry.join("source.fp"), &sf.fp)?;
}
eprintln!("shard worker: ship done {}", sf.fname);
} else {
eprintln!(
"shard worker: ship {} FAILED integrity ({:016x} != {:016x})",
sf.fname, sf.fnv, want
);
let _ = std::fs::remove_file(&sf.part);
}
write_msg(&mut s, &[TAG_TOKEN, 1], &[f32::from_bits(u32::from(!ok))])?;
}
OP_INFO => {
let (start, end, total, fp) = loaded
.as_ref()
.map_or((0, 0, sh.total_layers, sh.ckpt_fp), |l| {
(l.start, l.end, l.total, l.src_fp)
});
let desc = format!(
"{start}|{end}|{total}|{}|{}|{}|{fp:016x}|{}",
u8::from(ctx.subgroups),
sh.vram_mb,
u8::from(sh.pinned),
ctx.backend
);
write_msg_bytes(&mut s, &[TAG_INFO, desc.len() as u32], desc.as_bytes())?;
}
OP_VOCAB => {
let raw = read_bytes(&mut s, n)?;
vocab = crate::vocab::ByteVocab::from_blob(&raw);
}
OP_SIGNAL => {
let _payload = read_bytes(&mut s, n)?;
let msg: &[u8] = b"webrtc endpoint not built on this worker";
write_msg_bytes(&mut s, &[1, msg.len() as u32], msg)?;
}
OP_SHUTDOWN => {
*chain = None;
break;
}
other => return Err(anyhow!("bad op {other}")),
}
}
Ok(())
}
pub struct ShardClient {
s: Conn,
}
impl ShardClient {
fn from_stream(s: Conn) -> Self {
Self { s }
}
pub fn connect(addr: &str) -> Result<Self> {
let s = dial(addr, std::time::Duration::from_secs(5))
.with_context(|| format!("connect {addr}"))?;
let mut c = Self { s };
if let Ok(t) = std::env::var("OSFKB_SHARD_TOKEN")
&& !t.is_empty()
{
hop_hello(&mut c.s, &t).with_context(|| format!("authenticate to {addr}"))?;
}
Ok(c)
}
pub fn load_send(&mut self, start: usize, end: usize, fp: u64) -> Result<()> {
write_msg_bytes(
&mut self.s,
&[OP_LOAD, start as u32, end as u32],
format!("{fp:016x}").as_bytes(),
)
}
pub fn load_wait(&mut self) -> Result<LoadAck> {
let hdr = read_u32s(&mut self.s, 2)?;
let p = read_f32s(&mut self.s, hdr[1] as usize)?;
Ok(match p.first().map(|x| x.to_bits()) {
Some(0) => LoadAck::Ok,
Some(2) => LoadAck::NeedWeights,
Some(3) => LoadAck::WrongCheckpoint,
_ => LoadAck::Failed,
})
}
fn ship_begin(
&mut self,
fpx: &str,
range: (usize, usize),
fname: &str,
total: u64,
) -> Result<(u64, u64)> {
let meta = format!("{fpx}|{}|{}|{fname}|{total}", range.0, range.1);
write_msg_bytes(
&mut self.s,
&[OP_SHIP_BEGIN, 0, meta.len() as u32],
meta.as_bytes(),
)?;
let r = read_u32s(&mut self.s, 4)?;
let have = r[0] as u64 | ((r[1] as u64) << 32);
let fnv = r[2] as u64 | ((r[3] as u64) << 32);
anyhow::ensure!(
have != u64::MAX,
"ship target {fname} busy on the stage (another coordinator is shipping it); retry"
);
Ok((have, fnv))
}
fn ship_chunk(&mut self, fnv: &mut u64, raw: &[u8]) -> Result<()> {
*fnv = fnv64_update(*fnv, raw);
write_msg_bytes(&mut self.s, &[OP_SHIP_CHUNK, 0, raw.len() as u32], raw)?;
let _ack = read_u32s(&mut self.s, 2)?; Ok(())
}
fn ship_end(&mut self, fnv: u64) -> Result<()> {
write_msg(
&mut self.s,
&[OP_SHIP_END, fnv as u32, (fnv >> 32) as u32],
&[],
)?;
let hdr = read_u32s(&mut self.s, 2)?;
let p = read_f32s(&mut self.s, hdr[1] as usize)?;
anyhow::ensure!(
p.first().map(|x| x.to_bits()) == Some(0),
"stage rejected the shipped file (integrity)"
);
Ok(())
}
pub(crate) fn ship_bytes(
&mut self,
fpx: &str,
range: (usize, usize),
fname: &str,
bytes: &[u8],
) -> Result<()> {
let (have, mut fnv) = self.ship_begin(fpx, range, fname, bytes.len() as u64)?;
let have = (have as usize).min(bytes.len());
for c in bytes[have..].chunks(SHIP_CHUNK_BYTES) {
self.ship_chunk(&mut fnv, c)?;
}
self.ship_end(fnv)
}
pub(crate) fn ship_mini_ckpt(
&mut self,
fpx: &str,
range: (usize, usize),
src: &Path,
plan: &MiniCkpt,
) -> Result<()> {
use std::io::{Read as _, Seek as _};
let (have0, mut fnv) =
self.ship_begin(fpx, range, "model.safetensors", plan.total_bytes)?;
let have = if have0 > plan.total_bytes { 0 } else { have0 };
if have > 0 {
eprintln!(
"ship {:?}: resuming at {:.2} / {:.2} GB",
range,
have as f64 / 1e9,
plan.total_bytes as f64 / 1e9
);
}
let mut f = std::fs::File::open(src)?;
let mut buf = vec![0u8; SHIP_CHUNK_BYTES];
let mut vpos = 0u64; let (mut sent, mut mark) = (have, have);
let hlen = plan.header.len() as u64;
if have < hlen {
for c in plan.header[have as usize..].chunks(SHIP_CHUNK_BYTES) {
self.ship_chunk(&mut fnv, c)?;
sent += c.len() as u64;
}
}
vpos += hlen;
for (off, len) in &plan.slices {
let vstart = vpos;
vpos += len;
if have >= vpos {
continue; }
let skip = have.saturating_sub(vstart); f.seek(std::io::SeekFrom::Start(off + skip))?;
let mut left = len - skip;
while left > 0 {
let take = left.min(buf.len() as u64) as usize;
f.read_exact(&mut buf[..take])?;
self.ship_chunk(&mut fnv, &buf[..take])?;
left -= take as u64;
sent += take as u64;
if sent - mark >= 512 * 1024 * 1024 {
mark = sent;
eprintln!(
"ship {:?}: {:.1} / {:.1} GB",
range,
sent as f64 / 1e9,
plan.total_bytes as f64 / 1e9
);
}
}
}
self.ship_end(fnv)
}
pub fn step(&mut self, pos: usize, hidden: &[f32]) -> Result<StageOut> {
write_msg(
&mut self.s,
&[OP_STEP, pos as u32, hidden.len() as u32],
hidden,
)?;
let hdr = read_u32s(&mut self.s, 2)?;
let (tag, n) = (hdr[0], hdr[1] as usize);
let payload = read_f32s(&mut self.s, n)?;
Ok(match tag {
TAG_TOKEN => StageOut::Token(payload[0].to_bits()),
_ => StageOut::Hidden(payload),
})
}
pub fn reset(&mut self) -> Result<()> {
write_msg(&mut self.s, &[OP_RESET, 0, 0], &[])?;
let _ = read_u32s(&mut self.s, 2)?;
let _ = read_f32s(&mut self.s, 1)?;
Ok(())
}
pub fn shutdown(&mut self) {
let _ = write_msg(&mut self.s, &[OP_SHUTDOWN, 0, 0], &[]);
}
pub fn binit(&mut self, k: usize, kw: usize) -> Result<()> {
self.binit_grouped(k, kw, 1)
}
pub fn binit_grouped(&mut self, k: usize, kw: usize, n_groups: usize) -> Result<()> {
assert!(k <= 0xFFFF && n_groups <= 0xFFFF);
write_msg(
&mut self.s,
&[OP_BINIT, (k | (n_groups << 16)) as u32, kw as u32],
&[],
)?;
let _ = read_u32s(&mut self.s, 2)?;
let _ = read_f32s(&mut self.s, 1)?;
Ok(())
}
pub fn zslot(&mut self, slot: u32) -> Result<()> {
write_msg(&mut self.s, &[OP_ZSLOT, slot, 0], &[])
}
pub fn send_vocab(&mut self, blob: &[u8]) -> Result<()> {
write_msg_bytes(&mut self.s, &[OP_VOCAB, 0, blob.len() as u32], blob)
}
pub fn webrtc_offer(&mut self) -> Result<Vec<u8>> {
write_msg_bytes(&mut self.s, &[OP_SIGNAL, SIG_OFFER, 0], &[])?;
self.read_signal_reply("offer")
}
pub fn webrtc_answer(&mut self, offer: &[u8]) -> Result<Vec<u8>> {
write_msg_bytes(
&mut self.s,
&[OP_SIGNAL, SIG_ANSWER, offer.len() as u32],
offer,
)?;
self.read_signal_reply("answer")
}
pub fn webrtc_finish(&mut self, answer: &[u8]) -> Result<()> {
write_msg_bytes(
&mut self.s,
&[OP_SIGNAL, SIG_FINISH, answer.len() as u32],
answer,
)?;
self.read_signal_reply("finish").map(|_| ())
}
pub fn rtc_serve(&mut self) -> Result<()> {
write_msg_bytes(&mut self.s, &[OP_RTC_SERVE, 0, 0], &[])?;
self.read_signal_reply("rtc-serve").map(|_| ())
}
pub fn rtc_chain(&mut self, sink: bool) -> Result<()> {
write_msg_bytes(&mut self.s, &[OP_RTC_CHAIN, u32::from(sink), 0], &[])?;
self.read_signal_reply("rtc-chain").map(|_| ())
}
fn read_signal_reply(&mut self, step: &str) -> Result<Vec<u8>> {
let hdr = read_u32s(&mut self.s, 2)?;
let body = read_bytes(&mut self.s, hdr[1] as usize)?;
anyhow::ensure!(
hdr[0] == 0,
"webrtc {step} rejected by peer: {}",
String::from_utf8_lossy(&body)
);
Ok(body)
}
pub fn chain(&mut self, next: &str, sink: bool) -> Result<()> {
write_msg_bytes(
&mut self.s,
&[OP_CHAIN, u32::from(sink), next.len() as u32],
next.as_bytes(),
)?;
let hdr = read_u32s(&mut self.s, 2)?;
let payload = read_f32s(&mut self.s, hdr[1] as usize)?;
anyhow::ensure!(
payload.first().map(|x| x.to_bits()) == Some(0),
"stage failed to dial its next hop {next}"
);
Ok(())
}
pub fn chain_clear(&mut self) -> Result<()> {
write_msg_bytes(&mut self.s, &[OP_CHAIN, 0, 0], &[])?;
let hdr = read_u32s(&mut self.s, 2)?;
let _ = read_f32s(&mut self.s, hdr[1] as usize)?;
Ok(())
}
pub fn info(&mut self) -> Result<StageInfo> {
write_msg(&mut self.s, &[OP_INFO, 0, 0], &[])?;
let hdr = read_u32s(&mut self.s, 2)?;
anyhow::ensure!(hdr[0] == TAG_INFO, "info expects TAG_INFO");
let text = String::from_utf8(read_bytes(&mut self.s, hdr[1] as usize)?)
.context("stage info utf8")?;
let parts: Vec<&str> = text.splitn(8, '|').collect();
anyhow::ensure!(parts.len() == 8, "malformed stage info: {text}");
Ok(StageInfo {
start: parts[0].parse()?,
end: parts[1].parse()?,
total_layers: parts[2].parse()?,
subgroups: parts[3] == "1",
vram_mb: parts[4].parse()?,
pinned: parts[5] == "1",
ckpt_fp: u64::from_str_radix(parts[6], 16).context("ckpt fingerprint hex")?,
backend: parts[7].to_string(),
})
}
pub fn step_send(&mut self, pos: usize, hidden: &[f32]) -> Result<()> {
write_msg(
&mut self.s,
&[OP_STEP, pos as u32, hidden.len() as u32],
hidden,
)
}
pub fn mtp_draft(
&mut self,
col_j: usize,
first_tok: u32,
pos0: usize,
k: usize,
) -> Result<Vec<u32>> {
write_msg(
&mut self.s,
&[OP_MTP, pos0 as u32, ((col_j << 8) | k) as u32],
&[f32::from_bits(first_tok)],
)?;
let hdr = read_u32s(&mut self.s, 2)?;
anyhow::ensure!(hdr[0] == TAG_TOKEN, "mtp_draft expects tokens");
Ok(read_f32s(&mut self.s, hdr[1] as usize)?
.iter()
.map(|x| x.to_bits())
.collect())
}
pub fn mtp_draft_batch(
&mut self,
reqs: &[(u32, u32, usize)],
k: usize,
) -> Result<Vec<Vec<u32>>> {
let payload: Vec<f32> = reqs
.iter()
.flat_map(|(s, t, p)| {
[
f32::from_bits(*s),
f32::from_bits(*t),
f32::from_bits(*p as u32),
]
})
.collect();
write_msg(
&mut self.s,
&[OP_MTP_BATCH, reqs.len() as u32, k as u32],
&payload,
)?;
let hdr = read_u32s(&mut self.s, 2)?;
anyhow::ensure!(hdr[0] == TAG_TOKEN, "mtp_draft_batch expects tokens");
let flat = read_f32s(&mut self.s, hdr[1] as usize)?;
Ok(flat
.chunks(k)
.map(|c| c.iter().map(|x| x.to_bits()).collect())
.collect())
}
pub fn mtp_seed(&mut self, slot: u32, col: usize) -> Result<()> {
self.mtp_seed_grouped(slot, col, 0)
}
pub fn mtp_seed_grouped(&mut self, slot: u32, col: usize, group: usize) -> Result<()> {
write_msg(
&mut self.s,
&[OP_MTP_SEED, slot, (col | (group << 16)) as u32],
&[],
)
}
pub fn dn_restore(&mut self, slot: u32, col: usize) -> Result<()> {
self.dn_restore_grouped(slot, col, 0)
}
pub fn dn_restore_grouped(&mut self, slot: u32, col: usize, group: usize) -> Result<()> {
write_msg(
&mut self.s,
&[OP_DNRESTORE, slot, (col | (group << 16)) as u32],
&[],
)
}
pub fn btab(&mut self, row: u32, blocks: &[u32]) -> Result<()> {
let bits: Vec<f32> = blocks.iter().map(|b| f32::from_bits(*b)).collect();
write_msg(&mut self.s, &[OP_BTAB, row, blocks.len() as u32], &bits)
}
pub fn bstep(
&mut self,
cols: &[(u32, u32, u32, u32)],
hidden: &[f32],
wide: bool,
) -> Result<StageBatchOut> {
self.bstep_grouped(cols, hidden, wide, 0, (false, false, 0))
}
pub fn bstep_grouped(
&mut self,
cols: &[(u32, u32, u32, u32)],
hidden: &[f32],
wide: bool,
group: usize,
spec: (bool, bool, usize),
) -> Result<StageBatchOut> {
self.bstep_send_grouped(cols, hidden, wide, group, spec)?;
let hdr = read_u32s(&mut self.s, 2)?;
let (tag, n) = (hdr[0], hdr[1] as usize);
let payload = read_f32s(&mut self.s, n)?;
Ok(match tag {
TAG_TOKEN => StageBatchOut::Tokens(payload.iter().map(|x| x.to_bits()).collect()),
_ => StageBatchOut::Hidden(payload),
})
}
pub fn bstep_send_grouped(
&mut self,
cols: &[(u32, u32, u32, u32)],
hidden: &[f32],
wide: bool,
group: usize,
spec: (bool, bool, usize),
) -> Result<()> {
let (spec_on, spec_verify, spec_k) = spec;
assert!(cols.len() <= 0xFFFF && group <= 0xFF && spec_k <= 0x3F);
let op = if wide { OP_BSTEPW } else { OP_BSTEP };
let mut header = Vec::with_capacity(3 + cols.len() * 4);
header.extend_from_slice(&[
op,
(cols.len()
| (group << 16)
| (spec_k << 24)
| (usize::from(spec_on) << 30)
| (usize::from(spec_verify) << 31)) as u32,
hidden.len() as u32,
]);
for &(pos, row, nl, tok) in cols {
header.extend_from_slice(&[pos, row, nl, tok]);
}
write_msg(&mut self.s, &header, hidden)
}
}
pub fn webrtc_pair(a: &mut ShardClient, b: &mut ShardClient) -> Result<()> {
let offer = a.webrtc_offer()?;
let answer = b.webrtc_answer(&offer)?;
a.webrtc_finish(&answer)?;
Ok(())
}
pub fn webrtc_chain(from: &mut ShardClient, to: &mut ShardClient, sink: bool) -> Result<()> {
webrtc_pair(from, to)?; to.rtc_serve()?; from.rtc_chain(sink)?; Ok(())
}
fn webrtc_hop(clients: &mut [ShardClient], i: usize) -> Result<()> {
#[cfg(feature = "webrtc")]
{
let (head, tail) = clients.split_at_mut(i + 1);
webrtc_chain(&mut head[i], &mut tail[0], false)
}
#[cfg(not(feature = "webrtc"))]
{
let _ = (clients, i);
anyhow::bail!("a WebRTC hop was requested but this build lacks the `webrtc` feature")
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LoadAck {
Ok,
NeedWeights,
WrongCheckpoint,
Failed,
}
#[derive(Debug, Clone)]
pub struct StageInfo {
pub start: usize,
pub end: usize,
pub total_layers: usize,
pub subgroups: bool,
pub vram_mb: u64,
pub pinned: bool,
pub ckpt_fp: u64,
pub backend: String,
}
pub fn validate_chain(split: usize, infos: &[StageInfo]) -> Result<String> {
anyhow::ensure!(!infos.is_empty(), "need at least one worker stage");
let total = infos[0].total_layers;
let mut cursor = split;
let mut map = if split > 0 {
format!("stage 0: coordinator, layers 0..{split} (local)")
} else {
"stage 0: coordinator, headless (tokenizer + scheduler only)".to_string()
};
for (i, inf) in infos.iter().enumerate() {
anyhow::ensure!(
inf.total_layers == total,
"stage {}: model depth {} != {total} — different checkpoints across the fleet?",
i + 1,
inf.total_layers
);
anyhow::ensure!(
inf.ckpt_fp == infos[0].ckpt_fp,
"stage {}: checkpoint fingerprint {:016x} != stage 1's {:016x} — same architecture, \
DIFFERENT weights (wrong model dir on that host?)",
i + 1,
inf.ckpt_fp,
infos[0].ckpt_fp
);
anyhow::ensure!(
inf.start == cursor,
"stage {} covers layers {}..{} but the chain is at layer {cursor} — gap or overlap",
i + 1,
inf.start,
inf.end
);
anyhow::ensure!(
inf.end > inf.start && inf.end <= total,
"stage {} range {}..{} is invalid for depth {total}",
i + 1,
inf.start,
inf.end
);
map.push_str(&format!(
"\nstage {}: layers {}..{} on {}",
i + 1,
inf.start,
inf.end,
inf.backend
));
cursor = inf.end;
}
anyhow::ensure!(
cursor == total,
"chain ends at layer {cursor} but the model has {total} — the last stage must own the head"
);
Ok(map)
}
pub fn plan_split(split: usize, total: usize, vram_mb: &[u64]) -> Result<Vec<(usize, usize)>> {
anyhow::ensure!(!vram_mb.is_empty(), "no worker stages to plan");
anyhow::ensure!(
split < total,
"coordinator owns 0..{split} of {total} layers — nothing left for the workers"
);
let (remaining, n) = (total - split, vram_mb.len());
anyhow::ensure!(
remaining >= n,
"{remaining} remaining layers cannot cover {n} stages at ≥1 layer each"
);
let weights: Vec<f64> = if vram_mb.iter().all(|v| *v > 0) {
vram_mb.iter().map(|v| *v as f64).collect()
} else {
vec![1.0; n] };
let wsum: f64 = weights.iter().sum();
let mut counts = Vec::with_capacity(n);
let mut rems = Vec::with_capacity(n);
for (i, w) in weights.iter().enumerate() {
let share = remaining as f64 * w / wsum;
counts.push(share as usize);
rems.push((i, share - share.floor()));
}
rems.sort_by(|a, b| b.1.partial_cmp(&a.1).expect("finite").then(a.0.cmp(&b.0)));
let mut leftover = remaining - counts.iter().sum::<usize>();
let mut ri = 0usize;
while leftover > 0 {
counts[rems[ri % n].0] += 1;
ri += 1;
leftover -= 1;
}
while let Some(z) = counts.iter().position(|c| *c == 0) {
let fat = (0..n).max_by_key(|i| counts[*i]).expect("non-empty counts");
counts[fat] -= 1;
counts[z] += 1;
}
let mut cursor = split;
Ok(counts
.into_iter()
.map(|c| {
let r = (cursor, cursor + c);
cursor += c;
r
})
.collect())
}
pub(crate) fn negotiate_fleet(
clients: &mut [ShardClient],
split: usize,
dir: &Path,
) -> Result<Vec<StageInfo>> {
let infos: Vec<StageInfo> = clients
.iter_mut()
.map(ShardClient::info)
.collect::<Result<_>>()?;
let local_fp = checkpoint_fingerprint(dir).ok();
let (fleet_fp, fleet_total) = match infos.iter().find(|i| i.ckpt_fp != 0) {
Some(i) => (i.ckpt_fp, i.total_layers),
None => (
local_fp.context(
"no stage holds the checkpoint and the coordinator's --model has no \
model.safetensors to ship from",
)?,
crate::weights::Lfm2Config::from_json(&std::fs::read(dir.join("config.json"))?)?
.n_layers,
),
};
if let Some(lfp) = local_fp {
anyhow::ensure!(
lfp == fleet_fp,
"coordinator --model fingerprint {lfp:016x} != the fleet's {fleet_fp:016x} — \
different checkpoints"
);
}
if infos.iter().all(|i| i.pinned) {
return Ok(infos);
}
let vram: Vec<u64> = infos.iter().map(|i| i.vram_mb).collect();
let ranges = plan_split(split, fleet_total, &vram)?;
eprintln!(
"auto-split: layers {split}..{fleet_total} over {} stages → {ranges:?}",
clients.len()
);
for (c, (s0, e0)) in clients.iter_mut().zip(&ranges) {
c.load_send(*s0, *e0, fleet_fp)?;
}
let mut needs_ship = Vec::new();
for (i, c) in clients.iter_mut().enumerate() {
match c
.load_wait()
.with_context(|| format!("stage {} loading layers {:?}", i + 1, ranges[i]))?
{
LoadAck::Ok => {}
LoadAck::NeedWeights => needs_ship.push(i),
LoadAck::WrongCheckpoint => anyhow::bail!(
"stage {} holds a DIFFERENT checkpoint than the fleet's {fleet_fp:016x}",
i + 1
),
LoadAck::Failed => anyhow::bail!(
"stage {} failed to load layers {:?} (see its log)",
i + 1,
ranges[i]
),
}
}
if !needs_ship.is_empty() {
let src = dir.join("model.safetensors");
anyhow::ensure!(
src.exists(),
"{} stage(s) need weights shipped, but the coordinator's --model has no single-file \
model.safetensors (HF-sharded sources are not shippable yet — copy those manually)",
needs_ship.len()
);
let cfg_bytes = std::fs::read(dir.join("config.json"))?;
let fpx = format!("{fleet_fp:016x}");
for i in needs_ship {
let range = ranges[i];
let plan = plan_mini_ckpt(&src, range.0, range.1, fleet_total)?;
eprintln!(
"shipping layers {range:?} to stage {} ({:.2} GB)…",
i + 1,
plan.total_bytes as f64 / 1e9
);
let c = &mut clients[i];
c.ship_bytes(&fpx, range, "config.json", &cfg_bytes)?;
c.ship_mini_ckpt(&fpx, range, &src, &plan)?;
c.load_send(range.0, range.1, fleet_fp)?;
anyhow::ensure!(
c.load_wait()? == LoadAck::Ok,
"stage {} still cannot load {range:?} after shipping (see its log)",
i + 1
);
}
}
clients.iter_mut().map(ShardClient::info).collect()
}
pub struct Fleet {
pub clients: Vec<ShardClient>,
pub data_addrs: Vec<String>,
}
fn admit_join(s: &mut Conn, peer: std::net::SocketAddr, tok: Option<&str>) -> Result<String> {
let mut hdr = read_u32s(s, 3)?;
if hdr[0] == OP_HELLO {
let raw = read_bytes(s, hdr[2] as usize)?;
let ok = tok.is_none_or(|t| t.as_bytes() == raw.as_slice());
write_msg(s, &[TAG_TOKEN, 1], &[f32::from_bits(u32::from(!ok))])?;
anyhow::ensure!(ok, "bad fleet token");
hdr = read_u32s(s, 3)?;
} else {
anyhow::ensure!(tok.is_none(), "tokenless join on a tokened fleet");
}
anyhow::ensure!(hdr[0] == OP_JOIN, "expected OP_JOIN, got op {}", hdr[0]);
let adv = String::from_utf8(read_bytes(s, hdr[2] as usize)?).context("advertise utf8")?;
Ok(if adv.is_empty() {
std::net::SocketAddr::new(peer.ip(), hdr[1] as u16).to_string()
} else {
adv
})
}
pub fn accept_fleet(listen: &str, n: usize) -> Result<Fleet> {
anyhow::ensure!(n >= 1, "need at least one worker");
let l = TcpListener::bind(listen)?;
let tok = std::env::var("OSFKB_SHARD_TOKEN")
.ok()
.filter(|t| !t.is_empty());
eprintln!("fleet registry on {listen}: waiting for {n} worker(s)…");
let mut clients = Vec::with_capacity(n);
let mut data_addrs = Vec::with_capacity(n);
while clients.len() < n {
let (raw, peer) = l.accept()?;
let mut s = match accept_conn(raw) {
Ok(c) => c,
Err(e) => {
eprintln!("fleet registry: handshake from {peer} failed: {e:#}");
continue;
}
};
match admit_join(&mut s, peer, tok.as_deref()) {
Ok(data) => {
eprintln!(
"fleet registry: worker {} joined from {peer} (data hops via {data})",
clients.len() + 1
);
data_addrs.push(data);
clients.push(ShardClient::from_stream(s));
}
Err(e) => eprintln!("fleet registry: rejected {peer}: {e:#}"),
}
}
Ok(Fleet {
clients,
data_addrs,
})
}
pub struct FleetRegistry {
pending: std::sync::Arc<std::sync::Mutex<std::collections::VecDeque<(ShardClient, String)>>>,
}
impl FleetRegistry {
pub fn bind(listen: &str) -> Result<Self> {
let l = TcpListener::bind(listen)?;
let tok = std::env::var("OSFKB_SHARD_TOKEN")
.ok()
.filter(|t| !t.is_empty());
eprintln!("fleet registry on {listen}: open — workers may join at any time");
let pending = std::sync::Arc::new(std::sync::Mutex::new(std::collections::VecDeque::new()));
let p = pending.clone();
std::thread::spawn(move || {
loop {
let (raw, peer) = match l.accept() {
Ok(x) => x,
Err(_) => continue,
};
let mut s = match accept_conn(raw) {
Ok(c) => c,
Err(e) => {
eprintln!("fleet registry: handshake from {peer} failed: {e:#}");
continue;
}
};
match admit_join(&mut s, peer, tok.as_deref()) {
Ok(data) => {
eprintln!(
"fleet registry: worker joined from {peer} (data hops via {data})"
);
if let Ok(mut q) = p.lock() {
q.push_back((ShardClient::from_stream(s), data));
}
}
Err(e) => eprintln!("fleet registry: rejected {peer}: {e:#}"),
}
}
});
Ok(Self { pending })
}
pub fn pending_len(&self) -> usize {
self.pending.lock().map(|q| q.len()).unwrap_or(0)
}
pub fn form(&self, min_n: usize) -> Result<Fleet> {
anyhow::ensure!(min_n >= 1, "need at least one worker");
let grace = std::time::Duration::from_millis(
std::env::var("OSFKB_FLEET_GRACE_MS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(2500),
);
eprintln!("fleet registry: waiting for ≥{min_n} worker(s)…");
while self.pending_len() < min_n {
std::thread::sleep(std::time::Duration::from_millis(50));
}
if !grace.is_zero() {
let mut last = self.pending_len();
let mut quiet = std::time::Instant::now();
loop {
std::thread::sleep(std::time::Duration::from_millis(100));
let now = self.pending_len();
if now > last {
last = now;
quiet = std::time::Instant::now();
} else if quiet.elapsed() >= grace {
break;
}
}
}
let drained: Vec<(ShardClient, String)> = self
.pending
.lock()
.map_err(|_| anyhow!("registry pool poisoned"))?
.drain(..)
.collect();
let (clients, data_addrs): (Vec<_>, Vec<_>) = drained.into_iter().unzip();
eprintln!(
"fleet registry: forming a fleet with {} worker(s)",
clients.len()
);
Ok(Fleet {
clients,
data_addrs,
})
}
}
pub(crate) fn chain_workers(
clients: &mut [ShardClient],
workers: &[&str],
ret: Option<&str>,
webrtc_hops: &[usize],
auto_webrtc: bool,
) -> Result<Conn> {
anyhow::ensure!(clients.len() == workers.len() && !clients.is_empty());
let (listener, advertise) = match ret {
Some(a) => {
let (_, port) = a
.rsplit_once(':')
.ok_or_else(|| anyhow!("return address must be host:port"))?;
let port: u16 = port.parse().context("return port")?;
(TcpListener::bind(("0.0.0.0", port))?, a.to_string())
}
None => {
let l = TcpListener::bind("0.0.0.0:0")?;
let port = l.local_addr()?.port();
(l, format!(":{port}"))
}
};
let n = clients.len();
for i in 0..n {
if i + 1 == n {
clients[i].chain(&advertise, true)?;
} else if webrtc_hops.contains(&i) {
webrtc_hop(clients, i)?;
} else {
match clients[i].chain(workers[i + 1], false) {
Ok(()) => {}
Err(direct_err) if auto_webrtc => {
eprintln!(
"chain: direct hop {i}→{} unreachable ({direct_err:#}); falling back to WebRTC",
i + 1
);
webrtc_hop(clients, i).map_err(|werr| {
direct_err.context(format!("webrtc fallback also failed: {werr:#}"))
})?;
}
Err(direct_err) => return Err(direct_err),
}
}
}
let (raw, _) = listener.accept()?;
let mut s = if advertise.starts_with("ws://") {
accept_conn(raw)? } else {
raw.set_nodelay(true).ok();
Conn::Tcp(raw)
};
if let Ok(tok) = std::env::var("OSFKB_SHARD_TOKEN")
&& !tok.is_empty()
{
let hdr = read_u32s(&mut s, 3)?;
anyhow::ensure!(
hdr[0] == OP_HELLO,
"sink: expected the fleet hello, got op {}",
hdr[0]
);
let raw = read_bytes(&mut s, hdr[2] as usize)?;
anyhow::ensure!(
raw == tok.as_bytes(),
"sink: fleet token rejected — a foreign connection reached the return port"
);
}
Ok(s)
}
pub(crate) fn read_done(s: &mut Conn) -> Result<(usize, Vec<u32>)> {
let hdr = read_u32s(s, 3)?;
anyhow::ensure!(hdr[0] == TAG_DONE, "sink expects TAG_DONE, got {}", hdr[0]);
let toks = read_f32s(s, hdr[2] as usize)?
.iter()
.map(|x| x.to_bits())
.collect();
Ok((hdr[1] as usize, toks))
}
pub fn decode_greedy_mtp_duo(
model_dir: &Path,
adapters: (usize, usize),
split: usize,
prompt: &[u32],
ngen: usize,
k_draft: usize,
) -> Result<(Vec<u32>, f64, f64)> {
anyhow::ensure!(!prompt.is_empty() && k_draft >= 1);
let kb = k_draft + 1;
let ctx0 = crate::GpuCtx::new_at(adapters.0)?;
let ctx1 = crate::GpuCtx::new_at(adapters.1)?;
let w0 = crate::Weights::load_shard(&ctx0, model_dir, 0, split)?;
let cfg_nl = {
let cfg =
crate::weights::Lfm2Config::from_json(&std::fs::read(model_dir.join("config.json"))?)?;
cfg.n_layers
};
let w1 = crate::Weights::load_shard(&ctx1, model_dir, split, cfg_nl)?;
let gpu0 = crate::Lfm2Gpu::new(&ctx0, w0);
let gpu1 = crate::Lfm2Gpu::new(&ctx1, w1);
let mtp = crate::forward::MtpEngine::new(&ctx1, &gpu1.w)
.ok_or_else(|| anyhow!("checkpoint has no mtp.* head"))?;
let slot = 1u32;
let total = prompt.len() + ngen + k_draft + 1;
let blocks: Vec<u32> = (0..(total as u32).div_ceil(16)).collect();
let bp0 = gpu0.make_batch_plan_spec(&ctx0, kb);
let bp1 = gpu1.make_batch_plan_spec(&ctx1, kb);
for (ctx, gpu) in [(&ctx0, &gpu0), (&ctx1, &gpu1)] {
gpu.reset(ctx);
gpu.zero_dn_slot(ctx, slot as usize);
gpu.write_btab_row(ctx, slot, &blocks);
}
mtp.gpu.write_btab_row(&ctx1, slot, &blocks);
let brk = std::cell::RefCell::new((0.0f64, 0.0f64, 0.0f64, 0.0f64, 0.0f64, 0.0f64, 0usize));
let dbrk = std::env::var("OSFKB_DUO_BREAK").is_ok();
let run_mb = |cols: &[BatchCol]| -> Result<Vec<u32>> {
let s0 = std::time::Instant::now();
let h = match gpu0.batch_stage_step(&ctx0, &bp0, cols, None)? {
crate::forward::StageBatchOut::Hidden(h) => h,
_ => return Err(anyhow!("stage 0 must produce hidden")),
};
let w0 = s0.elapsed().as_secs_f64();
let b0 = gpu0.take_step_breakdown();
let s1 = std::time::Instant::now();
let out = match gpu1.batch_stage_step(&ctx1, &bp1, cols, Some(&h))? {
crate::forward::StageBatchOut::Tokens(t) => t,
_ => return Err(anyhow!("stage 1 must produce tokens")),
};
let w1 = s1.elapsed().as_secs_f64();
let b1 = gpu1.take_step_breakdown();
if dbrk {
let mut a = brk.borrow_mut();
a.0 += w0;
a.1 += b0.2; a.2 += b0.0 + b0.1; a.3 += w1;
a.4 += b1.2;
a.5 += b1.0 + b1.1;
a.6 += 1;
}
Ok(out)
};
let mut col_seed = 0usize;
let mut t1 = 0u32;
for chunk in prompt.chunks(kb) {
let base = (chunk.as_ptr() as usize - prompt.as_ptr() as usize) / 4;
let cols: Vec<BatchCol> = chunk
.iter()
.enumerate()
.map(|(i, t)| BatchCol::text(*t, (base + i) as u32, slot, base + i + 1 == prompt.len()))
.collect();
let out = run_mb(&cols)?;
col_seed = cols.len() - 1;
t1 = out[col_seed];
for (i, _) in chunk.iter().enumerate() {
let q = base + i;
if q + 1 < prompt.len() {
let _ = mtp.draft_chain_seeded(&ctx1, &bp1, i, prompt[q + 1], q, 1)?;
}
}
}
let mut emitted: Vec<u32> = vec![t1];
let mut tok = t1;
let mut pos0 = prompt.len() - 1;
let (mut drafted, mut accepted) = (0usize, 0usize);
let t_start = std::time::Instant::now();
let (mut t_draft, mut t_verify, mut rounds) = (0f64, 0f64, 0usize);
while emitted.len() < ngen {
let td = std::time::Instant::now();
let drafts = mtp.draft_chain_seeded(&ctx1, &bp1, col_seed, tok, pos0, k_draft)?;
t_draft += td.elapsed().as_secs_f64();
let cols: Vec<BatchCol> = std::iter::once(tok)
.chain(drafts.iter().copied())
.enumerate()
.map(|(i, t)| BatchCol::text(t, (pos0 + 1 + i) as u32, slot, true))
.collect();
let tv = std::time::Instant::now();
let out = run_mb(&cols)?;
t_verify += tv.elapsed().as_secs_f64();
let mut j = 0usize;
while j < k_draft && out[j] == drafts[j] {
j += 1;
}
drafted += k_draft;
accepted += j;
for t in out.iter().take(j + 1) {
emitted.push(*t);
if emitted.len() == ngen {
break;
}
}
if j < k_draft {
gpu0.dn_restore(&ctx0, &bp0, slot, j);
gpu1.dn_restore(&ctx1, &bp1, slot, j);
}
rounds += 1;
pos0 = pos0 + 1 + j;
tok = out[j];
col_seed = j;
}
let secs = t_start.elapsed().as_secs_f64();
eprintln!(
"mtp duo round breakdown: {} rounds, draft {:.1}ms verify {:.1}ms per round",
rounds,
t_draft * 1000.0 / rounds.max(1) as f64,
t_verify * 1000.0 / rounds.max(1) as f64
);
if dbrk {
let a = *brk.borrow();
let n = a.6.max(1) as f64;
eprintln!(
"duo verify stage split (per run_mb, avg over {} calls):\n dev0: wall {:.2}ms [gpu+poll {:.2} cpu-prep+enc {:.2}]\n dev1: wall {:.2}ms [gpu+poll {:.2} cpu-prep+enc {:.2}]",
a.6,
a.0 * 1e3 / n,
a.1 * 1e3 / n,
a.2 * 1e3 / n,
a.3 * 1e3 / n,
a.4 * 1e3 / n,
a.5 * 1e3 / n,
);
}
Ok((
emitted,
ngen as f64 / secs,
accepted as f64 / drafted.max(1) as f64,
))
}
pub fn decode_greedy_mtp(
workers: &[&str],
prompt: &[u32],
ngen: usize,
k_draft: usize,
) -> Result<(Vec<u32>, f64, f64)> {
anyhow::ensure!(!workers.is_empty() && !prompt.is_empty() && k_draft >= 1);
let kb = k_draft + 1; let mut clients = workers
.iter()
.map(|a| ShardClient::connect(a))
.collect::<Result<Vec<_>>>()?;
let slot = 1u32;
let total = prompt.len() + ngen + k_draft + 1;
let blocks: Vec<u32> = (0..(total as u32).div_ceil(16)).collect();
for c in &mut clients {
c.chain_clear()?; c.reset()?;
c.binit(kb, 0)?;
c.zslot(slot)?;
c.btab(slot, &blocks)?;
}
let trace = std::env::var("OSFKB_MTP_TRACE").is_ok();
let run_mb = |clients: &mut [ShardClient], cols: &[(u32, u32, u32, u32)]| -> Result<Vec<u32>> {
let mut hidden: Vec<f32> = Vec::new();
let mut spans = Vec::with_capacity(clients.len());
for (i, c) in clients.iter_mut().enumerate() {
let t0 = std::time::Instant::now();
let out = c.bstep(cols, &hidden, false)?;
spans.push(t0.elapsed().as_secs_f64() * 1e3);
match out {
StageBatchOut::Hidden(h) => hidden = h,
StageBatchOut::Tokens(t) => {
anyhow::ensure!(i + 1 == workers.len(), "tokens before the last stage");
if trace {
eprintln!(
"verify stage spans: {}",
spans
.iter()
.map(|v| format!("{v:.1}ms"))
.collect::<Vec<_>>()
.join(" ")
);
}
return Ok(t);
}
}
}
Err(anyhow!("last stage returned no tokens"))
};
let mut col_seed = 0usize;
let mut t1 = 0u32;
for chunk in prompt.chunks(kb) {
let base = (chunk.as_ptr() as usize - prompt.as_ptr() as usize) / 4;
let cols: Vec<(u32, u32, u32, u32)> = chunk
.iter()
.enumerate()
.map(|(i, t)| {
let pos = (base + i) as u32;
let last = base + i + 1 == prompt.len();
(pos, slot, u32::from(last), *t)
})
.collect();
let out = run_mb(&mut clients, &cols)?;
col_seed = cols.len() - 1;
t1 = out[col_seed];
if std::env::var("OSFKB_MTP_DEBUG_ROUNDS").is_ok() {
eprintln!("prefill chunk base {base} ncols {} out {out:?}", cols.len());
}
let last = clients.len() - 1;
for (i, _) in chunk.iter().enumerate() {
let q = base + i;
if q + 1 < prompt.len() {
let _ = clients[last].mtp_draft(i, prompt[q + 1], q, 1)?;
}
}
}
let mut emitted: Vec<u32> = vec![t1];
let mut tok = t1;
let mut pos0 = prompt.len() - 1; let (mut drafted, mut accepted) = (0usize, 0usize);
let shift: i64 = std::env::var("OSFKB_MTP_SHIFT")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let mut depth_hit = vec![0usize; k_draft];
let mut depth_seen = vec![0usize; k_draft];
let beta_probe = std::env::var("OSFKB_MTP_BETA_PROBE").is_ok();
let mut cover = vec![[0usize; 3]; k_draft];
let mut cover_seen = vec![0usize; k_draft];
let t_start = std::time::Instant::now();
let last = clients.len() - 1;
let (mut t_draft, mut t_verify, mut t_restore, mut rounds) = (0f64, 0f64, 0f64, 0usize);
while emitted.len() < ngen {
let td = std::time::Instant::now();
let raw =
clients[last].mtp_draft(col_seed, tok, (pos0 as i64 + shift) as usize, k_draft)?;
let (drafts, top3) = raw.split_at(k_draft);
let drafts = drafts.to_vec();
t_draft += td.elapsed().as_secs_f64();
let cols: Vec<(u32, u32, u32, u32)> = std::iter::once(tok)
.chain(drafts.iter().copied())
.enumerate()
.map(|(i, t)| ((pos0 + 1 + i) as u32, slot, 1u32, t))
.collect();
let tv = std::time::Instant::now();
let out = run_mb(&mut clients, &cols)?;
t_verify += tv.elapsed().as_secs_f64();
let mut j = 0usize;
while j < k_draft && out[j] == drafts[j] {
j += 1;
}
if rounds
< std::env::var("OSFKB_MTP_DEBUG_ROUNDS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0)
{
eprintln!(
"round {rounds}: pos0 {pos0} tok {tok} col_seed {col_seed} drafts {drafts:?} out {:?} j {j}",
&out[..=k_draft.min(out.len() - 1)]
);
}
for i in 0..k_draft {
if i <= j {
depth_seen[i] += 1;
}
if i < j {
depth_hit[i] += 1;
}
}
if beta_probe {
for i in 0..=j.min(k_draft - 1) {
cover_seen[i] += 1;
let target = out[i];
for r in 0..3 {
if top3.get(i * 3 + r) == Some(&target) {
for c in cover[i].iter_mut().skip(r) {
*c += 1;
}
break;
}
}
}
}
drafted += k_draft;
accepted += j;
for t in out.iter().take(j + 1) {
emitted.push(*t);
if emitted.len() == ngen {
break;
}
}
if j < k_draft {
let tr = std::time::Instant::now();
for c in &mut clients {
c.dn_restore(slot, j)?;
}
t_restore += tr.elapsed().as_secs_f64();
}
rounds += 1;
pos0 = pos0 + 1 + j;
tok = out[j];
col_seed = j;
}
let secs = t_start.elapsed().as_secs_f64();
let per_depth: Vec<String> = (0..k_draft)
.map(|i| format!("d{}:{}/{}", i + 1, depth_hit[i], depth_seen[i]))
.collect();
eprintln!("mtp per-depth acceptance: {}", per_depth.join(" "));
if beta_probe {
for d in 0..k_draft {
let n = cover_seen[d].max(1);
eprintln!(
"beta depth {}: top1 {:.2} top2 {:.2} top3 {:.2} (n={})",
d + 1,
cover[d][0] as f64 / n as f64,
cover[d][1] as f64 / n as f64,
cover[d][2] as f64 / n as f64,
cover_seen[d]
);
}
}
eprintln!(
"mtp round breakdown: {} rounds, draft {:.1}ms verify {:.1}ms restore {:.1}ms per round",
rounds,
t_draft * 1000.0 / rounds.max(1) as f64,
t_verify * 1000.0 / rounds.max(1) as f64,
t_restore * 1000.0 / rounds.max(1) as f64
);
for c in &mut clients {
c.shutdown();
}
Ok((
emitted,
ngen as f64 / secs,
accepted as f64 / drafted.max(1) as f64,
))
}
pub fn decode_greedy_mtp_multi(
workers: &[&str],
prompt: &[u32],
ngen: usize,
k_draft: usize,
n_streams: usize,
) -> Result<(Vec<Vec<u32>>, f64, f64)> {
anyhow::ensure!(!workers.is_empty() && !prompt.is_empty() && k_draft >= 1 && n_streams >= 1);
let span = 1 + k_draft;
let kb = span * n_streams;
let mut clients = workers
.iter()
.map(|a| ShardClient::connect(a))
.collect::<Result<Vec<_>>>()?;
let blocks_per = (prompt.len() + ngen + k_draft + 2).div_ceil(16);
for c in &mut clients {
c.chain_clear()?; c.reset()?;
c.binit(kb, 0)?;
for st in 0..n_streams {
let slot = (st + 1) as u32;
c.zslot(slot)?;
let blocks: Vec<u32> = (0..blocks_per as u32)
.map(|b| st as u32 * blocks_per as u32 + b)
.collect();
c.btab(slot, &blocks)?;
}
}
let run_mb = |clients: &mut [ShardClient], cols: &[(u32, u32, u32, u32)]| -> Result<Vec<u32>> {
let mut hidden: Vec<f32> = Vec::new();
for (i, c) in clients.iter_mut().enumerate() {
match c.bstep(cols, &hidden, false)? {
StageBatchOut::Hidden(h) => hidden = h,
StageBatchOut::Tokens(t) => {
anyhow::ensure!(i + 1 == clients.len(), "tokens before the last stage");
return Ok(t);
}
}
}
Err(anyhow!("last stage returned no tokens"))
};
let last = clients.len() - 1;
let mut t1 = vec![0u32; n_streams];
#[allow(clippy::needless_range_loop)] for st in 0..n_streams {
let slot = (st + 1) as u32;
for chunk in prompt.chunks(kb) {
let base = (chunk.as_ptr() as usize - prompt.as_ptr() as usize) / 4;
let cols: Vec<(u32, u32, u32, u32)> = chunk
.iter()
.enumerate()
.map(|(i, t)| {
let pos = (base + i) as u32;
let is_last = base + i + 1 == prompt.len();
(pos, slot, u32::from(is_last), *t)
})
.collect();
let out = run_mb(&mut clients, &cols)?;
for (i, _) in chunk.iter().enumerate() {
let q = base + i;
if q + 1 < prompt.len() {
let _ = clients[last].mtp_draft(i, prompt[q + 1], q, 1)?;
} else {
t1[st] = out[i];
clients[last].mtp_seed(slot, i)?;
}
}
}
}
let mut emitted: Vec<Vec<u32>> = t1.iter().map(|t| vec![*t]).collect();
let mut tok: Vec<u32> = t1.clone();
let mut pos0: Vec<usize> = vec![prompt.len() - 1; n_streams];
let (mut drafted, mut accepted) = (0usize, 0usize);
let t_start = std::time::Instant::now();
while emitted.iter().any(|e| e.len() < ngen) {
let live: Vec<usize> = (0..n_streams)
.filter(|s| emitted[*s].len() < ngen)
.collect();
let reqs: Vec<(u32, u32, usize)> = live
.iter()
.map(|s| ((s + 1) as u32, tok[*s], pos0[*s]))
.collect();
let drafts = clients[last].mtp_draft_batch(&reqs, k_draft)?;
let mut cols: Vec<(u32, u32, u32, u32)> = Vec::with_capacity(span * live.len());
for (li, s) in live.iter().enumerate() {
let slot = (s + 1) as u32;
cols.push((pos0[*s] as u32 + 1, slot, 1, tok[*s]));
for (d, t) in drafts[li].iter().enumerate() {
cols.push(((pos0[*s] + 2 + d) as u32, slot, 1, *t));
}
}
let out = run_mb(&mut clients, &cols)?;
for (li, s) in live.iter().enumerate() {
let base = li * span;
let slot = (s + 1) as u32;
let mut j = 0usize;
while j < k_draft && out[base + j] == drafts[li][j] {
j += 1;
}
drafted += k_draft;
accepted += j;
for t in out[base..=base + j].iter() {
if emitted[*s].len() < ngen {
emitted[*s].push(*t);
}
}
if j < k_draft {
for c in &mut clients {
c.dn_restore(slot, base + j)?;
}
}
clients[last].mtp_seed(slot, base + j)?;
pos0[*s] += 1 + j;
tok[*s] = out[base + j];
}
}
let secs = t_start.elapsed().as_secs_f64();
let total: usize = emitted.iter().map(Vec::len).sum();
for c in &mut clients {
c.shutdown();
}
Ok((
emitted,
total as f64 / secs,
accepted as f64 / drafted.max(1) as f64,
))
}
pub struct ShardedPipeline {
ctx: GpuCtx,
first: Lfm2Gpu,
remotes: Vec<ShardClient>,
sink: Option<Conn>,
}
impl ShardedPipeline {
pub fn connect(model_dir: &Path, split: usize, workers: &[&str]) -> Result<Self> {
anyhow::ensure!(!workers.is_empty(), "need at least one remote stage");
let ctx = GpuCtx::new()?;
let w = Weights::load_shard(&ctx, model_dir, 0, split)?;
let first = Lfm2Gpu::new(&ctx, w);
let mut remotes = workers
.iter()
.map(|a| ShardClient::connect(a))
.collect::<Result<Vec<_>>>()?;
for r in &mut remotes {
r.chain_clear()?; }
let infos = negotiate_fleet(&mut remotes, split, model_dir)?;
eprintln!("{}", validate_chain(split, &infos)?);
Ok(Self {
ctx,
first,
remotes,
sink: None,
})
}
pub fn connect_p2p(
model_dir: &Path,
split: usize,
workers: &[&str],
ret: Option<&str>,
) -> Result<Self> {
Self::connect_p2p_impl(model_dir, split, workers, workers, ret, &[], false)
}
#[cfg(feature = "webrtc")]
pub fn connect_p2p_webrtc(
model_dir: &Path,
split: usize,
workers: &[&str],
ret: Option<&str>,
webrtc_hops: &[usize],
) -> Result<Self> {
Self::connect_p2p_impl(model_dir, split, workers, workers, ret, webrtc_hops, false)
}
#[cfg(feature = "webrtc")]
pub fn connect_p2p_auto(
model_dir: &Path,
split: usize,
control_addrs: &[&str],
dial_addrs: &[&str],
ret: Option<&str>,
) -> Result<Self> {
Self::connect_p2p_impl(model_dir, split, control_addrs, dial_addrs, ret, &[], true)
}
#[allow(clippy::too_many_arguments)]
fn connect_p2p_impl(
model_dir: &Path,
split: usize,
control_addrs: &[&str],
dial_addrs: &[&str],
ret: Option<&str>,
webrtc_hops: &[usize],
auto_webrtc: bool,
) -> Result<Self> {
anyhow::ensure!(!control_addrs.is_empty(), "need at least one remote stage");
anyhow::ensure!(
control_addrs.len() == dial_addrs.len(),
"control and dial address lists must have the same length"
);
anyhow::ensure!(split > 0, "the M=1 coordinator owns stage 0 (split ≥ 1)");
let ctx = GpuCtx::new()?;
let w = Weights::load_shard(&ctx, model_dir, 0, split)?;
let first = Lfm2Gpu::new(&ctx, w);
let mut remotes = control_addrs
.iter()
.map(|a| ShardClient::connect(a))
.collect::<Result<Vec<_>>>()?;
let infos = negotiate_fleet(&mut remotes, split, model_dir)?;
eprintln!("{}", validate_chain(split, &infos)?);
let sink = chain_workers(&mut remotes, dial_addrs, ret, webrtc_hops, auto_webrtc)?;
Ok(Self {
ctx,
first,
remotes,
sink: Some(sink),
})
}
pub fn decode_greedy(&mut self, prompt: &[u32], ngen: usize) -> Result<(Vec<u32>, f64)> {
self.first.reset(&self.ctx);
for r in &mut self.remotes {
r.reset()?;
}
let total = prompt.len() + ngen;
let mut out = Vec::with_capacity(ngen);
let mut tok = prompt[0];
let t0 = std::time::Instant::now();
for pos in 0..total - 1 {
let mut cur = match self.first.stage_forward(&self.ctx, pos, tok, None)? {
StageOut::Hidden(h) => h,
StageOut::Token(_) => return Err(anyhow!("stage 0 must not be last")),
};
let t = if self.sink.is_some() {
self.remotes[0].step_send(pos, &cur)?;
let sink = self.sink.as_mut().expect("checked");
let (_, toks) = read_done(sink)?;
*toks.first().ok_or_else(|| anyhow!("empty sink message"))?
} else {
let mut produced = None;
for r in &mut self.remotes {
match r.step(pos, &cur)? {
StageOut::Hidden(h) => cur = h,
StageOut::Token(t) => produced = Some(t),
}
}
produced.ok_or_else(|| anyhow!("last stage returned no token"))?
};
tok = if pos + 1 < prompt.len() {
prompt[pos + 1]
} else {
out.push(t);
t
};
}
let secs = t0.elapsed().as_secs_f64();
Ok((out, total as f64 / secs))
}
pub fn shutdown_workers(&mut self) {
for r in &mut self.remotes {
r.shutdown();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn inf(start: usize, end: usize, total: usize, backend: &str) -> StageInfo {
StageInfo {
start,
end,
total_layers: total,
subgroups: true,
vram_mb: 0,
pinned: true,
ckpt_fp: 0xF00D,
backend: backend.to_string(),
}
}
#[test]
fn should_reject_checkpoint_fingerprint_mismatch() {
let mut b = inf(3, 6, 6, "Vulkan/B");
b.ckpt_fp = 0xBEEF;
let err = validate_chain(0, &[inf(0, 3, 6, "Metal/A"), b]).unwrap_err();
assert!(err.to_string().contains("DIFFERENT weights"), "{err}");
}
#[test]
fn should_accept_contiguous_heterogeneous_chain() {
let map = validate_chain(
2,
&[inf(2, 4, 6, "Metal/Apple M3"), inf(4, 6, 6, "Vulkan/V100S")],
)
.expect("valid chain");
assert!(map.contains("Metal/Apple M3") && map.contains("Vulkan/V100S"));
assert!(map.contains("layers 0..2 (local)"));
}
#[test]
fn should_accept_headless_chain() {
let map =
validate_chain(0, &[inf(0, 3, 6, "Metal/A"), inf(3, 6, 6, "Metal/B")]).expect("valid");
assert!(map.contains("headless"));
}
#[test]
fn should_reject_gap_between_stages() {
let err = validate_chain(2, &[inf(3, 6, 6, "Metal/A")]).unwrap_err();
assert!(err.to_string().contains("gap or overlap"), "{err}");
}
#[test]
fn should_reject_uncovered_tail() {
let err = validate_chain(2, &[inf(2, 5, 6, "Metal/A")]).unwrap_err();
assert!(err.to_string().contains("must own the head"), "{err}");
}
#[test]
fn should_reject_depth_mismatch_across_fleet() {
let err =
validate_chain(0, &[inf(0, 3, 6, "Metal/A"), inf(3, 8, 8, "Vulkan/B")]).unwrap_err();
assert!(err.to_string().contains("different checkpoints"), "{err}");
}
#[test]
fn should_plan_split_proportional_to_vram() {
let r = plan_split(0, 6, &[1024, 2048]).expect("plan");
assert_eq!(r, vec![(0, 2), (2, 6)]);
}
#[test]
fn should_plan_split_equally_when_any_vram_unknown() {
let r = plan_split(0, 6, &[0, 1024]).expect("plan");
assert_eq!(r, vec![(0, 3), (3, 6)]);
}
#[test]
fn should_plan_split_start_after_coordinator_share() {
let r = plan_split(2, 6, &[0, 0]).expect("plan");
assert_eq!(r, vec![(2, 4), (4, 6)]);
}
#[test]
fn should_plan_split_floor_every_stage_at_one_layer() {
let r = plan_split(0, 5, &[10 * 1024, 1]).expect("plan");
assert_eq!(r, vec![(0, 4), (4, 5)]);
}
#[test]
fn should_plan_split_cover_exactly_with_remainders() {
let r = plan_split(0, 7, &[1, 1, 1]).expect("plan");
assert_eq!(*r.last().map(|(_, e)| e).expect("stages"), 7);
assert!(r.windows(2).all(|w| w[0].1 == w[1].0), "contiguous: {r:?}");
assert!(r.iter().all(|(s, e)| e > s), "≥1 layer each: {r:?}");
}
#[test]
fn should_plan_split_reject_more_stages_than_layers() {
let err = plan_split(2, 3, &[0, 0]).unwrap_err();
assert!(err.to_string().contains("cannot cover"), "{err}");
}
fn write_toy_st(path: &Path) -> Vec<(String, Vec<u8>)> {
use safetensors::tensor::TensorView;
let mut tensors: Vec<(String, Vec<f32>)> = vec![
("model.embed_tokens.weight".into(), vec![1.0, 2.0]),
("model.norm.weight".into(), vec![3.0]),
("mtp.fc.weight".into(), vec![9.0, 8.0]),
];
for i in 0..4 {
tensors.push((
format!("model.layers.{i}.w"),
vec![i as f32, 10.0 + i as f32],
));
}
let bytes: Vec<(String, Vec<u8>)> = tensors
.iter()
.map(|(k, v)| (k.clone(), bytemuck::cast_slice(v).to_vec()))
.collect();
let views: Vec<(String, TensorView<'_>)> = bytes
.iter()
.map(|(k, b)| {
(
k.clone(),
TensorView::new(safetensors::Dtype::F32, vec![b.len() / 4], b).unwrap(),
)
})
.collect();
std::fs::write(path, safetensors::tensor::serialize(views, &None).unwrap()).unwrap();
bytes
}
fn materialize(src: &Path, plan: &MiniCkpt) -> Vec<u8> {
use std::io::{Read as _, Seek as _};
let mut out = plan.header.clone();
let mut f = std::fs::File::open(src).unwrap();
for (off, len) in &plan.slices {
let mut b = vec![0u8; *len as usize];
f.seek(std::io::SeekFrom::Start(*off)).unwrap();
f.read_exact(&mut b).unwrap();
out.extend_from_slice(&b);
}
out
}
#[test]
fn should_plan_mini_ckpt_with_exact_middle_range_subset() {
let dir = std::env::temp_dir().join(format!("lfm2-mini-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let src = dir.join("model.safetensors");
let source = write_toy_st(&src);
let plan = plan_mini_ckpt(&src, 1, 3, 4).expect("plan");
let mini = materialize(&src, &plan);
assert_eq!(mini.len() as u64, plan.total_bytes);
let st = safetensors::SafeTensors::deserialize(&mini).expect("valid safetensors");
let mut names: Vec<&str> = st.names().into_iter().map(String::as_str).collect();
names.sort_unstable();
assert_eq!(names, vec!["model.layers.1.w", "model.layers.2.w"]);
for n in names {
let want = &source.iter().find(|(k, _)| k == n).unwrap().1;
assert_eq!(st.tensor(n).unwrap().data(), want.as_slice(), "{n} data");
}
}
#[test]
fn should_plan_mini_ckpt_include_nonlayer_tensors_at_stack_ends() {
let dir = std::env::temp_dir().join(format!("lfm2-mini2-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let src = dir.join("model.safetensors");
write_toy_st(&src);
let first = plan_mini_ckpt(&src, 0, 2, 4).expect("plan");
let last = plan_mini_ckpt(&src, 2, 4, 4).expect("plan");
for (plan, layer) in [(&first, "model.layers.0.w"), (&last, "model.layers.3.w")] {
let mini = materialize(&src, plan);
let st = safetensors::SafeTensors::deserialize(&mini).expect("valid");
let names: Vec<String> = st.names().into_iter().cloned().collect();
assert!(names.iter().any(|n| n == "model.embed_tokens.weight"));
assert!(names.iter().any(|n| n == "model.norm.weight"));
assert!(names.iter().any(|n| n == "mtp.fc.weight"));
assert!(names.iter().any(|n| n == layer));
}
}
#[test]
fn should_classify_mtp_layers_as_a_non_layer_end_tensor() {
assert_eq!(
tensor_layer("model.layers.7.self_attn.q_proj.weight"),
Some(7)
);
assert_eq!(
tensor_layer("model.language_model.layers.12.mlp.gate.weight"),
Some(12)
);
assert_eq!(tensor_layer("mtp.layers.0.self_attn.q_proj.weight"), None);
assert_eq!(tensor_layer("mtp.fc.weight"), None);
assert_eq!(tensor_layer("model.embed_tokens.weight"), None);
}
#[test]
fn should_ship_mtp_layer_tensors_to_the_last_stage_not_layer_zero() {
use safetensors::tensor::TensorView;
let dir = std::env::temp_dir().join(format!("lfm2-mtp-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let src = dir.join("model.safetensors");
let mut t: Vec<(String, Vec<f32>)> = vec![(
"mtp.layers.0.self_attn.q_proj.weight".into(),
vec![1.0, 2.0],
)];
for i in 0..4 {
t.push((format!("model.layers.{i}.w"), vec![i as f32]));
}
let bytes: Vec<(String, Vec<u8>)> = t
.iter()
.map(|(k, v)| (k.clone(), bytemuck::cast_slice(v).to_vec()))
.collect();
let views: Vec<(String, TensorView<'_>)> = bytes
.iter()
.map(|(k, b)| {
(
k.clone(),
TensorView::new(safetensors::Dtype::F32, vec![b.len() / 4], b).unwrap(),
)
})
.collect();
std::fs::write(&src, safetensors::tensor::serialize(views, &None).unwrap()).unwrap();
let last = materialize(&src, &plan_mini_ckpt(&src, 2, 4, 4).expect("plan last"));
let middle = materialize(&src, &plan_mini_ckpt(&src, 1, 3, 4).expect("plan middle"));
let names = |m: &[u8]| -> Vec<String> {
safetensors::SafeTensors::deserialize(m)
.unwrap()
.names()
.into_iter()
.cloned()
.collect()
};
assert!(
names(&last)
.iter()
.any(|n| n == "mtp.layers.0.self_attn.q_proj.weight"),
"last stage must carry the MTP head"
);
assert!(
!names(&middle).iter().any(|n| n.starts_with("mtp.")),
"a middle stage must NOT carry the MTP head (it's an end-tensor, not layer 0)"
);
}
#[test]
fn should_reconstruct_resume_state_from_a_partial() {
let dir = std::env::temp_dir().join(format!("lfm2-resume-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let missing = dir.join("absent.part");
assert_eq!(resume_state(&missing).unwrap(), (0, SHIP_FNV_INIT));
let part = dir.join("present.part");
let bytes: Vec<u8> = (0..5000u32).map(|i| (i * 7 + 1) as u8).collect();
std::fs::write(&part, &bytes).unwrap();
let (have, fnv) = resume_state(&part).unwrap();
assert_eq!(have, bytes.len() as u64);
assert_eq!(fnv, fnv64_update(SHIP_FNV_INIT, &bytes));
}
#[test]
fn should_resume_shipping_from_a_partial_and_backpressure() {
use std::io::{Read as _, Write as _};
use std::net::{TcpListener, TcpStream};
let dir = std::env::temp_dir().join(format!("lfm2-shipresume-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let src = dir.join("model.safetensors");
write_toy_st(&src);
let plan = plan_mini_ckpt(&src, 0, 4, 4).expect("plan");
let full = materialize(&src, &plan);
let skip = (full.len() / 3) as u64; let fnv_prefix = fnv64_update(SHIP_FNV_INIT, &full[..skip as usize]);
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let full_c = full.clone();
let recv = std::thread::spawn(move || {
let (mut s, _) = listener.accept().unwrap();
let mut hdr = [0u8; 12];
let mut got: Vec<u8> = Vec::new();
let mut acks = 0u32;
loop {
s.read_exact(&mut hdr).unwrap();
let op = u32::from_le_bytes(hdr[0..4].try_into().unwrap());
let a = u32::from_le_bytes(hdr[4..8].try_into().unwrap());
let b = u32::from_le_bytes(hdr[8..12].try_into().unwrap());
if op == OP_SHIP_BEGIN {
let mut meta = vec![0u8; b as usize];
s.read_exact(&mut meta).unwrap();
let reply = [
skip as u32,
(skip >> 32) as u32,
fnv_prefix as u32,
(fnv_prefix >> 32) as u32,
];
for w in reply {
s.write_all(&w.to_le_bytes()).unwrap();
}
} else if op == OP_SHIP_CHUNK {
let mut c = vec![0u8; b as usize];
s.read_exact(&mut c).unwrap();
got.extend_from_slice(&c);
s.write_all(&TAG_TOKEN.to_le_bytes()).unwrap();
s.write_all(&0u32.to_le_bytes()).unwrap();
acks += 1;
} else if op == OP_SHIP_END {
let fnv = (a as u64) | ((b as u64) << 32);
let reassembled: Vec<u8> = full_c[..skip as usize]
.iter()
.chain(got.iter())
.copied()
.collect();
let ok =
fnv == fnv64_update(SHIP_FNV_INIT, &reassembled) && reassembled == full_c;
s.write_all(&TAG_TOKEN.to_le_bytes()).unwrap();
s.write_all(&1u32.to_le_bytes()).unwrap();
s.write_all(&f32::from_bits(u32::from(!ok)).to_le_bytes())
.unwrap();
return (got, ok, acks);
}
}
});
let mut client = ShardClient::from_stream(Conn::Tcp(TcpStream::connect(addr).unwrap()));
client
.ship_mini_ckpt("dead00", (0, 4), &src, &plan)
.expect("ship");
let (got, ok, acks) = recv.join().unwrap();
assert!(ok, "receiver rejected integrity");
assert_eq!(
got,
full[skip as usize..],
"shipped exactly the post-resume remainder"
);
assert!(acks >= 1, "chunks were acked (back-pressure handshake ran)");
}
#[cfg(feature = "ws")]
#[test]
fn should_roundtrip_large_messages_over_the_ws_byte_pipe() {
let l = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = format!("ws://{}", l.local_addr().expect("addr"));
let server = std::thread::spawn(move || {
let (s, _) = l.accept().expect("accept");
let mut c = accept_conn(s).expect("ws upgrade");
let hdr = read_u32s(&mut c, 3).expect("hdr");
let body = read_bytes(&mut c, hdr[2] as usize).expect("body");
write_msg_bytes(&mut c, &[hdr[0], 0, body.len() as u32], &body).expect("echo");
});
let mut c = dial(&addr, std::time::Duration::from_secs(5)).expect("dial ws");
let big: Vec<u8> = (0..20 * 1024 * 1024).map(|i| (i % 251) as u8).collect();
write_msg_bytes(&mut c, &[0xBEEF, 0, big.len() as u32], &big).expect("send");
let hdr = read_u32s(&mut c, 3).expect("hdr back");
assert_eq!(hdr[0], 0xBEEF);
let echoed = read_bytes(&mut c, hdr[2] as usize).expect("echo body");
assert_eq!(echoed, big, "byte pipe must be transparent");
server.join().expect("server thread");
}
}