pub trait StrategyContext {
// Required methods
fn num_workers(&self) -> usize;
fn execute_on_worker(
&self,
worker_idx: usize,
plan: &Value,
input: &Value,
y: Option<&Value>,
) -> Result<HashMap<String, Value>>;
fn get_state(
&self,
worker_idx: usize,
node_ids: &[String],
) -> Result<HashMap<String, Value>>;
fn set_state(
&self,
worker_idx: usize,
states: &HashMap<String, Value>,
) -> Result<()>;
fn get_gradients(
&self,
worker_idx: usize,
node_ids: &[String],
) -> Result<HashMap<String, Value>>;
fn apply_gradients(
&self,
worker_idx: usize,
gradients: &HashMap<String, Value>,
) -> Result<()>;
// Provided methods
fn read_back_state(
&self,
worker_idx: usize,
node_ids: &[String],
) -> Result<HashMap<String, Value>> { ... }
fn execute_partition(
&self,
_worker_idx: usize,
_node_ids: &[String],
_input: &Value,
_y: Option<&Value>,
) -> Result<(Value, HashMap<String, Value>)> { ... }
fn worker_for(&self, target: &RemoteTarget) -> Result<usize> { ... }
}Expand description
Context provided to strategy executors. Abstracts worker communication — the strategy doesn’t know about WS/HTTP.
Required Methods§
Sourcefn num_workers(&self) -> usize
fn num_workers(&self) -> usize
Number of available workers.
Sourcefn execute_on_worker(
&self,
worker_idx: usize,
plan: &Value,
input: &Value,
y: Option<&Value>,
) -> Result<HashMap<String, Value>>
fn execute_on_worker( &self, worker_idx: usize, plan: &Value, input: &Value, y: Option<&Value>, ) -> Result<HashMap<String, Value>>
Execute a plan on a specific worker (by index). Returns trained states.
Sourcefn get_state(
&self,
worker_idx: usize,
node_ids: &[String],
) -> Result<HashMap<String, Value>>
fn get_state( &self, worker_idx: usize, node_ids: &[String], ) -> Result<HashMap<String, Value>>
Get trained states from a worker.
Sourcefn set_state(
&self,
worker_idx: usize,
states: &HashMap<String, Value>,
) -> Result<()>
fn set_state( &self, worker_idx: usize, states: &HashMap<String, Value>, ) -> Result<()>
Set states on a worker (e.g. after aggregation).
Provided Methods§
Sourcefn read_back_state(
&self,
worker_idx: usize,
node_ids: &[String],
) -> Result<HashMap<String, Value>>
fn read_back_state( &self, worker_idx: usize, node_ids: &[String], ) -> Result<HashMap<String, Value>>
Read a worker’s state now, over the wire, rather than recalling what its last fit returned.
The two differ exactly when something changed the model after the
fit — which is what apply_gradients does.
A data-parallel round that finished with get_state handed back the
weights each replica had before the averaged gradient was applied,
so the training it had just done was discarded on the way out.
Defaults to get_state, for a context whose two answers cannot
differ.
Sourcefn execute_partition(
&self,
_worker_idx: usize,
_node_ids: &[String],
_input: &Value,
_y: Option<&Value>,
) -> Result<(Value, HashMap<String, Value>)>
fn execute_partition( &self, _worker_idx: usize, _node_ids: &[String], _input: &Value, _y: Option<&Value>, ) -> Result<(Value, HashMap<String, Value>)>
Run part of the graph on a worker, returning the activation and the states it learned.
This is what model parallelism needs and data parallelism does not: every other strategy runs the whole plan on each worker and only ever wants the states back. Here each worker holds a slice of the model, so its output is the next worker’s input.
Defaults to refusing, so a context that cannot address part of a plan says so instead of silently running all of it.
Sourcefn worker_for(&self, target: &RemoteTarget) -> Result<usize>
fn worker_for(&self, target: &RemoteTarget) -> Result<usize>
Which worker answers to target.
Every other strategy indexes workers by position, because every worker is interchangeable to it. A partition is pinned to one, so it has to be found by id or tag.
Dyn Compatibility§
This trait is dyn compatible.
In older versions of Rust, dyn compatibility was called "object safety".