Skip to content

Commit 579ebc1

Browse files
committed
fix formatting for ruff
1 parent 6634043 commit 579ebc1

28 files changed

Lines changed: 1012 additions & 822 deletions

File tree

algoperf/jax_sharding_utils.py

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,8 @@ def shard_along_batch_dim(x):
2020
"""Shards a tensor across all devices."""
2121
mesh = jax.sharding.Mesh(jax.devices(), ('batch',))
2222
return jax.tree.map(
23-
lambda x: jax.device_put(x, NamedSharding(mesh, P('batch'))), x)
23+
lambda x: jax.device_put(x, NamedSharding(mesh, P('batch'))), x
24+
)
2425

2526

2627
def replicate(x):
@@ -32,5 +33,7 @@ def replicate(x):
3233
def display_shard_info(x: jax.Array):
3334
"""Displays shard info of a jax array."""
3435
for shard in x.addressable_shards:
35-
print(f"shard.device: {shard.device}, index: {shard.index}, replica_id:"
36-
f" {shard.replica_id}.\n")
36+
print(
37+
f'shard.device: {shard.device}, index: {shard.index}, replica_id:'
38+
f' {shard.replica_id}.\n'
39+
)

algoperf/workloads/cifar/cifar_jax/workload.py

Lines changed: 25 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -176,35 +176,38 @@ def _compute_metrics(
176176
return metrics
177177

178178
def _eval_model(
179-
self,
180-
params: spec.ParameterContainer,
181-
batch: Dict[str, spec.Tensor],
182-
model_state: spec.ModelAuxiliaryState,
183-
rng: spec.RandomState) -> Dict[spec.Tensor, spec.ModelAuxiliaryState]:
179+
self,
180+
params: spec.ParameterContainer,
181+
batch: Dict[str, spec.Tensor],
182+
model_state: spec.ModelAuxiliaryState,
183+
rng: spec.RandomState,
184+
) -> Dict[spec.Tensor, spec.ModelAuxiliaryState]:
184185
"""Return the mean accuracy and loss as a dict."""
185186

186187
@functools.partial(
187-
jax.jit,
188-
in_shardings=(
189-
jax_sharding_utils.get_replicate_sharding(), # params
190-
jax_sharding_utils.get_batch_dim_sharding(), # batch
191-
jax_sharding_utils.get_replicate_sharding(), # model_state
192-
jax_sharding_utils.get_batch_dim_sharding(), # rng
193-
),
188+
jax.jit,
189+
in_shardings=(
190+
jax_sharding_utils.get_replicate_sharding(), # params
191+
jax_sharding_utils.get_batch_dim_sharding(), # batch
192+
jax_sharding_utils.get_replicate_sharding(), # model_state
193+
jax_sharding_utils.get_batch_dim_sharding(), # rng
194+
),
194195
)
195196
def _eval_model_jitted(
196-
params: spec.ParameterContainer,
197-
batch: Dict[str, spec.Tensor],
198-
model_state: spec.ModelAuxiliaryState,
199-
rng: spec.RandomState) -> Dict[spec.Tensor, spec.ModelAuxiliaryState]:
197+
params: spec.ParameterContainer,
198+
batch: Dict[str, spec.Tensor],
199+
model_state: spec.ModelAuxiliaryState,
200+
rng: spec.RandomState,
201+
) -> Dict[spec.Tensor, spec.ModelAuxiliaryState]:
200202
"""Return the mean accuracy and loss as a dict."""
201203
logits, _ = self.model_fn(
202-
params,
203-
batch,
204-
model_state,
205-
spec.ForwardPassMode.EVAL,
206-
rng,
207-
update_batch_norm=False)
204+
params,
205+
batch,
206+
model_state,
207+
spec.ForwardPassMode.EVAL,
208+
rng,
209+
update_batch_norm=False,
210+
)
208211
weights = batch.get('weights')
209212
if weights is None:
210213
weights = jnp.ones(len(logits))

