Skip to content

Commit 03dd496

Browse files
committed
style: apply black formatting
Signed-off-by: Rusheel Sharma <rusheelhere@gmail.com>
1 parent 1a85bb0 commit 03dd496

2 files changed

Lines changed: 7 additions & 7 deletions

File tree

monai/losses/focal_loss.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -232,7 +232,9 @@ def forward(self, input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
232232
if mask is not None:
233233
spatial_dims = list(range(2, len(target.shape)))
234234
sum_mask = mask.sum(dim=spatial_dims, keepdim=True)
235-
loss = (loss * mask).sum(dim=spatial_dims, keepdim=True) / sum_mask.clamp(min=torch.finfo(mask.dtype).eps)
235+
loss = (loss * mask).sum(dim=spatial_dims, keepdim=True) / sum_mask.clamp(
236+
min=torch.finfo(mask.dtype).eps
237+
)
236238
else:
237239
loss = loss.mean(dim=list(range(2, len(target.shape))))
238240
loss = loss.sum()

tests/metrics/test_ignore_index_metrics.py

Lines changed: 4 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -80,9 +80,8 @@ def test_metric_ignore_consistency(self, metric_class, kwargs):
8080
torch.testing.assert_close(res1, res2, msg=f"Failed for {metric_class.__name__}")
8181

8282
@parameterized.expand(
83-
[(metric_class, kwargs, ignore_index)
84-
for metric_class, kwargs in TEST_METRICS
85-
for ignore_index in (0, 1)])
83+
[(metric_class, kwargs, ignore_index) for metric_class, kwargs in TEST_METRICS for ignore_index in (0, 1)]
84+
)
8685
def test_metric_ignore_class_index(self, metric_class, kwargs, ignore_index):
8786
metric = metric_class(ignore_index=ignore_index, **kwargs)
8887

@@ -175,9 +174,8 @@ def test_metric_ignore_consistency(self, metric_class, kwargs):
175174
torch.testing.assert_close(res1, res2, msg=f"Failed for {metric_class.__name__}")
176175

177176
@parameterized.expand(
178-
[(metric_class, kwargs, ignore_index)
179-
for metric_class, kwargs in SCIPY_METRICS
180-
for ignore_index in (0, 1)])
177+
[(metric_class, kwargs, ignore_index) for metric_class, kwargs in SCIPY_METRICS for ignore_index in (0, 1)]
178+
)
181179
def test_metric_ignore_class_index(self, metric_class, kwargs, ignore_index):
182180
metric = metric_class(ignore_index=ignore_index, **kwargs)
183181

0 commit comments

Comments
 (0)