-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathcompute_metrics_avalanche.py
More file actions
56 lines (53 loc) · 1.94 KB
/
Copy pathcompute_metrics_avalanche.py
File metadata and controls
56 lines (53 loc) · 1.94 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
import torch
from torch.nn import CrossEntropyLoss
from torch.optim import SGD
from avalanche.benchmarks.classic import SplitMNIST
from avalanche.models import SimpleMLP
from avalanche.evaluation.metrics import accuracy_metrics, loss_metrics, forgetting_metrics, timing_metrics, cpu_usage_metrics
from avalanche.logging import InteractiveLogger
from avalanche.training.plugins import EvaluationPlugin
from avalanche.training.supervised import Naive
def main():
device = torch.device("cuda" if torch.cuda.is_available()
else "cpu")
print(f"Using device: {device}")
# 1. Benchmark: Split MNIST into 5 experiences
benchmark = SplitMNIST(n_experiences=5)
# 2. Model: simple MLP
model = SimpleMLP(num_classes=benchmark.n_classes).to(device)
# 3. Optimizer & loss
optimizer = SGD(model.parameters(), lr=0.001, momentum=0.9)
criterion = CrossEntropyLoss()
# 4. Logger: only interactive console logging
interactive_logger = InteractiveLogger()
# 5. Evaluation plugin with metrics + interactive logger
eval_plugin = EvaluationPlugin(
accuracy_metrics(minibatch=True, epoch=True,
experience=True, stream=True),
loss_metrics(minibatch=True, epoch=True,
experience=True, stream=True),
forgetting_metrics(experience=True, stream=True),
timing_metrics(epoch=True, epoch_running=True),
cpu_usage_metrics(experience=True),
loggers=[interactive_logger],
collect_all=True
)
# 6. Naive strategy
strategy = Naive(
model=model,
optimizer=optimizer,
criterion=criterion,
train_mb_size=500,
train_epochs=1,
eval_mb_size=128,
evaluator=eval_plugin,
device=device
)
# 7. Train + evaluate
for experience in benchmark.train_stream:
print(f"\n Experience {experience.current_experience} training")
strategy.train(experience)
print("Evaluating on test stream")
strategy.eval(benchmark.test_stream)
if __name__ == "__main__":
main()