algoperf/workloads/criteo1tb/criteo1tb_jax/workload.py

Lines changed: 11 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -134,16 +134,17 @@ def model_fn(
134134
return logits_batch, None
135135

136136
@functools.partial(
137-
jax.jit,
138-
in_shardings=(
139-
jax_sharding_utils.get_replicate_sharding(),
140-
jax_sharding_utils.get_batch_dim_sharding(),
141-
),
142-
static_argnums=(0,),
143-
out_shardings=jax_sharding_utils.get_replicate_sharding())
144-
def _eval_batch_jitted(self,
145-
params: spec.ParameterContainer,
146-
batch: Dict[str, spec.Tensor]) -> spec.Tensor:
137+
jax.jit,
138+
in_shardings=(
139+
jax_sharding_utils.get_replicate_sharding(),
140+
jax_sharding_utils.get_batch_dim_sharding(),
141+
),
142+
static_argnums=(0,),
143+
out_shardings=jax_sharding_utils.get_replicate_sharding(),
144+
)
145+
def _eval_batch_jitted(
146+
self, params: spec.ParameterContainer, batch: Dict[str, spec.Tensor]
147+
) -> spec.Tensor:
147148
logits, _ = self.model_fn(
148149
params,
149150
batch,

algoperf/workloads/fastmri/fastmri_jax/workload.py

Lines changed: 17 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -101,16 +101,21 @@ def loss_fn(
101101
}
102102

103103
@functools.partial(
104-
jax.jit,
105-
in_shardings=(jax_sharding_utils.get_replicate_sharding(),
106-
jax_sharding_utils.get_batch_dim_sharding(),
107-
jax_sharding_utils.get_replicate_sharding()),
108-
static_argnums=(0,),
109-
out_shardings=jax_sharding_utils.get_replicate_sharding())
110-
def _eval_model(self,
111-
params: spec.Tensor,
112-
batch: Dict[str, spec.Tensor],
113-
rng: spec.RandomState) -> Dict[str, spec.Tensor]:
104+
jax.jit,
105+
in_shardings=(
106+
jax_sharding_utils.get_replicate_sharding(),
107+
jax_sharding_utils.get_batch_dim_sharding(),
108+
jax_sharding_utils.get_replicate_sharding(),
109+
),
110+
static_argnums=(0,),
111+
out_shardings=jax_sharding_utils.get_replicate_sharding(),
112+
)
113+
def _eval_model(
114+
self,
115+
params: spec.Tensor,
116+
batch: Dict[str, spec.Tensor],
117+
rng: spec.RandomState,
118+
) -> Dict[str, spec.Tensor]:
114119
"""Return the SSIM and loss as a dict."""
115120
logits, _ = self.model_fn(
116121
params,
@@ -166,13 +171,13 @@ def _eval_model_on_split(
166171
num_batches=num_batches,
167172
)
168173

169-
total_metrics = {'ssim': 0., 'loss': 0.}
174+
total_metrics = {'ssim': 0.0, 'loss': 0.0}
170175
for _ in range(num_batches):
171176
batch = next(self._eval_iters[split])
172177
# We already sum these metrics across devices inside _eval_model.
173178
synced_metrics = self._eval_model(params, batch, model_rng)
174179
total_metrics = {
175-
k: v + synced_metrics[k] for k, v in total_metrics.items()
180+
k: v + synced_metrics[k] for k, v in total_metrics.items()
176181
}
177182
return {k: float(v.item() / num_examples) for k, v in total_metrics.items()}
178183

algoperf/workloads/imagenet_resnet/imagenet_jax/workload.py

Lines changed: 23 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -107,33 +107,36 @@ def init_model_fn(
107107
self._param_shapes = param_utils.jax_param_shapes(params)
108108
self._param_types = param_utils.jax_param_types(self._param_shapes)
109109
params = jax.tree.map(
110-
lambda x: jax.device_put(x, jax_sharding_utils.get_replicate_sharding()
111-
),
112-
params)
110+
lambda x: jax.device_put(x, jax_sharding_utils.get_replicate_sharding()),
111+
params,
112+
)
113113
model_state = jax.tree.map(
114-
lambda x: jax.device_put(x, jax_sharding_utils.get_replicate_sharding()
115-
),
116-
model_state)
114+
lambda x: jax.device_put(x, jax_sharding_utils.get_replicate_sharding()),
115+
model_state,
116+
)
117117
return params, model_state
118118

