Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions examples/emnist_fed_avg.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@ def loss(params, batch, rng):
server_optimizer = fedjax.optimizers.adam(
learning_rate=10**(-2.5), b1=0.9, b2=0.999, eps=10**(-4))
# Hyperparameters for client local traing dataset preparation.
client_batch_hparams = fedjax.ShuffleRepeatBatchHParams(batch_size=20)
client_batch_hparams = fedjax.ShuffleRepeatBatchHParams(batch_size=20) # pyrefly: ignore[unexpected-keyword]
algorithm = fed_avg.federated_averaging(grad_fn, client_optimizer,
server_optimizer,
client_batch_hparams)
Expand Down Expand Up @@ -88,9 +88,9 @@ def loss(params, batch, rng):

# Run evaluation metrics defined in `model.eval_metrics`.
train_metrics = fedjax.evaluate_model(model, server_state.params, # pytype: disable=wrong-arg-types # jax-ndarray
train_eval_batches)
train_eval_batches) # pyrefly: ignore[bad-argument-type]
test_metrics = fedjax.evaluate_model(model, server_state.params, # pytype: disable=wrong-arg-types # jax-ndarray
test_eval_batches)
test_eval_batches) # pyrefly: ignore[bad-argument-type]
print(f'[round {round_num}] train_metrics={train_metrics}')
print(f'[round {round_num}] test_metrics={test_metrics}')

