Skip to content

Commit b8e1219

Browse files
Fix Accuracy API usage for torch metrics v 1.x+
1 parent c19ae67 commit b8e1219

File tree

1 file changed

+2
-2
lines changed

1 file changed

+2
-2
lines changed

tests/e2e/mnist.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -99,8 +99,8 @@ def __init__(self, data_dir=PATH_DATASETS, hidden_size=64, learning_rate=2e-4):
9999
nn.Linear(hidden_size, self.num_classes),
100100
)
101101

102-
self.val_accuracy = Accuracy()
103-
self.test_accuracy = Accuracy()
102+
self.val_accuracy = Accuracy(task="multiclass", num_classes=self.num_classes)
103+
self.test_accuracy = Accuracy(task="multiclass", num_classes=self.num_classes)
104104

105105
def forward(self, x):
106106
x = self.model(x)

0 commit comments

Comments
 (0)