test_strategies.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464
  1. ################################################################################
  2. # Copyright (c) 2021 ContinualAI. #
  3. # Copyrights licensed under the MIT License. #
  4. # See the accompanying LICENSE file for terms. #
  5. # #
  6. # Date: 1-06-2020 #
  7. # Author(s): Andrea Cossu #
  8. # E-mail: contact@continualai.org #
  9. # Website: avalanche.continualai.org #
  10. ################################################################################
  11. import torch
  12. import unittest
  13. import os
  14. import sys
  15. from torch.optim import SGD
  16. from torch.nn import CrossEntropyLoss, Linear
  17. from avalanche.logging import TextLogger
  18. from avalanche.models import SimpleMLP
  19. from avalanche.training.plugins import EvaluationPlugin, StrategyPlugin, \
  20. LwFPlugin, ReplayPlugin
  21. from avalanche.training.strategies import Naive, Replay, CWRStar, \
  22. GDumb, LwF, AGEM, GEM, EWC, \
  23. SynapticIntelligence, JointTraining, CoPE, StreamingLDA, BaseStrategy
  24. from avalanche.training.strategies.cumulative import Cumulative
  25. from avalanche.training.strategies.joint_training import AlreadyTrainedError
  26. from avalanche.training.strategies.strategy_wrappers import PNNStrategy
  27. from avalanche.training.strategies.icarl import ICaRL
  28. from avalanche.training.utils import get_last_fc_layer
  29. from avalanche.evaluation.metrics import StreamAccuracy
  30. from tests.unit_tests_utils import get_fast_benchmark, get_device
  31. class BaseStrategyTest(unittest.TestCase):
  32. def test_periodic_eval(self):
  33. model = SimpleMLP(input_size=6, hidden_size=10)
  34. benchmark = get_fast_benchmark()
  35. optimizer = SGD(model.parameters(), lr=1e-3)
  36. criterion = CrossEntropyLoss()
  37. curve_key = 'Top1_Acc_Stream/eval_phase/train_stream/Task000'
  38. ###################
  39. # Case #1: No eval
  40. ###################
  41. # we use stream acc. because it emits a single value
  42. # for each eval loop.
  43. acc = StreamAccuracy()
  44. strategy = Naive(model, optimizer, criterion, train_epochs=2,
  45. eval_every=-1, evaluator=EvaluationPlugin(acc))
  46. strategy.train(benchmark.train_stream[0])
  47. # eval is not called in this case
  48. assert len(strategy.evaluator.get_all_metrics()) == 0
  49. ###################
  50. # Case #2: Eval at the end only and before training
  51. ###################
  52. acc = StreamAccuracy()
  53. strategy = Naive(model, optimizer, criterion, train_epochs=2,
  54. eval_every=0, evaluator=EvaluationPlugin(acc))
  55. strategy.train(benchmark.train_stream[0])
  56. # eval is called once at the end of the training loop
  57. curve = strategy.evaluator.get_all_metrics()[curve_key][1]
  58. assert len(curve) == 2
  59. ###################
  60. # Case #3: Eval after every epoch and before training
  61. ###################
  62. acc = StreamAccuracy()
  63. strategy = Naive(model, optimizer, criterion, train_epochs=2,
  64. eval_every=1, evaluator=EvaluationPlugin(acc))
  65. strategy.train(benchmark.train_stream[0])
  66. # eval is called after every epoch + the end of the training loop
  67. curve = strategy.evaluator.get_all_metrics()[curve_key][1]
  68. assert len(curve) == 4
  69. def test_forward_hooks(self):
  70. model = SimpleMLP(input_size=6, hidden_size=10)
  71. optimizer = SGD(model.parameters(), lr=1e-3)
  72. criterion = CrossEntropyLoss()
  73. strategy = Naive(model, optimizer, criterion,
  74. train_epochs=2, eval_every=0)
  75. was_hook_called = False
  76. def hook(a, b, c):
  77. nonlocal was_hook_called
  78. was_hook_called = True
  79. model.register_forward_hook(hook)
  80. mb_x = torch.randn(32, 6, device=strategy.device)
  81. strategy.mbatch = mb_x, None, None
  82. strategy.forward()
  83. assert was_hook_called
  84. def test_early_stop(self):
  85. class EarlyStopP(StrategyPlugin):
  86. def after_training_iteration(self, strategy: 'BaseStrategy',
  87. **kwargs):
  88. if strategy.mb_it == 10:
  89. strategy.stop_training()
  90. model = SimpleMLP(input_size=6, hidden_size=100)
  91. criterion = CrossEntropyLoss()
  92. optimizer = SGD(model.parameters(), lr=1)
  93. strategy = Cumulative(
  94. model, optimizer, criterion, train_mb_size=1, device=get_device(),
  95. eval_mb_size=512, train_epochs=1, evaluator=None,
  96. plugins=[EarlyStopP()])
  97. benchmark = get_fast_benchmark()
  98. for train_batch_info in benchmark.train_stream:
  99. strategy.train(train_batch_info)
  100. assert strategy.mb_it == 11
  101. class StrategyTest(unittest.TestCase):
  102. if "FAST_TEST" in os.environ:
  103. fast_test = os.environ['FAST_TEST'].lower() in ["true"]
  104. else:
  105. fast_test = False
  106. if "USE_GPU" in os.environ:
  107. use_gpu = os.environ['USE_GPU'].lower() in ["true"]
  108. else:
  109. use_gpu = False
  110. print("Fast Test:", fast_test)
  111. print("Test on GPU:", use_gpu)
  112. if use_gpu:
  113. device = "cuda"
  114. else:
  115. device = "cpu"
  116. def init_sit(self):
  117. model = self.get_model(fast_test=True)
  118. optimizer = SGD(model.parameters(), lr=1e-3)
  119. criterion = CrossEntropyLoss()
  120. benchmark = self.load_benchmark(use_task_labels=False)
  121. return model, optimizer, criterion, benchmark
  122. def test_naive(self):
  123. # SIT scenario
  124. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  125. strategy = Naive(model, optimizer, criterion, train_mb_size=64,
  126. device=self.device, eval_mb_size=50, train_epochs=2)
  127. self.run_strategy(my_nc_benchmark, strategy)
  128. # MT scenario
  129. strategy = Naive(model, optimizer, criterion, train_mb_size=64,
  130. device=self.device, eval_mb_size=50, train_epochs=2)
  131. benchmark = self.load_benchmark(use_task_labels=True)
  132. self.run_strategy(benchmark, strategy)
  133. def test_joint(self):
  134. class JointSTestPlugin(StrategyPlugin):
  135. def __init__(self, benchmark):
  136. super().__init__()
  137. self.benchmark = benchmark
  138. def after_train_dataset_adaptation(self, strategy: 'BaseStrategy',
  139. **kwargs):
  140. """
  141. Check that the dataset used for training contains the
  142. correct number of samples.
  143. """
  144. cum_len = sum([len(exp.dataset) for exp
  145. in self.benchmark.train_stream])
  146. assert len(strategy.adapted_dataset) == cum_len
  147. # SIT scenario
  148. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  149. strategy = JointTraining(model, optimizer, criterion, train_mb_size=64,
  150. device=self.device, eval_mb_size=50,
  151. train_epochs=2,
  152. plugins=[JointSTestPlugin(my_nc_benchmark)])
  153. strategy.evaluator.loggers = [TextLogger(sys.stdout)]
  154. strategy.train(my_nc_benchmark.train_stream)
  155. # MT scenario
  156. my_nc_benchmark = self.load_benchmark(use_task_labels=True)
  157. strategy = JointTraining(
  158. model, optimizer, criterion, train_mb_size=64,
  159. device=self.device, eval_mb_size=50, train_epochs=2,
  160. plugins=[JointSTestPlugin(my_nc_benchmark)])
  161. strategy.evaluator.loggers = [TextLogger(sys.stdout)]
  162. strategy.train(my_nc_benchmark.train_stream)
  163. # Raise error when retraining
  164. self.assertRaises(AlreadyTrainedError,
  165. lambda: strategy.train(my_nc_benchmark.train_stream))
  166. def test_cwrstar(self):
  167. # SIT scenario
  168. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  169. last_fc_name, _ = get_last_fc_layer(model)
  170. strategy = CWRStar(model, optimizer, criterion, last_fc_name,
  171. train_mb_size=64, device=self.device)
  172. self.run_strategy(my_nc_benchmark, strategy)
  173. # MT scenario
  174. strategy = CWRStar(model, optimizer, criterion, last_fc_name,
  175. train_mb_size=64, device=self.device)
  176. benchmark = self.load_benchmark(use_task_labels=True)
  177. self.run_strategy(benchmark, strategy)
  178. def test_replay(self):
  179. # SIT scenario
  180. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  181. strategy = Replay(model, optimizer, criterion,
  182. mem_size=10, train_mb_size=64, device=self.device,
  183. eval_mb_size=50, train_epochs=2)
  184. self.run_strategy(my_nc_benchmark, strategy)
  185. # MT scenario
  186. strategy = Replay(model, optimizer, criterion,
  187. mem_size=10, train_mb_size=64, device=self.device,
  188. eval_mb_size=50, train_epochs=2)
  189. benchmark = self.load_benchmark(use_task_labels=True)
  190. self.run_strategy(benchmark, strategy)
  191. def test_gdumb(self):
  192. # SIT scenario
  193. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  194. strategy = GDumb(
  195. model, optimizer, criterion,
  196. mem_size=200, train_mb_size=64, device=self.device,
  197. eval_mb_size=50, train_epochs=2
  198. )
  199. self.run_strategy(my_nc_benchmark, strategy)
  200. # MT scenario
  201. strategy = GDumb(
  202. model, optimizer, criterion,
  203. mem_size=200, train_mb_size=64, device=self.device,
  204. eval_mb_size=50, train_epochs=2
  205. )
  206. benchmark = self.load_benchmark(use_task_labels=True)
  207. self.run_strategy(benchmark, strategy)
  208. def test_cumulative(self):
  209. # SIT scenario
  210. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  211. strategy = Cumulative(model, optimizer, criterion, train_mb_size=64,
  212. device=self.device, eval_mb_size=50,
  213. train_epochs=2)
  214. self.run_strategy(my_nc_benchmark, strategy)
  215. # MT scenario
  216. strategy = Cumulative(model, optimizer, criterion, train_mb_size=64,
  217. device=self.device, eval_mb_size=50,
  218. train_epochs=2)
  219. benchmark = self.load_benchmark(use_task_labels=True)
  220. self.run_strategy(benchmark, strategy)
  221. def test_slda(self):
  222. model, _, criterion, my_nc_benchmark = self.init_sit()
  223. strategy = StreamingLDA(model, criterion, input_size=10,
  224. output_layer_name='features',
  225. num_classes=10, eval_mb_size=7,
  226. train_epochs=1, device=self.device,
  227. train_mb_size=7)
  228. self.run_strategy(my_nc_benchmark, strategy)
  229. def test_warning_slda_lwf(self):
  230. model, _, criterion, my_nc_benchmark = self.init_sit()
  231. with self.assertLogs('avalanche.training.strategies', "WARNING") as cm:
  232. StreamingLDA(model, criterion, input_size=10,
  233. output_layer_name='features', num_classes=10,
  234. plugins=[LwFPlugin(), ReplayPlugin()])
  235. self.assertEqual(1, len(cm.output))
  236. self.assertIn(
  237. "LwFPlugin seems to use the callback before_backward"
  238. " which is disabled by StreamingLDA",
  239. cm.output[0]
  240. )
  241. def test_lwf(self):
  242. # SIT scenario
  243. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  244. strategy = LwF(model, optimizer, criterion,
  245. alpha=[0, 1 / 2, 2 * (2 / 3), 3 * (3 / 4), 4 * (4 / 5)],
  246. temperature=2, device=self.device,
  247. train_mb_size=10, eval_mb_size=50,
  248. train_epochs=2)
  249. self.run_strategy(my_nc_benchmark, strategy)
  250. # MT scenario
  251. strategy = LwF(model, optimizer, criterion,
  252. alpha=[0, 1 / 2, 2 * (2 / 3), 3 * (3 / 4), 4 * (4 / 5)],
  253. temperature=2, device=self.device,
  254. train_mb_size=10, eval_mb_size=50,
  255. train_epochs=2)
  256. benchmark = self.load_benchmark(use_task_labels=True)
  257. self.run_strategy(benchmark, strategy)
  258. def test_agem(self):
  259. # SIT scenario
  260. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  261. strategy = AGEM(model, optimizer, criterion,
  262. patterns_per_exp=250, sample_size=256,
  263. train_mb_size=10, eval_mb_size=50,
  264. train_epochs=2)
  265. self.run_strategy(my_nc_benchmark, strategy)
  266. # MT scenario
  267. strategy = AGEM(model, optimizer, criterion,
  268. patterns_per_exp=250, sample_size=256,
  269. train_mb_size=10, eval_mb_size=50,
  270. train_epochs=2)
  271. benchmark = self.load_benchmark(use_task_labels=True)
  272. self.run_strategy(benchmark, strategy)
  273. def test_gem(self):
  274. # SIT scenario
  275. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  276. strategy = GEM(model, optimizer, criterion,
  277. patterns_per_exp=256,
  278. train_mb_size=10, eval_mb_size=50,
  279. train_epochs=2)
  280. self.run_strategy(my_nc_benchmark, strategy)
  281. # MT scenario
  282. strategy = GEM(model, optimizer, criterion,
  283. patterns_per_exp=256,
  284. train_mb_size=10, eval_mb_size=50,
  285. train_epochs=2)
  286. self.run_strategy(my_nc_benchmark, strategy)
  287. benchmark = self.load_benchmark(use_task_labels=True)
  288. self.run_strategy(benchmark, strategy)
  289. def test_ewc(self):
  290. # SIT scenario
  291. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  292. strategy = EWC(model, optimizer, criterion, ewc_lambda=0.4,
  293. mode='separate',
  294. train_mb_size=10, eval_mb_size=50,
  295. train_epochs=2)
  296. self.run_strategy(my_nc_benchmark, strategy)
  297. # MT scenario
  298. strategy = EWC(model, optimizer, criterion, ewc_lambda=0.4,
  299. mode='separate',
  300. train_mb_size=10, eval_mb_size=50,
  301. train_epochs=2)
  302. benchmark = self.load_benchmark(use_task_labels=True)
  303. self.run_strategy(benchmark, strategy)
  304. def test_ewc_online(self):
  305. # SIT scenario
  306. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  307. strategy = EWC(model, optimizer, criterion, ewc_lambda=0.4,
  308. mode='online', decay_factor=0.1,
  309. train_mb_size=10, eval_mb_size=50,
  310. train_epochs=2)
  311. self.run_strategy(my_nc_benchmark, strategy)
  312. # MT scenario
  313. strategy = EWC(model, optimizer, criterion, ewc_lambda=0.4,
  314. mode='online', decay_factor=0.1,
  315. train_mb_size=10, eval_mb_size=50,
  316. train_epochs=2)
  317. benchmark = self.load_benchmark(use_task_labels=True)
  318. self.run_strategy(benchmark, strategy)
  319. def test_synaptic_intelligence(self):
  320. # SIT scenario
  321. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  322. strategy = SynapticIntelligence(
  323. model, optimizer, criterion, si_lambda=0.0001,
  324. train_epochs=1, train_mb_size=10, eval_mb_size=10)
  325. benchmark = self.load_benchmark(use_task_labels=False)
  326. self.run_strategy(benchmark, strategy)
  327. # MT scenario
  328. strategy = SynapticIntelligence(
  329. model, optimizer, criterion, si_lambda=0.0001,
  330. train_epochs=1, train_mb_size=10, eval_mb_size=10)
  331. benchmark = self.load_benchmark(use_task_labels=True)
  332. self.run_strategy(benchmark, strategy)
  333. def test_cope(self):
  334. # Fast benchmark (hardcoded)
  335. n_classes = 10
  336. emb_size = n_classes # Embedding size
  337. # SIT scenario
  338. model, optimizer, criterion, my_nc_benchmark = self.init_sit()
  339. strategy = CoPE(model, optimizer, criterion,
  340. mem_size=10, n_classes=n_classes, p_size=emb_size,
  341. train_mb_size=10, device=self.device,
  342. eval_mb_size=50, train_epochs=2)
  343. self.run_strategy(my_nc_benchmark, strategy)
  344. # MT scenario
  345. strategy = CoPE(model, optimizer, criterion,
  346. mem_size=10, n_classes=n_classes, p_size=emb_size,
  347. train_mb_size=10, device=self.device,
  348. eval_mb_size=50, train_epochs=2)
  349. benchmark = self.load_benchmark(use_task_labels=True)
  350. self.run_strategy(benchmark, strategy)
  351. def test_pnn(self):
  352. # only multi-task scenarios.
  353. # eval on future tasks is not allowed.
  354. strategy = PNNStrategy(
  355. num_layers=3, in_features=6, hidden_features_per_column=10,
  356. lr=0.1, train_mb_size=10, device=self.device, eval_mb_size=50,
  357. train_epochs=2)
  358. # train and test loop
  359. benchmark = self.load_benchmark(use_task_labels=True)
  360. for train_task in benchmark.train_stream:
  361. strategy.train(train_task)
  362. strategy.eval(benchmark.test_stream)
  363. def test_icarl(self):
  364. model, optimizer, criterion, benchmark = self.init_sit()
  365. strategy = ICaRL(
  366. model.features, model.classifier, optimizer, 20,
  367. buffer_transform=None, criterion=criterion,
  368. fixed_memory=True, train_mb_size=10,
  369. train_epochs=2, eval_mb_size=50,
  370. device=self.device,)
  371. self.run_strategy(benchmark, strategy)
  372. def load_benchmark(self, use_task_labels=False):
  373. """
  374. Returns a NC benchmark from a fake dataset of 10 classes, 5 experiences,
  375. 2 classes per experience.
  376. :param fast_test: if True loads fake data, MNIST otherwise.
  377. """
  378. return get_fast_benchmark(use_task_labels=use_task_labels)
  379. def get_model(self, fast_test=False):
  380. if fast_test:
  381. return SimpleMLP(input_size=6, hidden_size=10)
  382. else:
  383. return SimpleMLP()
  384. def run_strategy(self, benchmark, cl_strategy):
  385. print('Starting experiment...')
  386. cl_strategy.evaluator.loggers = [TextLogger(sys.stdout)]
  387. results = []
  388. for train_batch_info in benchmark.train_stream:
  389. print("Start of experience ", train_batch_info.current_experience)
  390. cl_strategy.train(train_batch_info)
  391. print('Training completed')
  392. print('Computing accuracy on the current test set')
  393. results.append(cl_strategy.eval(benchmark.test_stream[:]))
  394. if __name__ == '__main__':
  395. unittest.main()