@@ -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