diff --git a/enn/networks/combiners.py b/enn/networks/combiners.py index 73cf964..e618f45 100644 --- a/enn/networks/combiners.py +++ b/enn/networks/combiners.py @@ -145,4 +145,4 @@ def sample_logits(sub_key: chex.PRNGKey) -> chex.Array: return jax.vmap(sample_logits)(enn_keys) - return jax.jit(enn_batch_fwd) + return jax.jit(enn_batch_fwd) # pyrefly: ignore[bad-return] diff --git a/enn/networks/forwarders.py b/enn/networks/forwarders.py index 08809ae..a3adcfc 100644 --- a/enn/networks/forwarders.py +++ b/enn/networks/forwarders.py @@ -58,4 +58,4 @@ def forward(params: hk.Params, state: hk.State, x: base.Input) -> chex.Array: net_out, unused_state = batch_apply(params, state, x, indices) return utils.parse_net_output(net_out) - return jax.jit(forward) + return jax.jit(forward) # pyrefly: ignore[bad-return] diff --git a/enn/supervised/multiloss_experiment.py b/enn/supervised/multiloss_experiment.py index 86bfa65..ee8ba7e 100644 --- a/enn/supervised/multiloss_experiment.py +++ b/enn/supervised/multiloss_experiment.py @@ -146,7 +146,7 @@ def train(self, num_batches: int): # Periodically log this performance as dataset=train. if self.step % self._train_log_freq == 0: - loss_metrics.update({ + loss_metrics.update({ # pyrefly: ignore[no-matching-overload] 'dataset': 'train', 'step': self.step, 'sgd': True, diff --git a/enn/supervised/sgd_experiment.py b/enn/supervised/sgd_experiment.py index 51fa339..5d30be2 100644 --- a/enn/supervised/sgd_experiment.py +++ b/enn/supervised/sgd_experiment.py @@ -160,7 +160,7 @@ def train(self, num_batches: int): # Periodically log this performance as dataset=train. if self.step % self._train_log_freq == 0: - loss_metrics.update( + loss_metrics.update( # pyrefly: ignore[no-matching-overload] {'dataset': 'train', 'step': self.step, 'sgd': True}) self.logger.write(loss_metrics)