1use somatize_core::error::{Result, SomaError};
13use somatize_core::strategy::{FederatedAggregation, GradientAggregation, TrainingStrategy};
14use somatize_core::value::Value;
15use std::collections::HashMap;
16
17pub trait StrategyContext {
26 fn num_workers(&self) -> usize;
28
29 fn execute_on_worker(
31 &self,
32 worker_idx: usize,
33 plan: &serde_json::Value,
34 input: &Value,
35 y: Option<&Value>,
36 ) -> Result<HashMap<String, Value>>;
37
38 fn get_state(&self, worker_idx: usize, node_ids: &[String]) -> Result<HashMap<String, Value>>;
40
41 fn set_state(&self, worker_idx: usize, states: &HashMap<String, Value>) -> Result<()>;
43
44 fn get_gradients(
46 &self,
47 worker_idx: usize,
48 node_ids: &[String],
49 ) -> Result<HashMap<String, Value>>;
50
51 fn apply_gradients(&self, worker_idx: usize, gradients: &HashMap<String, Value>) -> Result<()>;
53}
54
55pub trait StrategyExecutor {
58 fn fit(
60 &self,
61 ctx: &dyn StrategyContext,
62 input: &Value,
63 y: Option<&Value>,
64 node_ids: &[String],
65 ) -> Result<HashMap<String, Value>>;
66}
67
68pub trait GradientAggregator {
70 fn aggregate(&self, gradients: &[HashMap<String, Value>]) -> Result<HashMap<String, Value>>;
73}
74
75pub trait StateAggregator {
77 fn aggregate(&self, states: &[HashMap<String, Value>]) -> Result<HashMap<String, Value>>;
80}
81
82impl StrategyExecutor for TrainingStrategy {
83 fn fit(
84 &self,
85 ctx: &dyn StrategyContext,
86 input: &Value,
87 y: Option<&Value>,
88 node_ids: &[String],
89 ) -> Result<HashMap<String, Value>> {
90 match self {
91 TrainingStrategy::Local => {
92 ctx.execute_on_worker(0, &serde_json::json!({}), input, y)
94 }
95
96 TrainingStrategy::DataParallel {
97 num_replicas,
98 aggregation,
99 } => {
100 let n = (*num_replicas).min(ctx.num_workers());
101 let shards = shard_value(input, n);
102
103 for (i, shard) in shards.iter().enumerate() {
105 ctx.execute_on_worker(i, &serde_json::json!({}), shard, y)?;
106 }
107
108 let mut all_grads = Vec::new();
110 for i in 0..n {
111 all_grads.push(ctx.get_gradients(i, node_ids)?);
112 }
113 let averaged = aggregation.aggregate(&all_grads)?;
114
115 for i in 0..n {
117 ctx.apply_gradients(i, &averaged)?;
118 }
119
120 ctx.get_state(0, node_ids)
122 }
123
124 TrainingStrategy::Federated {
125 num_clients,
126 rounds,
127 aggregation,
128 ..
129 } => {
130 let n = (*num_clients).min(ctx.num_workers());
131 let shards = shard_value(input, n);
132
133 for _round in 0..*rounds {
134 for (i, shard) in shards.iter().enumerate().take(n) {
136 ctx.execute_on_worker(i, &serde_json::json!({}), shard, y)?;
137 }
138
139 let mut all_states = Vec::new();
141 for i in 0..n {
142 all_states.push(ctx.get_state(i, node_ids)?);
143 }
144 let aggregated = aggregation.aggregate(&all_states)?;
145
146 for i in 0..n {
148 ctx.set_state(i, &aggregated)?;
149 }
150 }
151
152 ctx.get_state(0, node_ids)
153 }
154
155 TrainingStrategy::ModelParallel { .. } => {
156 Err(SomaError::Other(
158 "ModelParallel strategy execution not yet implemented".into(),
159 ))
160 }
161
162 TrainingStrategy::PopulationBased { .. } => {
163 Err(SomaError::Other(
165 "PopulationBased strategy execution not yet implemented".into(),
166 ))
167 }
168
169 TrainingStrategy::Custom { .. } => Err(SomaError::Other(
170 "Custom strategy requires a user-provided coordinator".into(),
171 )),
172
173 other => Err(SomaError::Other(format!(
179 "this runtime does not know how to run {other:?}. It was \
180 probably described by a newer version"
181 ))),
182 }
183 }
184}
185
186impl GradientAggregator for GradientAggregation {
193 fn aggregate(&self, gradients: &[HashMap<String, Value>]) -> Result<HashMap<String, Value>> {
194 if gradients.len() == 1 {
197 return Ok(gradients[0].clone());
198 }
199 Err(SomaError::Other(format!(
200 "{self:?} gradient aggregation over {} workers is not implemented yet; \
201 it would need element-wise tensor averaging",
202 gradients.len()
203 )))
204 }
205}
206
207impl StateAggregator for FederatedAggregation {
208 fn aggregate(&self, states: &[HashMap<String, Value>]) -> Result<HashMap<String, Value>> {
209 if states.len() == 1 {
210 return Ok(states[0].clone());
211 }
212 Err(SomaError::Other(format!(
213 "{self:?} state aggregation over {} clients is not implemented yet; \
214 it would need element-wise tensor averaging",
215 states.len()
216 )))
217 }
218}
219
220fn shard_value(value: &Value, n: usize) -> Vec<Value> {
222 match value {
223 Value::Tensor { values, shape } if !shape.is_empty() && shape[0] >= n => {
224 let rows = shape[0];
225 let row_size: usize = shape[1..].iter().product::<usize>().max(1);
226 let shard_rows = rows / n;
227 let mut shards = Vec::new();
228 for i in 0..n {
229 let start = i * shard_rows;
230 let end = if i == n - 1 { rows } else { start + shard_rows };
231 let flat_start = start * row_size;
232 let flat_end = end * row_size;
233 let shard_vals = values[flat_start..flat_end].to_vec();
234 let mut shard_shape = shape.clone();
235 shard_shape[0] = end - start;
236 shards.push(Value::tensor(shard_vals, shard_shape));
237 }
238 shards
239 }
240 _ => (0..n).map(|_| value.clone()).collect(),
241 }
242}
243
244#[cfg(test)]
245mod tests {
246 use super::*;
247
248 #[test]
252 fn multi_worker_aggregation_refuses_instead_of_guessing() {
253 let grads = |v: f64| HashMap::from([("w".to_string(), Value::tensor(vec![v], vec![1]))]);
254
255 let err = GradientAggregation::AllReduce
256 .aggregate(&[grads(1.0), grads(3.0)])
257 .expect_err("aggregating two workers must not silently succeed");
258 assert!(err.to_string().contains("not implemented"), "{err}");
259
260 let err = FederatedAggregation::FedAvg
261 .aggregate(&[grads(1.0), grads(3.0)])
262 .expect_err("aggregating two clients must not silently succeed");
263 assert!(err.to_string().contains("not implemented"), "{err}");
264 }
265
266 #[test]
269 fn single_worker_aggregation_is_the_identity() {
270 let only = HashMap::from([("w".to_string(), Value::tensor(vec![2.0], vec![1]))]);
271 let out = GradientAggregation::AllReduce
272 .aggregate(std::slice::from_ref(&only))
273 .unwrap();
274 assert_eq!(out, only);
275 }
276}