119119
def is_output_params(self, param_key: spec.ParameterKey) -> bool:
120120
return param_key == 'Dense_0'
121121

122122
@functools.partial(
123-
jax.jit,
124-
in_shardings=(
125-
jax_sharding_utils.get_replicate_sharding(), # params
126-
jax_sharding_utils.get_batch_dim_sharding(), # batch
127-
jax_sharding_utils.get_replicate_sharding(), # model_state
128-
jax_sharding_utils.get_replicate_sharding(), # rng
129-
),
130-
static_argnums=(0,),
131-
out_shardings=jax_sharding_utils.get_replicate_sharding())
132-
def _eval_model(self,
133-
params: spec.ParameterContainer,
134-
batch: Dict[str, spec.Tensor],
135-
model_state: spec.ModelAuxiliaryState,
136-
rng: spec.RandomState) -> Dict[str, spec.Tensor]:
123+
jax.jit,
124+
in_shardings=(
125+
jax_sharding_utils.get_replicate_sharding(), # params
126+
jax_sharding_utils.get_batch_dim_sharding(), # batch
127+
jax_sharding_utils.get_replicate_sharding(), # model_state
128+
jax_sharding_utils.get_replicate_sharding(), # rng
129+
),
130+
static_argnums=(0,),
131+
out_shardings=jax_sharding_utils.get_replicate_sharding(),
132+
)
133+
def _eval_model(
134+
self,
135+
params: spec.ParameterContainer,
136+
batch: Dict[str, spec.Tensor],
137+
model_state: spec.ModelAuxiliaryState,
138+
rng: spec.RandomState,
139+
) -> Dict[str, spec.Tensor]:
137140
logits, _ = self.model_fn(
138141
params,
139142
batch,

algoperf/workloads/imagenet_vit/imagenet_jax/workload.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010

1111
from algoperf import param_utils, spec, jax_sharding_utils
1212
from algoperf.workloads.imagenet_resnet.imagenet_jax.workload import (
13-
ImagenetResNetWorkload,
13+
ImagenetResNetWorkload,
1414
)
1515
from algoperf.workloads.imagenet_vit.imagenet_jax import models
1616
from algoperf.workloads.imagenet_vit.workload import (

algoperf/workloads/librispeech_conformer/librispeech_jax/workload.py

Lines changed: 43 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -15,8 +15,9 @@
1515
from algoperf import spec
1616
from algoperf.workloads.librispeech_conformer import metrics
1717
from algoperf.workloads.librispeech_conformer import workload
18-
from algoperf.workloads.librispeech_conformer.input_pipeline import \
19-
LibriSpeechDataset
18+
from algoperf.workloads.librispeech_conformer.input_pipeline import (
19+
LibriSpeechDataset,
20+
)
2021
from algoperf.workloads.librispeech_conformer.librispeech_jax import models
2122

2223

@@ -317,21 +318,23 @@ def greedy_decode(
317318
return hyp, hyp_paddings
318319

319320
@functools.partial(
320-
jax.jit,
321-
in_shardings=(
322-
jax_sharding_utils.get_replicate_sharding(), # params
323-
jax_sharding_utils.get_batch_dim_sharding(), # batch
324-
jax_sharding_utils.get_replicate_sharding(), # model_state
325-
jax_sharding_utils.get_replicate_sharding(), # rng
326-
),
327-
out_shardings=jax_sharding_utils.get_batch_dim_sharding(),
328-
static_argnums=(0,))
321+
jax.jit,
322+
in_shardings=(
323+
jax_sharding_utils.get_replicate_sharding(), # params
324+
jax_sharding_utils.get_batch_dim_sharding(), # batch
325+
jax_sharding_utils.get_replicate_sharding(), # model_state
326+
jax_sharding_utils.get_replicate_sharding(), # rng
327+
),
328+
out_shardings=jax_sharding_utils.get_batch_dim_sharding(),
329+
static_argnums=(0,),
330+
)
329331
def _eval_step(
330-
self,
331-
params: spec.ParameterContainer,
332-
batch: Dict[str, spec.Tensor],
333-
model_state: spec.ModelAuxiliaryState,
334-
rng: spec.RandomState) -> Tuple[spec.Tensor, spec.ModelAuxiliaryState]:
332+
self,
333+
params: spec.ParameterContainer,
334+
batch: Dict[str, spec.Tensor],
335+
model_state: spec.ModelAuxiliaryState,
336+
rng: spec.RandomState,
337+
) -> Tuple[spec.Tensor, spec.ModelAuxiliaryState]:
335338
(logits, logit_paddings), _ = self.model_fn(
336339
params,
337340
batch,
@@ -346,40 +349,38 @@ def _eval_step(
346349
targets, target_paddings = batch['targets']
347350
# Convert metrics bundle to dictionary
348351
metrics_dict = {
349-
'loss_per_example':
350-
loss['per_example'],
351-
'decoded':
352-
decoded,
353-
'decoded_paddings':
354-
decoded_paddings,
355-
'targets':
356-
targets,
357-
'target_paddings':
358-
target_paddings,
359-
'n_valid_examples':
360-
jnp.zeros((len(jax.devices()), 1)) + loss['n_valid_examples']
352+
'loss_per_example': loss['per_example'],
353+
'decoded': decoded,
354+
'decoded_paddings': decoded_paddings,
355+
'targets': targets,
356+
'target_paddings': target_paddings,
357+
'n_valid_examples': jnp.zeros((len(jax.devices()), 1))
358+
+ loss['n_valid_examples'],
361359
}
362360
return metrics_dict
363361

364-
def eval_step(self,
365-
params: spec.ParameterContainer,
366-
batch: Dict[str, spec.Tensor],
367-
model_state: spec.ModelAuxiliaryState,
368-
rng: spec.RandomState):
362+
def eval_step(
363+
self,
364+
params: spec.ParameterContainer,
365+
batch: Dict[str, spec.Tensor],
366+
model_state: spec.ModelAuxiliaryState,
367+
rng: spec.RandomState,
368+
):
369369
"""Evaluates the model and returns a metrics bundle."""
370370
metrics_dict = self._eval_step(params, batch, model_state, rng)
371371

372372
# Convert dictionary back to metrics bundle
373373
metrics_bundle = self.metrics_bundle.single_from_model_output(
374-
loss_dict={
375-
'summed': metrics_dict['loss_per_example'].sum(),
376-
'per_example': metrics_dict['loss_per_example'],
377-
'n_valid_examples': metrics_dict['n_valid_examples'].sum()
378-
},
379-
decoded=metrics_dict['decoded'],
380-
decoded_paddings=metrics_dict['decoded_paddings'],
381-
targets=metrics_dict['targets'],
382-
target_paddings=metrics_dict['target_paddings'])
374+
loss_dict={
375+
'summed': metrics_dict['loss_per_example'].sum(),
376+
'per_example': metrics_dict['loss_per_example'],
377+
'n_valid_examples': metrics_dict['n_valid_examples'].sum(),
378+
},
379+
decoded=metrics_dict['decoded'],
380+
decoded_paddings=metrics_dict['decoded_paddings'],
381+
targets=metrics_dict['targets'],
382+
target_paddings=metrics_dict['target_paddings'],
383+
)
383384

384385
return metrics_bundle
385386

0 commit comments

Comments
 (0)