Expand Down
6 changes: 3 additions & 3 deletions examples/fed_avg.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ def federated_averaging(

def init(params: fedjax.Params) -> ServerState:
opt_state = server_optimizer.init(params)
return ServerState(params, opt_state)
return ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

def apply(
server_state: ServerState,
Expand Down Expand Up @@ -98,6 +98,6 @@ def server_update(server_state, mean_delta_params):
opt_state, params = server_optimizer.apply(mean_delta_params,
server_state.opt_state,
server_state.params)
return ServerState(params, opt_state)
return ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

return fedjax.FederatedAlgorithm(init, apply)
return fedjax.FederatedAlgorithm(init, apply) # pyrefly: ignore[bad-argument-count]
10 changes: 5 additions & 5 deletions examples/stateful_fed_avg.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ def client_step(client_step_state, batch):
'rng': rng,
# Add to count of total number of steps of training for the client.
'state': ClientState(
num_steps=client_step_state['state'].num_steps + 1
num_steps=client_step_state['state'].num_steps + 1 # pyrefly: ignore[unexpected-keyword]
),
}
return next_client_step_state
Expand Down Expand Up @@ -134,7 +134,7 @@ def stateful_federated_averaging(
def init(params: fedjax.Params) -> ServerState:
opt_state = server_optimizer.init(params)
client_states = {}
return ServerState(params, opt_state, client_states)
return ServerState(params, opt_state, client_states) # pyrefly: ignore[bad-argument-count]

def apply(
server_state: ServerState,
Expand All @@ -148,7 +148,7 @@ def apply(
batch_cds = cds.shuffle_repeat_batch(client_batch_hparams)
if cid not in server_state.client_states:
# Initialize client state total training steps counter.
server_state.client_states[cid] = ClientState(num_steps=0)
server_state.client_states[cid] = ClientState(num_steps=0) # pyrefly: ignore[unexpected-keyword]
client_input = {'rng': crng, 'state': server_state.client_states[cid]}
batch_clients.append((cid, batch_cds, client_input))

Expand Down Expand Up @@ -184,6 +184,6 @@ def server_update(server_state, mean_delta_params):
opt_state, params = server_optimizer.apply(mean_delta_params,
server_state.opt_state,
server_state.params)
return ServerState(params, opt_state, server_state.client_states)
return ServerState(params, opt_state, server_state.client_states) # pyrefly: ignore[bad-argument-count]

return fedjax.FederatedAlgorithm(init, apply)
return fedjax.FederatedAlgorithm(init, apply) # pyrefly: ignore[bad-argument-count]
2 changes: 1 addition & 1 deletion experiments/fed_avg/run_fed_avg.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,7 +100,7 @@ def main(argv: Sequence[str]) -> None:
eval_batch_hparams)
}
if run_full_periodic_eval:
periodic_eval_fn_map.update(final_eval_fn_map)
periodic_eval_fn_map.update(final_eval_fn_map) # pyrefly: ignore[no-matching-overload]

init_state = algorithm.init(model.init(jax.random.PRNGKey(FLAGS.params_seed)))
fedjax.training.run_federated_experiment(
Expand Down
2 changes: 1 addition & 1 deletion fedjax/aggregators/aggregator.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,4 +72,4 @@ def extract_params_and_weight(clients_params_and_weight):
clients_params_and_weights)
return tree_util.tree_mean(params_and_weights), state

return Aggregator(init, apply)
return Aggregator(init, apply) # pyrefly: ignore[bad-argument-count]
40 changes: 20 additions & 20 deletions fedjax/aggregators/compression.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,13 +57,13 @@ def binary_stochastic_quantize(v: jnp.ndarray,
Quantized array.
"""
if v_min is None:
v_min = jnp.amin(v)
v_min = jnp.amin(v) # pyrefly: ignore[bad-assignment]
if v_max is None:
v_max = jnp.amax(v)
v = jnp.nan_to_num((v - v_min) / (v_max - v_min))
v_max = jnp.amax(v) # pyrefly: ignore[bad-assignment]
v = jnp.nan_to_num((v - v_min) / (v_max - v_min)) # pyrefly: ignore[unsupported-operation]
v = jnp.maximum(0., jnp.minimum(v, 1.))
rand = jax.random.uniform(key=rng, shape=v.shape)
return jnp.where(rand > v, v_min, v_max)
return jnp.where(rand > v, v_min, v_max) # pyrefly: ignore[bad-argument-type]


def uniform_stochastic_quantize(v: jnp.ndarray,
Expand All @@ -85,10 +85,10 @@ def uniform_stochastic_quantize(v: jnp.ndarray,
"""
# Rescale the vector to be between zero to one.
if v_min is None:
v_min = jnp.amin(v)
v_min = jnp.amin(v) # pyrefly: ignore[bad-assignment]
if v_max is None:
v_max = jnp.amax(v)
v = jnp.nan_to_num((v - v_min) / (v_max - v_min))
v_max = jnp.amax(v) # pyrefly: ignore[bad-assignment]
v = jnp.nan_to_num((v - v_min) / (v_max - v_min)) # pyrefly: ignore[unsupported-operation]
v = jnp.maximum(0., jnp.minimum(v, 1.))
# Compute the upper and lower boundary of each value.
v_ceil = jnp.ceil(v * (num_levels - 1)) / (num_levels - 1)
Expand All @@ -98,7 +98,7 @@ def uniform_stochastic_quantize(v: jnp.ndarray,
threshold = jnp.nan_to_num((v - v_floor) / (v_ceil - v_floor))
quantized = jnp.where(rand > threshold, v_floor, v_ceil)
# Rescale the values and return it.
return v_min + quantized * (v_max - v_min)
return v_min + quantized * (v_max - v_min) # pyrefly: ignore[unsupported-operation]


@jax.jit
Expand Down Expand Up @@ -171,7 +171,7 @@ def uniform_stochastic_quantizer(
"""

def init():
return CompressionState(0.0, rng)
return CompressionState(0.0, rng) # pyrefly: ignore[bad-argument-count]

def apply(
clients_params_and_weights: Iterable[Tuple[ClientId, Params, float]],
Expand Down Expand Up @@ -214,10 +214,10 @@ def arithmetic_encoding_num_bits_pytree(params, weights):
# 32 bits for every float and log2(num_levels) bit for every parameter.
new_bits = math.log2(
num_levels) * total_num_params + 32 * total_num_floats
new_state = CompressionState(aggregator_state.num_bits + new_bits, rng)
new_state = CompressionState(aggregator_state.num_bits + new_bits, rng) # pyrefly: ignore[bad-argument-count]
return aggregated_params, new_state

return aggregator.Aggregator(init, apply)
return aggregator.Aggregator(init, apply) # pyrefly: ignore[bad-argument-count]


def rotated_uniform_stochastic_quantizer(num_levels: int,
Expand All @@ -235,7 +235,7 @@ def rotated_uniform_stochastic_quantizer(num_levels: int,
"""

def init():
return CompressionState(0.0, rng)
return CompressionState(0.0, rng) # pyrefly: ignore[bad-argument-count]

def apply(
clients_params_and_weights: Iterable[Tuple[ClientId, Params, float]],
Expand Down Expand Up @@ -263,10 +263,10 @@ def quantize_params_and_weight(client_params_and_weight, rng):
total_num_floats = 2 * num_leaves(aggregated_params)
# 32 bits for every float used and log2(num_levels) bit for every parameter.
new_bits = math.log2(num_levels) * total_num_params + 32 * total_num_floats
new_state = CompressionState(aggregator_state.num_bits + new_bits, rng)
new_state = CompressionState(aggregator_state.num_bits + new_bits, rng) # pyrefly: ignore[bad-argument-count]
return aggregated_params, new_state

return aggregator.Aggregator(init, apply)
return aggregator.Aggregator(init, apply) # pyrefly: ignore[bad-argument-count]


@jax.jit
Expand All @@ -293,7 +293,7 @@ def structured_drive_quantizer(rng: PRNGKey) -> aggregator.Aggregator:
"""

def init():
return CompressionState(0.0, rng)
return CompressionState(0.0, rng) # pyrefly: ignore[bad-argument-count]

def apply(
clients_params_and_weights: Iterable[Tuple[ClientId, Params, float]],
Expand All @@ -319,10 +319,10 @@ def quantize_params_and_weight(client_params_and_weight, client_rng):
total_num_floats = 2 * num_leaves(aggregated_params)
# 32 bits for every float used and one bit for every parameter.
new_bits = total_num_params + 32 * total_num_floats
new_state = CompressionState(aggregator_state.num_bits + new_bits, rng)
new_state = CompressionState(aggregator_state.num_bits + new_bits, rng) # pyrefly: ignore[bad-argument-count]
return aggregated_params, new_state

return aggregator.Aggregator(init, apply)
return aggregator.Aggregator(init, apply) # pyrefly: ignore[bad-argument-count]


def terngrad_quantize(v: jnp.ndarray, rng: PRNGKey) -> jnp.ndarray:
Expand Down Expand Up @@ -373,7 +373,7 @@ def terngrad_quantizer(rng: PRNGKey) -> aggregator.Aggregator:
"""

def init():
return CompressionState(0.0, rng)
return CompressionState(0.0, rng) # pyrefly: ignore[bad-argument-count]

def apply(
clients_params_and_weights: Iterable[Tuple[ClientId, Params, float]],
Expand All @@ -394,7 +394,7 @@ def quantize_params_and_weight(client_params_and_weight, rng):
total_num_floats = 2 * num_leaves(aggregated_params)
# 32 bits for every float used and log2(3) bit for every parameter.
new_bits = math.log2(3) * total_num_params + 32 * total_num_floats
new_state = CompressionState(aggregator_state.num_bits + new_bits, rng)
new_state = CompressionState(aggregator_state.num_bits + new_bits, rng) # pyrefly: ignore[bad-argument-count]
return aggregated_params, new_state

return aggregator.Aggregator(init, apply)
return aggregator.Aggregator(init, apply) # pyrefly: ignore[bad-argument-count]
9 changes: 5 additions & 4 deletions fedjax/algorithms/agnostic_fed_avg.py
Original file line number Diff line number Diff line change
Expand Up @@ -230,8 +230,9 @@ def agnostic_federated_averaging(
if init_domain_window is None:
init_domain_window = jnp.ones_like(init_domain_weights) # pytype: disable=wrong-arg-types # jnp-type

if len(init_domain_weights) != len(init_domain_window):
if len(init_domain_weights) != len(init_domain_window): # pyrefly: ignore[bad-argument-type]
raise ValueError(
# pyrefly: ignore[bad-argument-type]
f'init_domain_weights and init_domain_window must be equal lengths.'
f' {len(init_domain_weights)} != {len(init_domain_window)}'
)
Expand All @@ -247,7 +248,7 @@ def init(params: Params) -> ServerState:
opt_state = server_optimizer.init(params)
domain_weights = jnp.array(init_domain_weights)
domain_window = [jnp.array(init_domain_window)] * domain_window_size
return ServerState(params, opt_state, domain_weights, domain_window)
return ServerState(params, opt_state, domain_weights, domain_window) # pyrefly: ignore[bad-argument-count]

def apply(
server_state: ServerState,
Expand Down Expand Up @@ -308,6 +309,6 @@ def server_update(server_state, mean_delta_params, sum_domain_loss,
domain_learning_rate,
domain_algorithm)
domain_window = server_state.domain_window[1:] + [sum_domain_num]
return ServerState(params, opt_state, domain_weights, domain_window)
return ServerState(params, opt_state, domain_weights, domain_window) # pyrefly: ignore[bad-argument-count]

return federated_algorithm.FederatedAlgorithm(init, apply)
return federated_algorithm.FederatedAlgorithm(init, apply) # pyrefly: ignore[bad-argument-count]
6 changes: 3 additions & 3 deletions fedjax/algorithms/fed_avg.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,7 +115,7 @@ def federated_averaging(

def init(params: Params) -> ServerState:
opt_state = server_optimizer.init(params)
return ServerState(params, opt_state)
return ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

def apply(
server_state: ServerState,
Expand Down Expand Up @@ -151,6 +151,6 @@ def server_update(server_state, mean_delta_params):
opt_state, params = server_optimizer.apply(mean_delta_params,
server_state.opt_state,
server_state.params)
return ServerState(params, opt_state)
return ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

return federated_algorithm.FederatedAlgorithm(init, apply)
return federated_algorithm.FederatedAlgorithm(init, apply) # pyrefly: ignore[bad-argument-count]
6 changes: 3 additions & 3 deletions fedjax/algorithms/fed_prox.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,7 @@ def fed_prox_loss(params, server_params, batch, rng):

def init(params: Params) -> ServerState:
opt_state = server_optimizer.init(params)
return ServerState(params, opt_state)
return ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

def apply(
server_state: ServerState,
Expand Down Expand Up @@ -146,6 +146,6 @@ def server_update(server_state, mean_delta_params):
opt_state, params = server_optimizer.apply(mean_delta_params,
server_state.opt_state,
server_state.params)
return ServerState(params, opt_state)
return ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

return federated_algorithm.FederatedAlgorithm(init, apply)
return federated_algorithm.FederatedAlgorithm(init, apply) # pyrefly: ignore[bad-argument-count]
8 changes: 4 additions & 4 deletions fedjax/algorithms/hyp_cluster.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,7 +81,7 @@ def hyp_cluster(

def init(cluster_params: List[Params]) -> ServerState:
return ServerState(
cluster_params,
cluster_params, # pyrefly: ignore[bad-argument-count]
[server_optimizer.init(params) for params in cluster_params])

# Creating these objects outside apply() can speed up repeated apply() calls.
Expand Down Expand Up @@ -131,9 +131,9 @@ def apply(
client_diagnostics = {}
for client_id, cluster_id in client_cluster_ids.items():
client_diagnostics[client_id] = {'cluster_id': cluster_id}
return ServerState(cluster_params, opt_states), client_diagnostics
return ServerState(cluster_params, opt_states), client_diagnostics # pyrefly: ignore[bad-argument-count]

return federated_algorithm.FederatedAlgorithm(init, apply)
return federated_algorithm.FederatedAlgorithm(init, apply) # pyrefly: ignore[bad-argument-count]


class _BaseClientTrainer:
Expand Down Expand Up @@ -250,7 +250,7 @@ def _cluster_losses(
for i, params in enumerate(cluster_params):
for client_id, average_loss in evaluator.evaluate_global_params(
params,
[(client_id, dataset.padded_batch(batch_hparams), rng[i])
[(client_id, dataset.padded_batch(batch_hparams), rng[i]) # pyrefly: ignore[bad-argument-type]
for (client_id, dataset, _), rng in zip(clients, client_rngs)]): # pytype: disable=wrong-arg-types # jax-ndarray
cluster_losses[client_id].append(average_loss)
return cluster_losses
Expand Down
6 changes: 3 additions & 3 deletions fedjax/algorithms/mime.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,7 +158,7 @@ def mime(

def init(params: Params) -> ServerState:
opt_state = base_optimizer.init(params)
return ServerState(params, opt_state)
return ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

def apply(
server_state: ServerState,
Expand Down Expand Up @@ -208,6 +208,6 @@ def server_update(server_state, server_grads, mean_delta_params):
mean_delta_params)
opt_state, _ = base_optimizer.apply(server_grads, server_state.opt_state,
server_state.params)
return ServerState(params, opt_state)
return ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

return federated_algorithm.FederatedAlgorithm(init, apply)
return federated_algorithm.FederatedAlgorithm(init, apply) # pyrefly: ignore[bad-argument-count]
6 changes: 3 additions & 3 deletions fedjax/algorithms/mime_lite.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,7 +109,7 @@ def mime_lite(

def init(params: Params) -> mime.ServerState:
opt_state = base_optimizer.init(params)
return mime.ServerState(params, opt_state)
return mime.ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

def apply(
server_state: mime.ServerState,
Expand Down Expand Up @@ -167,6 +167,6 @@ def server_update(server_state, server_grads, mean_delta_params):
mean_delta_params)
opt_state, _ = base_optimizer.apply(server_grads, server_state.opt_state,
server_state.params)
return mime.ServerState(params, opt_state)
return mime.ServerState(params, opt_state) # pyrefly: ignore[bad-argument-count]

return federated_algorithm.FederatedAlgorithm(init, apply)
return federated_algorithm.FederatedAlgorithm(init, apply) # pyrefly: ignore[bad-argument-count]
10 changes: 5 additions & 5 deletions fedjax/core/client_datasets.py
Original file line number Diff line number Diff line change
Expand Up @@ -323,7 +323,7 @@ def a_library_function(client_dataset, hparams):
if hparams is None:
hparams = PaddedBatchHParams(**kwargs)
elif kwargs:
hparams = hparams.replace(**kwargs)
hparams = hparams.replace(**kwargs) # pyrefly: ignore[missing-attribute]
return PaddedBatchView(self, hparams)

def shuffle_repeat_batch(self,
Expand Down Expand Up @@ -389,7 +389,7 @@ def a_library_function(client_dataset, hparams):
if hparams is None:
hparams = ShuffleRepeatBatchHParams(**kwargs)
elif kwargs:
hparams = hparams.replace(**kwargs)
hparams = hparams.replace(**kwargs) # pyrefly: ignore[missing-attribute]
return ShuffleRepeatBatchView(self, hparams)

def batch(self,
Expand Down Expand Up @@ -427,7 +427,7 @@ def a_library_function(client_dataset, hparams):
if hparams is None:
hparams = BatchHParams(**kwargs)
elif kwargs:
hparams = hparams.replace(**kwargs)
hparams = hparams.replace(**kwargs) # pyrefly: ignore[missing-attribute]
return BatchView(self, hparams)


Expand Down Expand Up @@ -600,7 +600,7 @@ def a_library_function(datasets, hparams):
if hparams is None:
hparams = PaddedBatchHParams(**kwargs)
elif kwargs:
hparams = hparams.replace(**kwargs)
hparams = hparams.replace(**kwargs) # pyrefly: ignore[missing-attribute]
preprocessor = None
features = None
# Pieces of examples whose total size is < batch_size
Expand Down Expand Up @@ -650,7 +650,7 @@ def a_library_function(datasets, hparams):
buf.append(slice_examples(examples, slice(start, size)))
buf_size += size - start
if buf:
final_examples = preprocessor(concat_examples(buf))
final_examples = preprocessor(concat_examples(buf)) # pyrefly: ignore[not-callable]
final_batch_size = _pick_final_batch_size(buf_size, hparams.batch_size,
hparams.num_batch_size_buckets)
yield pad_examples(final_examples, final_batch_size)
Expand Down
Loading
Loading