pnn.py 8.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240
  1. import torch
  2. import torch.nn.functional as F
  3. from torch import nn
  4. from avalanche.benchmarks.utils import AvalancheDataset
  5. from avalanche.benchmarks.utils.dataset_utils import ConstantSequence
  6. from avalanche.models import MultiTaskModule, DynamicModule
  7. from avalanche.models import MultiHeadClassifier
  8. class LinearAdapter(nn.Module):
  9. def __init__(self, in_features, out_features_per_column, num_prev_modules):
  10. """ Linear adapter for Progressive Neural Networks.
  11. :param in_features: size of each input sample
  12. :param out_features_per_column: size of each output sample
  13. :param num_prev_modules: number of previous modules
  14. """
  15. super().__init__()
  16. # Eq. 1 - lateral connections
  17. # one layer for each previous column. Empty for the first task.
  18. self.lat_layers = nn.ModuleList([])
  19. for _ in range(num_prev_modules):
  20. m = nn.Linear(in_features, out_features_per_column)
  21. self.lat_layers.append(m)
  22. def forward(self, x):
  23. assert len(x) == self.num_prev_modules
  24. hs = []
  25. for ii, lat in enumerate(self.lat_layers):
  26. hs.append(lat(x[ii]))
  27. return sum(hs)
  28. class MLPAdapter(nn.Module):
  29. def __init__(self, in_features, out_features_per_column, num_prev_modules,
  30. activation=F.relu):
  31. """ MLP adapter for Progressive Neural Networks.
  32. :param in_features: size of each input sample
  33. :param out_features_per_column: size of each output sample
  34. :param num_prev_modules: number of previous modules
  35. :param activation: activation function (default=ReLU)
  36. """
  37. super().__init__()
  38. self.num_prev_modules = num_prev_modules
  39. self.activation = activation
  40. if num_prev_modules == 0:
  41. return # first adapter is empty
  42. # Eq. 2 - MLP adapter. Not needed for the first task.
  43. self.V = nn.Linear(in_features * num_prev_modules,
  44. out_features_per_column)
  45. self.alphas = nn.Parameter(torch.randn(num_prev_modules))
  46. self.U = nn.Linear(out_features_per_column, out_features_per_column)
  47. def forward(self, x):
  48. if self.num_prev_modules == 0:
  49. return 0 # first adapter is empty
  50. assert len(x) == self.num_prev_modules
  51. assert len(x[0].shape) == 2, \
  52. "Inputs to MLPAdapter should have two dimensions: " \
  53. "<batch_size, num_features>."
  54. for i, el in enumerate(x):
  55. x[i] = self.alphas[i] * el
  56. x = torch.cat(x, dim=1)
  57. x = self.U(self.activation(self.V(x)))
  58. return x
  59. class PNNColumn(nn.Module):
  60. def __init__(self, in_features, out_features_per_column, num_prev_modules,
  61. adapter='mlp'):
  62. """ Progressive Neural Network column.
  63. :param in_features: size of each input sample
  64. :param out_features_per_column:
  65. size of each output sample (single column)
  66. :param num_prev_modules: number of previous columns
  67. :param adapter: adapter type. One of {'linear', 'mlp'} (default='mlp')
  68. """
  69. super().__init__()
  70. self.in_features = in_features
  71. self.out_features_per_column = out_features_per_column
  72. self.num_prev_modules = num_prev_modules
  73. self.itoh = nn.Linear(in_features, out_features_per_column)
  74. if adapter == 'linear':
  75. self.adapter = LinearAdapter(in_features, out_features_per_column,
  76. num_prev_modules)
  77. elif adapter == 'mlp':
  78. self.adapter = MLPAdapter(in_features, out_features_per_column,
  79. num_prev_modules)
  80. else:
  81. raise ValueError("`adapter` must be one of: {'mlp', `linear'}.")
  82. def freeze(self):
  83. for param in self.parameters():
  84. param.requires_grad = False
  85. def forward(self, x):
  86. prev_xs, last_x = x[:-1], x[-1]
  87. hs = self.adapter(prev_xs)
  88. hs += self.itoh(last_x)
  89. return hs
  90. class PNNLayer(MultiTaskModule, DynamicModule):
  91. def __init__(self, in_features, out_features_per_column, adapter='mlp'):
  92. """ Progressive Neural Network layer.
  93. The adaptation phase assumes that each experience is a separate task.
  94. Multiple experiences with the same task label or multiple task labels
  95. within the same experience will result in a runtime error.
  96. :param in_features: size of each input sample
  97. :param out_features_per_column:
  98. size of each output sample (single column)
  99. :param adapter: adapter type. One of {'linear', 'mlp'} (default='mlp')
  100. """
  101. super().__init__()
  102. self.in_features = in_features
  103. self.out_features_per_column = out_features_per_column
  104. self.adapter = adapter
  105. # convert from task label to module list order
  106. self.task_to_module_idx = {}
  107. first_col = PNNColumn(in_features, out_features_per_column,
  108. 0, adapter=adapter)
  109. self.columns = nn.ModuleList([first_col])
  110. @property
  111. def num_columns(self):
  112. return len(self.columns)
  113. def train_adaptation(self, dataset: AvalancheDataset):
  114. """ Training adaptation for PNN layer.
  115. Adds an additional column to the layer.
  116. :param dataset:
  117. :return:
  118. """
  119. task_labels = dataset.targets_task_labels
  120. if isinstance(task_labels, ConstantSequence):
  121. # task label is unique. Don't check duplicates.
  122. task_labels = [task_labels[0]]
  123. else:
  124. task_labels = set(task_labels)
  125. assert len(task_labels) == 1, \
  126. "PNN assumes a single task for each experience. Please use a " \
  127. "compatible benchmark."
  128. # extract task label from set
  129. task_label = next(iter(task_labels))
  130. assert task_label not in self.task_to_module_idx, \
  131. "A new experience is using a previously seen task label. This is " \
  132. "not compatible with PNN, which assumes different task labels for" \
  133. " each training experience."
  134. if len(self.task_to_module_idx) == 0:
  135. # we have already initialized the first column.
  136. # No need to call add_column here.
  137. self.task_to_module_idx[task_label] = 0
  138. else:
  139. self.task_to_module_idx[task_label] = self.num_columns
  140. self._add_column()
  141. def _add_column(self):
  142. """ Add a new column. """
  143. # Freeze old parameters
  144. for param in self.parameters():
  145. param.requires_grad = False
  146. self.columns.append(PNNColumn(self.in_features,
  147. self.out_features_per_column,
  148. self.num_columns,
  149. adapter=self.adapter))
  150. def forward_single_task(self, x, task_label):
  151. """ Forward.
  152. :param x: list of inputs.
  153. :param task_label:
  154. :return:
  155. """
  156. col_idx = self.task_to_module_idx[task_label]
  157. hs = []
  158. for ii in range(col_idx + 1):
  159. hs.append(self.columns[ii](x[:ii+1]))
  160. return hs
  161. class PNN(MultiTaskModule):
  162. def __init__(self, num_layers=1, in_features=784,
  163. hidden_features_per_column=100, adapter='mlp'):
  164. """ Progressive Neural Network.
  165. The model assumes that each experience is a separate task.
  166. Multiple experiences with the same task label or multiple task labels
  167. within the same experience will result in a runtime error.
  168. :param num_layers: number of layers (default=1)
  169. :param in_features: size of each input sample
  170. :param hidden_features_per_column:
  171. number of hidden units for each column
  172. :param adapter: adapter type. One of {'linear', 'mlp'} (default='mlp')
  173. """
  174. super().__init__()
  175. assert num_layers >= 1
  176. self.num_layers = num_layers
  177. self.in_features = in_features
  178. self.out_features_per_columns = hidden_features_per_column
  179. self.layers = nn.ModuleList()
  180. self.layers.append(PNNLayer(in_features, hidden_features_per_column))
  181. for _ in range(num_layers - 1):
  182. lay = PNNLayer(hidden_features_per_column,
  183. hidden_features_per_column,
  184. adapter=adapter)
  185. self.layers.append(lay)
  186. self.classifier = MultiHeadClassifier(hidden_features_per_column)
  187. def forward_single_task(self, x, task_label):
  188. """ Forward.
  189. :param x:
  190. :param task_label:
  191. :return:
  192. """
  193. x = x.contiguous()
  194. x = x.view(x.size(0), self.in_features)
  195. num_columns = self.layers[0].num_columns
  196. col_idx = self.layers[-1].task_to_module_idx[task_label]
  197. x = [x for _ in range(num_columns)]
  198. for lay in self.layers:
  199. x = [F.relu(el) for el in lay(x, task_label)]
  200. return self.classifier(x[col_idx], task_label)