diff --git a/scripts/compare_diloco.py b/scripts/compare_diloco.py index b18e97e..87644bc 100644 --- a/scripts/compare_diloco.py +++ b/scripts/compare_diloco.py @@ -285,6 +285,7 @@ def syncer_command( *, probe_capture: bool = False, probe_capture_every: int = 1, + merge_policy: str = "production", ) -> list[str]: # The syncer takes no fragment count: the layout arrives in HELLO. cmd = [ @@ -296,6 +297,7 @@ def syncer_command( "--total-steps", str(total_steps), "--pipeline", str(arm.pipeline), "--delta-correction", arm.delta_correction, + "--merge-policy", merge_policy, "--outer-lr", str(arm.outer_lr), "--outer-momentum", str(arm.outer_momentum), "--checkpoint-path", str(arm_dir / "state.ckpt"), @@ -525,6 +527,7 @@ def run_diloco(args, arm: Arm, work: Path) -> tuple[Path, float]: total_steps=steps * arm.m * 4, probe_capture=getattr(args, "syncer_probe_capture", False), probe_capture_every=getattr(args, "syncer_probe_capture_every", 1), + merge_policy=getattr(args, "syncer_merge_policy", "production"), ), stdout=open(arm_dir / "syncer.log", "w"), stderr=subprocess.STDOUT, ) @@ -689,6 +692,19 @@ def main() -> int: default=1, help="capture every Nth outer step when --syncer-probe-capture is set", ) + p.add_argument( + "--syncer-merge-policy", + default="production", + choices=[ + "production", + "coord-midpoint-normmatch", + "consensus-rda-sqrt", + "consensus-rda-linear", + "consensus-rda-affine50", + "consensus-rda-floor50", + ], + help="syncer aggregation policy for async arms", + ) p.add_argument("--work-dir", type=Path, default=REPO_ROOT / "compare-work") p.add_argument("--report-dir", type=Path, default=REPO_ROOT / "compare-report") p.add_argument("--dry-run", action="store_true", help="print the plan; run nothing") diff --git a/scripts/manual_stale_push_check.py b/scripts/manual_stale_push_check.py new file mode 100644 index 0000000..5082bd3 --- /dev/null +++ b/scripts/manual_stale_push_check.py @@ -0,0 +1,193 @@ +#!/usr/bin/env python3 +"""Force a stale PUSH_FRAGMENT and verify the syncer drops it. + +This is a small manual/integration check for the stale-commit guard. It uses +only the Python standard library and talks to the syncer wire protocol directly: + +1. Start a one-fragment, one-learner syncer for two outer steps. +2. Send a fresh push for step 1 with base_version=0. +3. Send the step-2 push with stale base_version=0 after the next pull. +5. Assert the syncer logs "stale push dropped" and never "stale push admitted". +""" + +from __future__ import annotations + +import argparse +import socket +import struct +import subprocess +import tempfile +import time +from pathlib import Path + + +MAGIC = 0xD170_C0DE +MSG_HELLO = 1 +MSG_INIT_PARAMS = 2 +MSG_PULL_REQ = 3 +MSG_PUSH_FRAGMENT = 4 +MSG_BCAST_FRAGMENT = 5 +MSG_SHUTDOWN = 7 +DTYPE_F32 = 1 +MERGE_AVG = 0 + + +def frame(msg_type: int, payload: bytes) -> bytes: + return struct.pack(" bytes: + out = bytearray() + while len(out) < n: + chunk = sock.recv(n - len(out)) + if not chunk: + raise EOFError("socket closed") + out.extend(chunk) + return bytes(out) + + +def read_frame(sock: socket.socket) -> tuple[int, bytes]: + header = read_exact(sock, 13) + magic, msg_type, length = struct.unpack(" tuple[int, int]: + deadline = time.monotonic() + timeout_s + while time.monotonic() < deadline: + msg_type, payload = read_frame(sock) + if msg_type == MSG_PULL_REQ: + fid, global_step = struct.unpack(" None: + payload = bytearray() + payload += struct.pack(" None: + payload = struct.pack(" None: + payload = bytearray() + payload += struct.pack(" int: + parser = argparse.ArgumentParser() + parser.add_argument( + "--syncer-bin", + type=Path, + default=Path("syncer/target/debug/yeto-syncer"), + help="path to a built yeto-syncer binary", + ) + parser.add_argument("--host", default="localhost", help="host to connect to for the temporary syncer") + parser.add_argument("--port", type=int, default=29591, help="localhost port for the temporary syncer") + args = parser.parse_args() + syncer_bin = args.syncer_bin + if not syncer_bin.exists(): + raise SystemExit(f"{syncer_bin} does not exist; run `cd syncer && cargo build` first") + + port = args.port + with tempfile.TemporaryDirectory(prefix="yeto-stale-push-") as td: + root = Path(td) + log_path = root / "syncer.log" + with log_path.open("wb") as log: + proc = subprocess.Popen( + [ + str(syncer_bin), + "--port", + str(port), + "--learners", + "1", + "--quorum", + "1", + "--grace-ms", + "0", + "--sync-interval-steps", + "0", + "--total-steps", + "2", + "--outer-lr", + "1.0", + "--outer-momentum", + "0.0", + ], + stdout=log, + stderr=subprocess.STDOUT, + cwd=Path(__file__).resolve().parents[1], + ) + try: + deadline = time.monotonic() + 10.0 + sock = None + while True: + candidate = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + try: + candidate.connect((args.host, port)) + sock = candidate + break + except OSError: + candidate.close() + if time.monotonic() > deadline: + raise + time.sleep(0.05) + assert sock is not None + with sock: + send_hello(sock) + send_init(sock, 0.0) + wait_for_pull(sock, 1) + send_push(sock, step=1, base_version=0, value=-1.0) + wait_for_pull(sock, 2) + send_push(sock, step=2, base_version=0, value=-2.0) + + # Drain until shutdown or socket close so the syncer can log. + try: + while True: + msg_type, _ = read_frame(sock) + if msg_type == MSG_SHUTDOWN: + break + except EOFError: + pass + + proc.wait(timeout=10) + finally: + if proc.poll() is None: + proc.terminate() + try: + proc.wait(timeout=5) + except subprocess.TimeoutExpired: + proc.kill() + + text = log_path.read_text(errors="replace") + if "stale push admitted" in text: + print(text) + raise SystemExit("FAIL: stale push was admitted") + if "stale push dropped" not in text: + print(text) + raise SystemExit("FAIL: stale push drop was not observed") + if "round had no fresh pushes after stale-drop filter" not in text: + print(text) + raise SystemExit("FAIL: no-op rebroadcast path was not observed") + + print("PASS: stale full/f32 push was dropped before merge") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/syncer/src/main.rs b/syncer/src/main.rs index 9e190f7..9c39b28 100644 --- a/syncer/src/main.rs +++ b/syncer/src/main.rs @@ -4,6 +4,7 @@ mod server; mod state; use clap::Parser; +use state::MergePolicy; /// Yeto syncer: pull-driven fragment merging (weighted RDA/Avg) /// with an SGD+Nesterov outer optimizer. See docs/PROTOCOL.md. @@ -48,6 +49,9 @@ struct Args { /// Pre-merge learner-delta correction: "heloco" or "none". #[arg(long, default_value = "heloco")] delta_correction: String, + /// Fresh-push aggregation rule. + #[arg(long, default_value = "production")] + merge_policy: String, /// Give up waiting for quorum and re-send the pull after this long. #[arg(long, default_value_t = 900)] quorum_timeout_s: u64, @@ -99,6 +103,15 @@ fn main() -> anyhow::Result<()> { "none" => false, other => anyhow::bail!("--delta-correction must be 'heloco' or 'none', got {other:?}"), }; + let merge_policy = match args.merge_policy.as_str() { + "production" => MergePolicy::Production, + "coord-midpoint-normmatch" => MergePolicy::CoordMidpointNormmatch, + "consensus-rda-sqrt" => MergePolicy::ConsensusRdaSqrt, + "consensus-rda-linear" => MergePolicy::ConsensusRdaLinear, + "consensus-rda-affine50" => MergePolicy::ConsensusRdaAffine50, + "consensus-rda-floor50" => MergePolicy::ConsensusRdaFloor50, + other => anyhow::bail!("unknown --merge-policy {other:?}"), + }; let cfg = server::Config { port: args.port, learners: args.learners, @@ -110,6 +123,7 @@ fn main() -> anyhow::Result<()> { min_round_interval_ms: args.min_round_interval_ms, sync_interval_steps: args.sync_interval_steps, delta_correction, + merge_policy, quorum_timeout_s: args.quorum_timeout_s, total_steps: args.total_steps, outer_lr: args.outer_lr, diff --git a/syncer/src/merge.rs b/syncer/src/merge.rs index 042c81b..10bdf1b 100644 --- a/syncer/src/merge.rs +++ b/syncer/src/merge.rs @@ -38,6 +38,22 @@ fn l2_norm(anchor: &[f32], learner: &[f32]) -> f64 { .sqrt() } +fn norm(values: &[f32]) -> f64 { + values + .iter() + .map(|v| (*v as f64) * (*v as f64)) + .sum::() + .sqrt() +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum ConsensusScale { + Sqrt, + Linear, + Affine50, + Floor50, +} + /// Weighted direct averaging: out[i] = Σ_m w_m (anchor[i] − learner_m[i]) / Σ w. pub fn merge_avg(anchor: &[f32], learners: &[&[f32]], weights: &[f64], out: &mut [f32]) { let wsum: f64 = weights.iter().sum(); @@ -62,12 +78,7 @@ pub fn merge_rda(anchor: &[f32], learners: &[&[f32]], weights: &[f64], out: &mut return; } let norms: Vec = learners.iter().map(|l| l2_norm(anchor, l)).collect(); - let radial: f64 = norms - .iter() - .zip(weights) - .map(|(n, w)| n * w) - .sum::() - / wsum; + let radial: f64 = norms.iter().zip(weights).map(|(n, w)| n * w).sum::() / wsum; // Weighted mean of unit directions, φ(0) := 0. out.fill(0.0); @@ -80,7 +91,11 @@ pub fn merge_rda(anchor: &[f32], learners: &[&[f32]], weights: &[f64], out: &mut *o += coef * (*a - *l); } } - let mean_dir_norm = out.iter().map(|v| (*v as f64) * (*v as f64)).sum::().sqrt(); + let mean_dir_norm = out + .iter() + .map(|v| (*v as f64) * (*v as f64)) + .sum::() + .sqrt(); if mean_dir_norm < 1e-12 { // Degenerate (all-zero or cancelling directions): fall back to Avg. merge_avg(anchor, learners, weights, out); @@ -92,6 +107,90 @@ pub fn merge_rda(anchor: &[f32], learners: &[&[f32]], weights: &[f64], out: &mut } } +/// RDA with a consensus-dependent radial scale. Consensus is the norm of the +/// weighted mean unit direction; low agreement damps that tensor's update. +pub fn merge_consensus_rda( + anchor: &[f32], + learners: &[&[f32]], + weights: &[f64], + scale_mode: ConsensusScale, + out: &mut [f32], +) -> f64 { + let wsum: f64 = weights.iter().sum(); + if wsum <= 0.0 { + out.fill(0.0); + return 0.0; + } + let norms: Vec = learners.iter().map(|l| l2_norm(anchor, l)).collect(); + let radial: f64 = norms.iter().zip(weights).map(|(n, w)| n * w).sum::() / wsum; + + out.fill(0.0); + for ((learner, &w), &n) in learners.iter().zip(weights).zip(&norms) { + if n <= 1e-12 { + continue; + } + let coef = (w / wsum / n) as f32; + for ((o, a), l) in out.iter_mut().zip(anchor).zip(*learner) { + *o += coef * (*a - *l); + } + } + let consensus = norm(out).clamp(0.0, 1.0); + if consensus < 1e-12 { + merge_avg(anchor, learners, weights, out); + return 0.0; + } + let scale = match scale_mode { + ConsensusScale::Sqrt => consensus.sqrt(), + ConsensusScale::Linear => consensus, + ConsensusScale::Affine50 => 0.5 + 0.5 * consensus, + ConsensusScale::Floor50 => consensus.max(0.5), + }; + let factor = (radial * scale / consensus) as f32; + for o in out.iter_mut() { + *o *= factor; + } + scale +} + +/// Coordinate-wise midpoint/median over learner deltas, with the resulting +/// robust direction rescaled to the production merge norm for this tensor. +pub fn merge_coord_midpoint_normmatch( + anchor: &[f32], + learners: &[&[f32]], + target_delta: &[f32], + out: &mut [f32], +) { + if learners.is_empty() { + out.fill(0.0); + return; + } + debug_assert_eq!(anchor.len(), target_delta.len()); + debug_assert_eq!(anchor.len(), out.len()); + let mut coord = Vec::with_capacity(learners.len()); + for i in 0..anchor.len() { + coord.clear(); + for learner in learners { + coord.push(anchor[i] - learner[i]); + } + coord.sort_by(|a, b| a.total_cmp(b)); + let mid = coord.len() / 2; + out[i] = if coord.len() % 2 == 0 { + 0.5 * (coord[mid - 1] + coord[mid]) + } else { + coord[mid] + }; + } + let source_norm = norm(out); + if source_norm < 1e-12 { + return; + } + let target_norm = norm(target_delta); + let scale = (target_norm / source_norm) as f32; + for v in out.iter_mut() { + *v *= scale; + } +} + /// SGD + Nesterov momentum treating `delta` as the gradient: /// buf ← μ·buf + Δ; θ ← θ − lr·(Δ + μ·buf). pub fn nesterov_step(params: &mut [f32], buf: &mut [f32], delta: &[f32], lr: f32, mu: f32) { @@ -130,18 +229,37 @@ pub struct Heloco { impl Default for Heloco { fn default() -> Self { // Table 3 of the paper. - Self { c_ok: 0.2, k_s: 0.5, k_d: 1.0, beta_max: 0.5, kappa: 3.0, eps: 1e-8 } + Self { + c_ok: 0.2, + k_s: 0.5, + k_d: 1.0, + beta_max: 0.5, + kappa: 3.0, + eps: 1e-8, + } } } pub fn heloco_correct(delta: &mut [f32], momentum: &[f32], h: &Heloco) { debug_assert_eq!(delta.len(), momentum.len()); - let du = delta.iter().map(|v| (*v as f64).powi(2)).sum::().sqrt(); - let dm = momentum.iter().map(|v| (*v as f64).powi(2)).sum::().sqrt(); + let du = delta + .iter() + .map(|v| (*v as f64).powi(2)) + .sum::() + .sqrt(); + let dm = momentum + .iter() + .map(|v| (*v as f64).powi(2)) + .sum::() + .sqrt(); if du < h.eps || dm < h.eps { return; } - let dot: f64 = delta.iter().zip(momentum).map(|(d, m)| *d as f64 * *m as f64).sum(); + let dot: f64 = delta + .iter() + .zip(momentum) + .map(|(d, m)| *d as f64 * *m as f64) + .sum(); let c = dot / (du * dm); if c >= h.c_ok { return; @@ -250,6 +368,84 @@ mod tests { assert!(out[0].abs() < 1e-6); } + #[test] + fn consensus_rda_floor50_damps_low_agreement() { + let anchor = [0.0f32, 0.0]; + let l0 = [-2.0f32, 0.0]; // delta (2, 0) + let l1 = [0.0f32, -2.0]; // delta (0, 2) + let mut rda = [0.0f32; 2]; + let mut out = [0.0f32; 2]; + + merge_rda(&anchor, &[&l0, &l1], &[1.0, 1.0], &mut rda); + let scale = merge_consensus_rda( + &anchor, + &[&l0, &l1], + &[1.0, 1.0], + ConsensusScale::Floor50, + &mut out, + ); + + assert!(scale > 0.5 && scale < 1.0); + assert!((norm(&out) - norm(&rda) * scale).abs() < 1e-5); + assert!((out[0] - out[1]).abs() < 1e-6); + } + + #[test] + fn consensus_rda_affine50_is_gentler_than_linear() { + let anchor = [0.0f32, 0.0]; + let l0 = [-2.0f32, 0.0]; + let l1 = [0.0f32, -2.0]; + let mut linear = [0.0f32; 2]; + let mut affine = [0.0f32; 2]; + + let linear_scale = merge_consensus_rda( + &anchor, + &[&l0, &l1], + &[1.0, 1.0], + ConsensusScale::Linear, + &mut linear, + ); + let affine_scale = merge_consensus_rda( + &anchor, + &[&l0, &l1], + &[1.0, 1.0], + ConsensusScale::Affine50, + &mut affine, + ); + + assert!(affine_scale > linear_scale); + assert!(norm(&affine) > norm(&linear)); + } + + #[test] + fn coord_midpoint_normmatch_uses_robust_direction_and_target_norm() { + let anchor = [0.0f32, 0.0]; + let l0 = [-10.0f32, 0.0]; // delta (10, 0) + let l1 = [0.0f32, -10.0]; // delta (0, 10) + let l2 = [-1.0f32, -1.0]; // median delta (1, 1) + let target = [3.0f32, 4.0]; // norm 5 + let mut out = [0.0f32; 2]; + + merge_coord_midpoint_normmatch(&anchor, &[&l0, &l1, &l2], &target, &mut out); + + assert!((norm(&out) - 5.0).abs() < 1e-5); + assert!((out[0] - out[1]).abs() < 1e-6); + assert!(out[0] > 0.0); + } + + #[test] + fn coord_midpoint_normmatch_midpoints_even_groups() { + let anchor = [0.0f32]; + let l0 = [-2.0f32]; // delta 2 + let l1 = [0.0f32]; // delta 0 + let target = [3.0f32]; + let mut out = [0.0f32; 1]; + + merge_coord_midpoint_normmatch(&anchor, &[&l0, &l1], &target, &mut out); + + assert!((out[0] - 3.0).abs() < 1e-6); + } + fn cosine(a: &[f32], b: &[f32]) -> f64 { let dot: f64 = a.iter().zip(b).map(|(x, y)| *x as f64 * *y as f64).sum(); dot / (norm(a) * norm(b)) @@ -294,7 +490,11 @@ mod tests { let mag = norm(&d); let before = cosine(&d, &m); heloco_correct(&mut d, &m, &h); - assert!((norm(&d) - mag).abs() < 1e-5, "magnitude changed: {mag} -> {}", norm(&d)); + assert!( + (norm(&d) - mag).abs() < 1e-5, + "magnitude changed: {mag} -> {}", + norm(&d) + ); assert!(cosine(&d, &m) > before); } @@ -308,8 +508,16 @@ mod tests { // Same directions, but huge momentum norm → low confidence → weaker // correction (closer to the original delta). let orig = [-1.0f32, 0.2]; - let moved_small: f64 = small_m.iter().zip(&orig).map(|(a, b)| (a - b).abs() as f64).sum(); - let moved_large: f64 = large_m.iter().zip(&orig).map(|(a, b)| (a - b).abs() as f64).sum(); + let moved_small: f64 = small_m + .iter() + .zip(&orig) + .map(|(a, b)| (a - b).abs() as f64) + .sum(); + let moved_large: f64 = large_m + .iter() + .zip(&orig) + .map(|(a, b)| (a - b).abs() as f64) + .sum(); assert!(moved_large < moved_small); } diff --git a/syncer/src/server.rs b/syncer/src/server.rs index db7befa..cec1a7f 100644 --- a/syncer/src/server.rs +++ b/syncer/src/server.rs @@ -19,7 +19,7 @@ use tokio::sync::mpsc; use tracing::{info, warn}; use crate::protocol::*; -use crate::state::{GlobalState, Layout}; +use crate::state::{GlobalState, Layout, MergePolicy}; const CHUNK_SIZE: usize = 4 * 1024 * 1024; const WRITE_TIMEOUT: Duration = Duration::from_secs(180); @@ -62,6 +62,8 @@ pub struct Config { pub sync_interval_steps: f64, /// HeLoCo per-tensor delta correction before merging. pub delta_correction: bool, + /// Aggregation policy for fresh learner pushes. + pub merge_policy: MergePolicy, pub quorum_timeout_s: u64, pub total_steps: u64, pub outer_lr: f32, @@ -761,25 +763,12 @@ async fn complete_round( .. } = round; let prev_version = st.versions[p]; + drop_stale_pushes(&mut pushes, prev_version, st.wire_dtype, t); if st.wire_dtype == DTYPE_Q4 { // Q4 pushes are deltas anchored at the learner's base_version; // reconstruction needs Θ at that exact version, and the syncer - // only holds the current value. A matching base is the steady - // state (learners anchor on the last broadcast); anything older - // is unreconstructable and dropped. - pushes.retain(|id, push| { - if push.base_version != prev_version { - warn!( - learner_id = id, - step = t, - base = push.base_version, - expected = prev_version, - "stale q4 delta dropped" - ); - return false; - } - true - }); + // only holds the current value. `drop_stale_pushes` guarantees + // that every remaining push matches the current fragment version. for push in pushes.values_mut() { for (v, a) in push.values.iter_mut().zip(&st.params[p]) { *v += *a; @@ -789,24 +778,22 @@ async fn complete_round( capture_round_candidates(cfg, st, p, t, prev_version, &pushes)?; let (mut learners, mut weights, mut ids) = (Vec::new(), Vec::new(), Vec::new()); for (id, push) in &pushes { - if push.base_version < prev_version { - // The learner had not yet applied this fragment's last merge; - // its delta is anchored further back. The weight formula - // compensates (larger c_steps); recorded for the event tape. - warn!( - learner_id = id, - step = t, - base = push.base_version, - expected = prev_version, - "stale push admitted" - ); - } learners.push(push.values.as_slice()); weights.push(crate::merge::learner_weight(push.c_tokens, push.c_steps)); ids.push(*id); } let sync_start = Instant::now(); - let gnorm = st.merge_and_step(p, &learners, &weights)?; + let gnorm = if learners.is_empty() { + warn!( + step = t, + fragment = p, + prev_version, + "round had no fresh pushes after stale-drop filter; rebroadcasting current fragment" + ); + 0.0 + } else { + st.merge_and_step(p, &learners, &weights)? + }; st.versions[p] = t; // Pipelined rounds can complete out of order; the global step only // moves forward. @@ -848,6 +835,37 @@ async fn complete_round( Ok(()) } +fn drop_stale_pushes( + pushes: &mut HashMap, + prev_version: u64, + wire_dtype: u8, + step: u64, +) { + pushes.retain(|id, push| { + if push.base_version == prev_version { + return true; + } + if wire_dtype == DTYPE_Q4 { + warn!( + learner_id = id, + step, + base = push.base_version, + expected = prev_version, + "stale q4 delta dropped" + ); + } else { + warn!( + learner_id = id, + step, + base = push.base_version, + expected = prev_version, + "stale push dropped" + ); + } + false + }); +} + fn capture_round_candidates( cfg: &Config, st: &GlobalState, @@ -979,6 +997,7 @@ fn new_state_for(group: &Arc, cfg: &Config) -> Result { if cfg.delta_correction { st.delta_correction = Some(crate::merge::Heloco::default()); } + st.merge_policy = cfg.merge_policy; Ok(st) } @@ -1095,6 +1114,19 @@ mod tests { const CAP: Duration = Duration::from_millis(1000); + fn push(id: u32, base_version: u64) -> Push { + Push { + learner_id: id, + fragment_id: 0, + global_step: 7, + base_version, + local_step: 10, + c_steps: 1, + c_tokens: 128, + values: vec![0.0], + } + } + #[test] fn grace_falls_back_to_cap_without_estimate() { assert_eq!(adaptive_grace(2.0, 0.8, None, 0.1, 0.1, CAP), CAP); @@ -1134,6 +1166,30 @@ mod tests { assert_eq!(launch_interval(floor, 0.0, 4, Some(1.0)), floor); } + #[test] + fn stale_full_pushes_are_dropped_before_merge() { + let mut pushes = HashMap::new(); + pushes.insert(1, push(1, 4)); + pushes.insert(2, push(2, 5)); + + drop_stale_pushes(&mut pushes, 5, DTYPE_BF16, 9); + + assert_eq!(pushes.len(), 1); + assert!(pushes.contains_key(&2)); + } + + #[test] + fn stale_q4_pushes_use_same_freshness_gate() { + let mut pushes = HashMap::new(); + pushes.insert(1, push(1, 4)); + pushes.insert(2, push(2, 6)); + + drop_stale_pushes(&mut pushes, 6, DTYPE_Q4, 9); + + assert_eq!(pushes.len(), 1); + assert!(pushes.contains_key(&2)); + } + #[test] fn step_rates_estimate_from_consecutive_pushes() { let mut rates = StepRates::default(); diff --git a/syncer/src/state.rs b/syncer/src/state.rs index 734186a..8dd52bb 100644 --- a/syncer/src/state.rs +++ b/syncer/src/state.rs @@ -9,6 +9,16 @@ use crate::protocol::Reader; pub const MERGE_AVG: u8 = 0; pub const MERGE_RDA: u8 = 1; +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum MergePolicy { + Production, + CoordMidpointNormmatch, + ConsensusRdaSqrt, + ConsensusRdaLinear, + ConsensusRdaAffine50, + ConsensusRdaFloor50, +} + #[derive(Clone, Debug, PartialEq)] pub struct FragmentInfo { pub merge_mode: u8, @@ -41,7 +51,10 @@ impl Layout { for _ in 0..num_tensors { tensor_numels.push(r.u64()?); } - fragments.push(FragmentInfo { merge_mode, tensor_numels }); + fragments.push(FragmentInfo { + merge_mode, + tensor_numels, + }); } Ok(Layout { fragments }) } @@ -75,6 +88,8 @@ pub struct GlobalState { /// HeLoCo per-tensor directional correction of learner deltas against /// the outer momentum before merging (None disables). pub delta_correction: Option, + /// Syncer-side aggregation rule applied after freshness filtering. + pub merge_policy: MergePolicy, } impl GlobalState { @@ -85,7 +100,11 @@ impl GlobalState { outer_momentum: f32, wire_dtype: u8, ) -> Self { - let params: Vec> = layout.fragments.iter().map(|f| vec![0.0; f.numel()]).collect(); + let params: Vec> = layout + .fragments + .iter() + .map(|f| vec![0.0; f.numel()]) + .collect(); let momentum = params.clone(); let initialized = vec![false; layout.fragments.len()]; let versions = vec![0; layout.fragments.len()]; @@ -102,6 +121,7 @@ impl GlobalState { outer_momentum, wire_dtype, delta_correction: None, + merge_policy: MergePolicy::Production, } } @@ -111,7 +131,11 @@ impl GlobalState { pub fn init_fragment(&mut self, fid: usize, values: Vec) -> Result<()> { if values.len() != self.params[fid].len() { - bail!("init fragment {fid}: got {} values, expected {}", values.len(), self.params[fid].len()); + bail!( + "init fragment {fid}: got {} values, expected {}", + values.len(), + self.params[fid].len() + ); } if !self.initialized[fid] { self.params[fid] = values; @@ -129,12 +153,20 @@ impl GlobalState { /// Merge learner copies of fragment `fid` and apply the outer step. /// Returns the l2 norm of the merged outer gradient (for logging). - pub fn merge_and_step(&mut self, fid: usize, learners: &[&[f32]], weights: &[f64]) -> Result { + pub fn merge_and_step( + &mut self, + fid: usize, + learners: &[&[f32]], + weights: &[f64], + ) -> Result { let frag = &self.layout.fragments[fid]; let numel = frag.numel(); for (i, l) in learners.iter().enumerate() { if l.len() != numel { - bail!("push for fragment {fid} from entry {i} has {} values, expected {numel}", l.len()); + bail!( + "push for fragment {fid} from entry {i} has {} values, expected {numel}", + l.len() + ); } } // HeLoCo: correct each learner's outer delta against the outer @@ -172,20 +204,82 @@ impl GlobalState { }; let learners = learners.as_slice(); let anchor = &self.params[fid]; - let mut delta = vec![0.0f32; numel]; + let mut production_delta = vec![0.0f32; numel]; // Merge per tensor slice within the fragment. let mut off = 0usize; for &tn in &frag.tensor_numels { let tn = tn as usize; let slice_learners: Vec<&[f32]> = learners.iter().map(|l| &l[off..off + tn]).collect(); - let out = &mut delta[off..off + tn]; + let out = &mut production_delta[off..off + tn]; match frag.merge_mode { - MERGE_AVG => merge::merge_avg(&anchor[off..off + tn], &slice_learners, weights, out), + MERGE_AVG => { + merge::merge_avg(&anchor[off..off + tn], &slice_learners, weights, out) + } _ => merge::merge_rda(&anchor[off..off + tn], &slice_learners, weights, out), } off += tn; } - let gnorm = delta.iter().map(|v| (*v as f64).powi(2)).sum::().sqrt(); + let delta = match self.merge_policy { + MergePolicy::Production => production_delta, + MergePolicy::ConsensusRdaSqrt + | MergePolicy::ConsensusRdaLinear + | MergePolicy::ConsensusRdaAffine50 + | MergePolicy::ConsensusRdaFloor50 => { + let mut consensus_delta = vec![0.0f32; numel]; + let scale_mode = match self.merge_policy { + MergePolicy::ConsensusRdaSqrt => merge::ConsensusScale::Sqrt, + MergePolicy::ConsensusRdaLinear => merge::ConsensusScale::Linear, + MergePolicy::ConsensusRdaAffine50 => merge::ConsensusScale::Affine50, + MergePolicy::ConsensusRdaFloor50 => merge::ConsensusScale::Floor50, + _ => unreachable!(), + }; + let mut off = 0usize; + for &tn in &frag.tensor_numels { + let tn = tn as usize; + let slice_learners: Vec<&[f32]> = + learners.iter().map(|l| &l[off..off + tn]).collect(); + let out = &mut consensus_delta[off..off + tn]; + match frag.merge_mode { + MERGE_AVG => { + merge::merge_avg(&anchor[off..off + tn], &slice_learners, weights, out) + } + _ => { + merge::merge_consensus_rda( + &anchor[off..off + tn], + &slice_learners, + weights, + scale_mode, + out, + ); + } + } + off += tn; + } + consensus_delta + } + MergePolicy::CoordMidpointNormmatch => { + let mut robust_delta = vec![0.0f32; numel]; + let mut off = 0usize; + for &tn in &frag.tensor_numels { + let tn = tn as usize; + let slice_learners: Vec<&[f32]> = + learners.iter().map(|l| &l[off..off + tn]).collect(); + merge::merge_coord_midpoint_normmatch( + &anchor[off..off + tn], + &slice_learners, + &production_delta[off..off + tn], + &mut robust_delta[off..off + tn], + ); + off += tn; + } + robust_delta + } + }; + let gnorm = delta + .iter() + .map(|v| (*v as f64).powi(2)) + .sum::() + .sqrt(); merge::nesterov_step( &mut self.params[fid], &mut self.momentum[fid], @@ -253,13 +347,19 @@ impl GlobalState { self.global_step = r.u64()?; let np = r.u32()? as usize; if np != self.params.len() { - bail!("checkpoint has {np} fragments, layout has {}", self.params.len()); + bail!( + "checkpoint has {np} fragments, layout has {}", + self.params.len() + ); } for p in 0..np { self.versions[p] = r.u64()?; let numel = r.u64()? as usize; if numel != self.params[p].len() { - bail!("checkpoint fragment {p} numel {numel} != layout {}", self.params[p].len()); + bail!( + "checkpoint fragment {p} numel {numel} != layout {}", + self.params[p].len() + ); } for slot in [&mut self.params[p], &mut self.momentum[p]] { for v in slot.iter_mut() { @@ -272,7 +372,11 @@ impl GlobalState { self.ledger.clear(); for _ in 0..nl { let id = r.u32()?; - let l = LearnerLedger { merges: r.u64()?, steps: r.u64()?, tokens: r.u64()? }; + let l = LearnerLedger { + merges: r.u64()?, + steps: r.u64()?, + tokens: r.u64()?, + }; self.ledger.insert(id, l); } if !r.0.is_empty() { @@ -300,8 +404,14 @@ mod tests { fn layout2() -> Layout { Layout { fragments: vec![ - FragmentInfo { merge_mode: MERGE_AVG, tensor_numels: vec![4] }, - FragmentInfo { merge_mode: MERGE_RDA, tensor_numels: vec![2, 2] }, + FragmentInfo { + merge_mode: MERGE_AVG, + tensor_numels: vec![4], + }, + FragmentInfo { + merge_mode: MERGE_RDA, + tensor_numels: vec![2, 2], + }, ], } } @@ -334,7 +444,13 @@ mod tests { let dir = std::env::temp_dir().join("yeto-ckpt-test"); std::fs::create_dir_all(&dir).unwrap(); let path = dir.join("state.ckpt"); - let mut st = GlobalState::new(layout2(), Some("{\"task\":\"nava\"}".to_string()), 0.7, 0.9, crate::protocol::DTYPE_F32); + let mut st = GlobalState::new( + layout2(), + Some("{\"task\":\"nava\"}".to_string()), + 0.7, + 0.9, + crate::protocol::DTYPE_F32, + ); st.init_fragment(0, vec![1.5; 4]).unwrap(); st.init_fragment(1, vec![-2.0; 4]).unwrap(); let learner = vec![0.0f32; 4]; @@ -344,7 +460,13 @@ mod tests { st.record_merge(3, 12, 4096); st.save_checkpoint(&path).unwrap(); - let mut st2 = GlobalState::new(layout2(), Some("{\"task\":\"nava\"}".to_string()), 0.7, 0.9, crate::protocol::DTYPE_F32); + let mut st2 = GlobalState::new( + layout2(), + Some("{\"task\":\"nava\"}".to_string()), + 0.7, + 0.9, + crate::protocol::DTYPE_F32, + ); st2.load_checkpoint(&path).unwrap(); assert_eq!(st2.global_step, 7); assert_eq!(st2.versions, vec![7, 0]); @@ -386,6 +508,75 @@ mod tests { ); } + #[test] + fn coord_midpoint_normmatch_changes_direction_but_keeps_production_norm() { + let anchor = vec![0.0f32; 4]; + let learners = [ + vec![-10.0f32, 0.0, 0.0, 0.0], // delta (10, 0, 0, 0) + vec![0.0f32, -10.0, 0.0, 0.0], // delta (0, 10, 0, 0) + vec![-1.0f32, -2.0, 0.0, 0.0], // midpoint direction (1, 2, 0, 0) + ]; + let refs: Vec<&[f32]> = learners.iter().map(|v| v.as_slice()).collect(); + let weights = [1.0, 1.0, 1.0]; + + let mut production = + GlobalState::new(layout2(), None, 1.0, 0.0, crate::protocol::DTYPE_F32); + production.init_fragment(0, anchor.clone()).unwrap(); + production.init_fragment(1, anchor.clone()).unwrap(); + production.merge_and_step(0, &refs, &weights).unwrap(); + + let mut robust = GlobalState::new(layout2(), None, 1.0, 0.0, crate::protocol::DTYPE_F32); + robust.merge_policy = MergePolicy::CoordMidpointNormmatch; + robust.init_fragment(0, anchor.clone()).unwrap(); + robust.init_fragment(1, anchor).unwrap(); + robust.merge_and_step(0, &refs, &weights).unwrap(); + + let production_delta: Vec = production.params[0].iter().map(|v| -v).collect(); + let robust_delta: Vec = robust.params[0].iter().map(|v| -v).collect(); + let production_norm = production_delta + .iter() + .map(|v| (*v as f64).powi(2)) + .sum::() + .sqrt(); + let robust_norm = robust_delta + .iter() + .map(|v| (*v as f64).powi(2)) + .sum::() + .sqrt(); + + assert!((production_norm - robust_norm).abs() < 1e-5); + assert!((robust_delta[1] / robust_delta[0] - 2.0).abs() < 1e-5); + assert!((production_delta[1] / production_delta[0] - 2.0).abs() > 0.1); + } + + #[test] + fn consensus_rda_floor50_damps_rda_fragment() { + let anchor = vec![0.0f32; 4]; + let learners = [vec![-2.0f32, 0.0, -2.0, 0.0], vec![0.0f32, -2.0, 0.0, -2.0]]; + let refs: Vec<&[f32]> = learners.iter().map(|v| v.as_slice()).collect(); + let weights = [1.0, 1.0]; + + let mut production = + GlobalState::new(layout2(), None, 1.0, 0.0, crate::protocol::DTYPE_F32); + production.init_fragment(0, anchor.clone()).unwrap(); + production.init_fragment(1, anchor.clone()).unwrap(); + production.merge_and_step(1, &refs, &weights).unwrap(); + + let mut consensus = GlobalState::new(layout2(), None, 1.0, 0.0, crate::protocol::DTYPE_F32); + consensus.merge_policy = MergePolicy::ConsensusRdaFloor50; + consensus.init_fragment(0, anchor.clone()).unwrap(); + consensus.init_fragment(1, anchor).unwrap(); + consensus.merge_and_step(1, &refs, &weights).unwrap(); + + let production_delta: Vec = production.params[1].iter().map(|v| -v).collect(); + let consensus_delta: Vec = consensus.params[1].iter().map(|v| -v).collect(); + let norm = |v: &[f32]| v.iter().map(|x| (*x as f64).powi(2)).sum::().sqrt(); + + assert!(norm(&consensus_delta) < norm(&production_delta)); + assert!((consensus_delta[0] - consensus_delta[1]).abs() < 1e-6); + assert!((consensus_delta[2] - consensus_delta[3]).abs() < 1e-6); + } + #[test] fn size_mismatch_rejected() { let mut st = GlobalState::new(layout2(), None, 0.7, 0.9, crate::protocol::DTYPE_F32);